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.
- {faster_diffbloch-0.1.6 → faster_diffbloch-0.1.7}/PKG-INFO +25 -27
- {faster_diffbloch-0.1.6 → faster_diffbloch-0.1.7}/README.md +24 -26
- {faster_diffbloch-0.1.6 → faster_diffbloch-0.1.7}/pyproject.toml +1 -1
- {faster_diffbloch-0.1.6 → faster_diffbloch-0.1.7}/src/faster_diffbloch/__init__.py +1 -1
- {faster_diffbloch-0.1.6 → faster_diffbloch-0.1.7}/src/faster_diffbloch/backend.py +30 -0
- faster_diffbloch-0.1.7/src/faster_diffbloch/scattering.py +271 -0
- {faster_diffbloch-0.1.6 → faster_diffbloch-0.1.7}/.gitignore +0 -0
- {faster_diffbloch-0.1.6 → faster_diffbloch-0.1.7}/LICENSE +0 -0
- {faster_diffbloch-0.1.6 → faster_diffbloch-0.1.7}/src/faster_diffbloch/builder.py +0 -0
- {faster_diffbloch-0.1.6 → faster_diffbloch-0.1.7}/src/faster_diffbloch/cli.py +0 -0
- {faster_diffbloch-0.1.6 → faster_diffbloch-0.1.7}/src/faster_diffbloch/native/batch_cgemm.c +0 -0
- {faster_diffbloch-0.1.6 → faster_diffbloch-0.1.7}/src/faster_diffbloch/native/batch_cgemm.h +0 -0
- {faster_diffbloch-0.1.6 → faster_diffbloch-0.1.7}/src/faster_diffbloch/native/batch_cgemm.metal +0 -0
- {faster_diffbloch-0.1.6 → faster_diffbloch-0.1.7}/src/faster_diffbloch/native/bridge_lib.c +0 -0
- {faster_diffbloch-0.1.6 → faster_diffbloch-0.1.7}/src/faster_diffbloch/native/metal_batch_cgemm.m +0 -0
- {faster_diffbloch-0.1.6 → faster_diffbloch-0.1.7}/src/faster_diffbloch/native/native_scattering.c +0 -0
- {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.
|
|
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
|
-
**
|
|
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.
|
|
88
|
-
| 61 | 2.
|
|
89
|
-
| 91 | 3.
|
|
90
|
-
| 163 | 10.
|
|
91
|
-
| 579 |
|
|
92
|
-
|
|
93
|
-
|
|
94
|
-
|
|
95
|
-
|
|
96
|
-
|
|
97
|
-
|
|
98
|
-
|
|
99
|
-
|
|
100
|
-
|
|
101
|
-
|
|
102
|
-
|
|
103
|
-
|
|
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
|
|
113
|
-
|
|
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
|
-
|
|
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
|
-
**
|
|
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.
|
|
55
|
-
| 61 | 2.
|
|
56
|
-
| 91 | 3.
|
|
57
|
-
| 163 | 10.
|
|
58
|
-
| 579 |
|
|
59
|
-
|
|
60
|
-
|
|
61
|
-
|
|
62
|
-
|
|
63
|
-
|
|
64
|
-
|
|
65
|
-
|
|
66
|
-
|
|
67
|
-
|
|
68
|
-
|
|
69
|
-
|
|
70
|
-
|
|
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
|
|
80
|
-
|
|
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
|
-
|
|
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.
|
|
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"
|
|
@@ -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
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
{faster_diffbloch-0.1.6 → faster_diffbloch-0.1.7}/src/faster_diffbloch/native/batch_cgemm.metal
RENAMED
|
File without changes
|
|
File without changes
|
{faster_diffbloch-0.1.6 → faster_diffbloch-0.1.7}/src/faster_diffbloch/native/metal_batch_cgemm.m
RENAMED
|
File without changes
|
{faster_diffbloch-0.1.6 → faster_diffbloch-0.1.7}/src/faster_diffbloch/native/native_scattering.c
RENAMED
|
File without changes
|
{faster_diffbloch-0.1.6 → faster_diffbloch-0.1.7}/src/faster_diffbloch/native/native_scattering.h
RENAMED
|
File without changes
|