faster-diffbloch 0.1.5__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.5 → faster_diffbloch-0.1.7}/PKG-INFO +29 -36
  2. {faster_diffbloch-0.1.5 → faster_diffbloch-0.1.7}/README.md +28 -35
  3. {faster_diffbloch-0.1.5 → faster_diffbloch-0.1.7}/pyproject.toml +1 -1
  4. {faster_diffbloch-0.1.5 → faster_diffbloch-0.1.7}/src/faster_diffbloch/__init__.py +1 -1
  5. {faster_diffbloch-0.1.5 → 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.5 → faster_diffbloch-0.1.7}/.gitignore +0 -0
  8. {faster_diffbloch-0.1.5 → faster_diffbloch-0.1.7}/LICENSE +0 -0
  9. {faster_diffbloch-0.1.5 → faster_diffbloch-0.1.7}/src/faster_diffbloch/builder.py +0 -0
  10. {faster_diffbloch-0.1.5 → faster_diffbloch-0.1.7}/src/faster_diffbloch/cli.py +0 -0
  11. {faster_diffbloch-0.1.5 → faster_diffbloch-0.1.7}/src/faster_diffbloch/native/batch_cgemm.c +0 -0
  12. {faster_diffbloch-0.1.5 → faster_diffbloch-0.1.7}/src/faster_diffbloch/native/batch_cgemm.h +0 -0
  13. {faster_diffbloch-0.1.5 → faster_diffbloch-0.1.7}/src/faster_diffbloch/native/batch_cgemm.metal +0 -0
  14. {faster_diffbloch-0.1.5 → faster_diffbloch-0.1.7}/src/faster_diffbloch/native/bridge_lib.c +0 -0
  15. {faster_diffbloch-0.1.5 → faster_diffbloch-0.1.7}/src/faster_diffbloch/native/metal_batch_cgemm.m +0 -0
  16. {faster_diffbloch-0.1.5 → faster_diffbloch-0.1.7}/src/faster_diffbloch/native/native_scattering.c +0 -0
  17. {faster_diffbloch-0.1.5 → 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.5
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
- A Metal GPU path for Apple Silicon and an optimized CPU path for macOS and Linux, measured against diffBloch running on PyTorch.
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
  | :--- | :--- | :--- | :--- |
@@ -78,28 +78,37 @@ example documents.
78
78
 
79
79
  ## Performance
80
80
 
81
- Forward-plus-backward timing on an Apple M4 Max, single-threaded PyTorch,
82
- running the same diffBloch workload through `bench/compare_faster.py`. This is
83
- what `pip install faster-diffbloch` gives you: the accelerated matrix
84
- exponential bridged into diffBloch, with PyTorch handling everything else.
81
+ Forward-plus-backward timing on an Apple M4 Max, single-threaded PyTorch, over
82
+ the same diffBloch quartz workload at five beam counts. Both device paths, best
83
+ of three runs each.
85
84
 
86
- | Beams | faster-diffbloch | PyTorch | Speedup |
87
- | :---: | :---: | :---: | :---: |
88
- | 31 | 1.48 ms | 1.54 ms | 1.05x |
89
- | 61 | 1.93 ms | 2.07 ms | 1.07x |
90
- | 91 | 3.48 ms | 3.69 ms | 1.06x |
91
- | 163 | 8.76 ms | 10.14 ms | 1.16x |
85
+ | Beams | PyTorch | CPU path | Metal path | CPU | Metal |
86
+ | :---: | :---: | :---: | :---: | :---: | :---: |
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
92
 
93
- Best of three alternating runs per mode. The gain grows with beam count,
94
- because the matrix exponential takes a larger share of the work as the system
95
- grows and the fixed ctypes and PyTorch overhead matters less.
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.
96
105
 
97
106
  ### The standalone Flow port
98
107
 
99
108
  The diffFlow repository also holds a standalone port of the whole calculation,
100
- compiled from Flow rather than bridged into PyTorch. It avoids the framework
101
- overhead entirely and is considerably faster. At 579 beams, forward plus
102
- backward, minimum of five 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:
103
112
 
104
113
  | Implementation | Forward + Backward |
105
114
  | :--- | :---: |
@@ -108,29 +117,13 @@ backward, minimum of five runs:
108
117
  | Flow port, C backend | 82.84 ms |
109
118
  | Flow port, MLIR backend | 79.74 ms |
110
119
 
111
- Those numbers are not what this package delivers. They are reachable by
112
- running the port directly, and they are why the package exists. 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
113
122
  [diffFlow documentation](https://godofecht.github.io/diffFlow/).
114
123
 
115
124
  These figures come from one workload on one machine. Other crystals, hardware
116
125
  configurations and beam counts will differ.
117
126
 
118
- ### Metal GPU
119
-
120
- The Metal path is available on Apple Silicon through `enable(device="gpu")`.
121
- Correctness on this path is checked: the quartz reproduction above runs through
122
- it. The timings below come from an earlier run and have not been reproduced
123
- against the current package, so treat those as indicative.
124
-
125
- ### Package benchmark snapshot
126
-
127
- | Implementation | Forward | Forward + Backward | Speedup vs PyTorch CPU | Speedup vs PyTorch MPS |
128
- | :--- | :---: | :---: | :---: | :---: |
129
- | PyTorch CPU | 25.7 ms | 130.3 ms | 1.00x | 1.17x |
130
- | PyTorch MPS (fallback) | 26.2 ms | 153.0 ms | 0.85x | 1.00x |
131
- | **faster-diffBloch CPU** | **24.4 ms** | **83.5 ms** | **1.56x** | **1.83x** |
132
- | **faster-diffBloch Metal GPU** | **13.1 ms** | **58.1 ms** | **2.24x** | **2.63x** |
133
-
134
127
  ---
135
128
 
136
129
  ## Installation
@@ -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
- A Metal GPU path for Apple Silicon and an optimized CPU path for macOS and Linux, measured against diffBloch running on PyTorch.
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
  | :--- | :--- | :--- | :--- |
@@ -45,28 +45,37 @@ example documents.
45
45
 
46
46
  ## Performance
47
47
 
48
- Forward-plus-backward timing on an Apple M4 Max, single-threaded PyTorch,
49
- running the same diffBloch workload through `bench/compare_faster.py`. This is
50
- what `pip install faster-diffbloch` gives you: the accelerated matrix
51
- exponential bridged into diffBloch, with PyTorch handling everything else.
48
+ Forward-plus-backward timing on an Apple M4 Max, single-threaded PyTorch, over
49
+ the same diffBloch quartz workload at five beam counts. Both device paths, best
50
+ of three runs each.
52
51
 
53
- | Beams | faster-diffbloch | PyTorch | Speedup |
54
- | :---: | :---: | :---: | :---: |
55
- | 31 | 1.48 ms | 1.54 ms | 1.05x |
56
- | 61 | 1.93 ms | 2.07 ms | 1.07x |
57
- | 91 | 3.48 ms | 3.69 ms | 1.06x |
58
- | 163 | 8.76 ms | 10.14 ms | 1.16x |
52
+ | Beams | PyTorch | CPU path | Metal path | CPU | Metal |
53
+ | :---: | :---: | :---: | :---: | :---: | :---: |
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
59
 
60
- Best of three alternating runs per mode. The gain grows with beam count,
61
- because the matrix exponential takes a larger share of the work as the system
62
- grows and the fixed ctypes and PyTorch overhead matters less.
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.
63
72
 
64
73
  ### The standalone Flow port
65
74
 
66
75
  The diffFlow repository also holds a standalone port of the whole calculation,
67
- compiled from Flow rather than bridged into PyTorch. It avoids the framework
68
- overhead entirely and is considerably faster. At 579 beams, forward plus
69
- backward, minimum of five 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:
70
79
 
71
80
  | Implementation | Forward + Backward |
72
81
  | :--- | :---: |
@@ -75,29 +84,13 @@ backward, minimum of five runs:
75
84
  | Flow port, C backend | 82.84 ms |
76
85
  | Flow port, MLIR backend | 79.74 ms |
77
86
 
78
- Those numbers are not what this package delivers. They are reachable by
79
- running the port directly, and they are why the package exists. 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
80
89
  [diffFlow documentation](https://godofecht.github.io/diffFlow/).
81
90
 
82
91
  These figures come from one workload on one machine. Other crystals, hardware
83
92
  configurations and beam counts will differ.
84
93
 
85
- ### Metal GPU
86
-
87
- The Metal path is available on Apple Silicon through `enable(device="gpu")`.
88
- Correctness on this path is checked: the quartz reproduction above runs through
89
- it. The timings below come from an earlier run and have not been reproduced
90
- against the current package, so treat those as indicative.
91
-
92
- ### Package benchmark snapshot
93
-
94
- | Implementation | Forward | Forward + Backward | Speedup vs PyTorch CPU | Speedup vs PyTorch MPS |
95
- | :--- | :---: | :---: | :---: | :---: |
96
- | PyTorch CPU | 25.7 ms | 130.3 ms | 1.00x | 1.17x |
97
- | PyTorch MPS (fallback) | 26.2 ms | 153.0 ms | 0.85x | 1.00x |
98
- | **faster-diffBloch CPU** | **24.4 ms** | **83.5 ms** | **1.56x** | **1.83x** |
99
- | **faster-diffBloch Metal GPU** | **13.1 ms** | **58.1 ms** | **2.24x** | **2.63x** |
100
-
101
94
  ---
102
95
 
103
96
  ## Installation
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
4
4
 
5
5
  [project]
6
6
  name = "faster-diffbloch"
7
- version = "0.1.5"
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.5"
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