StochasticForceInference 2.0.0__py3-none-any.whl
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.
- SFI/__init__.py +64 -0
- SFI/bases/__init__.py +85 -0
- SFI/bases/constants.py +492 -0
- SFI/bases/linear.py +325 -0
- SFI/bases/monomials.py +218 -0
- SFI/bases/pairs.py +998 -0
- SFI/bases/spde.py +1537 -0
- SFI/diagnostics/__init__.py +60 -0
- SFI/diagnostics/assess.py +87 -0
- SFI/diagnostics/dynamics_order.py +621 -0
- SFI/diagnostics/plotting.py +226 -0
- SFI/diagnostics/report.py +238 -0
- SFI/diagnostics/residual_tests.py +395 -0
- SFI/diagnostics/residuals.py +688 -0
- SFI/inference/__init__.py +58 -0
- SFI/inference/base.py +1460 -0
- SFI/inference/optimizers.py +200 -0
- SFI/inference/overdamped.py +1214 -0
- SFI/inference/parametric_core/__init__.py +34 -0
- SFI/inference/parametric_core/covariance.py +232 -0
- SFI/inference/parametric_core/driver.py +272 -0
- SFI/inference/parametric_core/flow.py +149 -0
- SFI/inference/parametric_core/flow_multi.py +362 -0
- SFI/inference/parametric_core/flow_ud.py +168 -0
- SFI/inference/parametric_core/jacobians.py +540 -0
- SFI/inference/parametric_core/objective.py +286 -0
- SFI/inference/parametric_core/objective_ud.py +253 -0
- SFI/inference/parametric_core/precision.py +229 -0
- SFI/inference/parametric_core/solve.py +763 -0
- SFI/inference/result.py +362 -0
- SFI/inference/serialization.py +245 -0
- SFI/inference/sparse/__init__.py +67 -0
- SFI/inference/sparse/base.py +43 -0
- SFI/inference/sparse/beam.py +303 -0
- SFI/inference/sparse/greedy.py +151 -0
- SFI/inference/sparse/hillclimb.py +307 -0
- SFI/inference/sparse/lasso.py +178 -0
- SFI/inference/sparse/metrics.py +78 -0
- SFI/inference/sparse/result.py +278 -0
- SFI/inference/sparse/scorer.py +323 -0
- SFI/inference/sparse/stlsq.py +165 -0
- SFI/inference/sparsity.py +40 -0
- SFI/inference/underdamped.py +1355 -0
- SFI/integrate/__init__.py +38 -0
- SFI/integrate/api.py +920 -0
- SFI/integrate/integrand.py +402 -0
- SFI/integrate/rk4.py +156 -0
- SFI/integrate/timeops.py +174 -0
- SFI/langevin/__init__.py +29 -0
- SFI/langevin/base.py +863 -0
- SFI/langevin/chunked.py +225 -0
- SFI/langevin/noise.py +446 -0
- SFI/langevin/overdamped.py +560 -0
- SFI/langevin/underdamped.py +448 -0
- SFI/statefunc/__init__.py +49 -0
- SFI/statefunc/basis.py +87 -0
- SFI/statefunc/core/runtime.py +47 -0
- SFI/statefunc/factory.py +346 -0
- SFI/statefunc/interactor.py +90 -0
- SFI/statefunc/layout/__init__.py +29 -0
- SFI/statefunc/layout/_base.py +186 -0
- SFI/statefunc/layout/_eval_compiler.py +453 -0
- SFI/statefunc/layout/_fd_atoms.py +196 -0
- SFI/statefunc/layout/_grid.py +878 -0
- SFI/statefunc/layout/_sectors.py +166 -0
- SFI/statefunc/memhint.py +222 -0
- SFI/statefunc/nodes/__init__.py +56 -0
- SFI/statefunc/nodes/base.py +318 -0
- SFI/statefunc/nodes/contract.py +422 -0
- SFI/statefunc/nodes/interactions/__init__.py +23 -0
- SFI/statefunc/nodes/interactions/dispatcher.py +1365 -0
- SFI/statefunc/nodes/interactions/prepare.py +161 -0
- SFI/statefunc/nodes/interactions/specs.py +362 -0
- SFI/statefunc/nodes/interactions/stencils.py +718 -0
- SFI/statefunc/nodes/leaf.py +530 -0
- SFI/statefunc/nodes/ops/__init__.py +27 -0
- SFI/statefunc/nodes/ops/concat.py +28 -0
- SFI/statefunc/nodes/ops/derivative.py +447 -0
- SFI/statefunc/nodes/ops/einsum.py +120 -0
- SFI/statefunc/nodes/ops/linear.py +153 -0
- SFI/statefunc/nodes/ops/mapn.py +138 -0
- SFI/statefunc/nodes/ops/reshape_rank.py +158 -0
- SFI/statefunc/nodes/ops/slice.py +83 -0
- SFI/statefunc/params.py +309 -0
- SFI/statefunc/psf.py +112 -0
- SFI/statefunc/sf.py +104 -0
- SFI/statefunc/stateexpr.py +1243 -0
- SFI/statefunc/structexpr.py +1003 -0
- SFI/trajectory/__init__.py +31 -0
- SFI/trajectory/collection.py +1122 -0
- SFI/trajectory/dataset.py +1164 -0
- SFI/trajectory/degrade.py +1108 -0
- SFI/trajectory/io.py +1027 -0
- SFI/trajectory/reserved_extras.py +129 -0
- SFI/utils/__init__.py +17 -0
- SFI/utils/formatting.py +308 -0
- SFI/utils/maths.py +185 -0
- SFI/utils/neighbors.py +162 -0
- SFI/utils/plotting.py +1936 -0
- stochasticforceinference-2.0.0.dist-info/METADATA +203 -0
- stochasticforceinference-2.0.0.dist-info/RECORD +104 -0
- stochasticforceinference-2.0.0.dist-info/WHEEL +5 -0
- stochasticforceinference-2.0.0.dist-info/licenses/LICENSE +21 -0
- stochasticforceinference-2.0.0.dist-info/top_level.txt +1 -0
SFI/utils/neighbors.py
ADDED
|
@@ -0,0 +1,162 @@
|
|
|
1
|
+
"""Cell-list neighbor builder for truncated-range pair interactions.
|
|
2
|
+
|
|
3
|
+
Provides :func:`build_neighbor_csr` which constructs a sparse CSR
|
|
4
|
+
neighbor list from particle positions and a cutoff radius, using
|
|
5
|
+
``scipy.spatial.cKDTree``. The returned ``(indptr, indices)`` arrays
|
|
6
|
+
plug directly into ``dispatch_pairs_from_extras``.
|
|
7
|
+
|
|
8
|
+
All routines run on the host (pure NumPy) and are meant to be called
|
|
9
|
+
*between* JIT-compiled simulation chunks, not inside ``jax.lax.scan``.
|
|
10
|
+
"""
|
|
11
|
+
|
|
12
|
+
from __future__ import annotations
|
|
13
|
+
|
|
14
|
+
from typing import Optional, Tuple
|
|
15
|
+
|
|
16
|
+
import jax.numpy as jnp
|
|
17
|
+
import numpy as np
|
|
18
|
+
from scipy.spatial import cKDTree
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
def build_neighbor_csr(
|
|
22
|
+
positions: np.ndarray,
|
|
23
|
+
cutoff: float,
|
|
24
|
+
box: Optional[np.ndarray] = None,
|
|
25
|
+
*,
|
|
26
|
+
exclude_self: bool = True,
|
|
27
|
+
) -> Tuple[np.ndarray, np.ndarray]:
|
|
28
|
+
"""Build a CSR neighbor list using ``scipy.spatial.cKDTree``.
|
|
29
|
+
|
|
30
|
+
Parameters
|
|
31
|
+
----------
|
|
32
|
+
positions : ndarray, shape ``(N, d)``
|
|
33
|
+
Particle positions (spatial coordinates only).
|
|
34
|
+
cutoff : float
|
|
35
|
+
Cutoff radius. Pairs with ``r_ij > cutoff`` are excluded.
|
|
36
|
+
box : ndarray, shape ``(d,)``, optional
|
|
37
|
+
Periodic box lengths. If *None*, open (non-periodic) boundaries.
|
|
38
|
+
exclude_self : bool
|
|
39
|
+
If *True* (default), self-pairs ``(i, i)`` are never included.
|
|
40
|
+
|
|
41
|
+
Returns
|
|
42
|
+
-------
|
|
43
|
+
indptr : ndarray, shape ``(N + 1,)``, dtype int32
|
|
44
|
+
CSR row pointers.
|
|
45
|
+
indices : ndarray, shape ``(nnz,)``, dtype int32
|
|
46
|
+
CSR column indices (neighbour particle indices).
|
|
47
|
+
"""
|
|
48
|
+
positions = np.asarray(positions, dtype=np.float64)
|
|
49
|
+
N, d = positions.shape
|
|
50
|
+
|
|
51
|
+
if N == 0:
|
|
52
|
+
return (
|
|
53
|
+
np.zeros(1, dtype=np.int32),
|
|
54
|
+
np.empty(0, dtype=np.int32),
|
|
55
|
+
)
|
|
56
|
+
|
|
57
|
+
# --- wrap positions into the primary box ---
|
|
58
|
+
if box is not None:
|
|
59
|
+
box = np.asarray(box, dtype=np.float64)
|
|
60
|
+
positions = positions % box
|
|
61
|
+
|
|
62
|
+
# --- build KD-tree and query pairs ---
|
|
63
|
+
boxsize = box if box is not None else None
|
|
64
|
+
tree = cKDTree(positions, boxsize=boxsize)
|
|
65
|
+
csr = tree.sparse_distance_matrix(tree, cutoff, output_type="coo_matrix")
|
|
66
|
+
csr = csr.tocsr()
|
|
67
|
+
|
|
68
|
+
if exclude_self:
|
|
69
|
+
csr.setdiag(0)
|
|
70
|
+
csr.eliminate_zeros()
|
|
71
|
+
|
|
72
|
+
indptr = csr.indptr.astype(np.int32)
|
|
73
|
+
indices = csr.indices.astype(np.int32)
|
|
74
|
+
|
|
75
|
+
return indptr, indices
|
|
76
|
+
|
|
77
|
+
|
|
78
|
+
def make_neighbor_extras(
|
|
79
|
+
positions: np.ndarray,
|
|
80
|
+
cutoff: float,
|
|
81
|
+
box: Optional[np.ndarray] = None,
|
|
82
|
+
*,
|
|
83
|
+
indptr_key: str = "indptr",
|
|
84
|
+
indices_key: str = "indices",
|
|
85
|
+
exclude_self: bool = True,
|
|
86
|
+
) -> dict:
|
|
87
|
+
"""Build a CSR neighbor list and return it as an extras dict.
|
|
88
|
+
|
|
89
|
+
Convenience wrapper around :func:`build_neighbor_csr`. The returned
|
|
90
|
+
dict is ready to be merged into ``extras_global`` for a process that
|
|
91
|
+
uses ``dispatch_pairs_from_extras(indptr_key, indices_key)``.
|
|
92
|
+
|
|
93
|
+
Parameters
|
|
94
|
+
----------
|
|
95
|
+
positions, cutoff, box, exclude_self
|
|
96
|
+
Forwarded to :func:`build_neighbor_csr`.
|
|
97
|
+
indptr_key, indices_key
|
|
98
|
+
Keys under which CSR arrays are stored.
|
|
99
|
+
|
|
100
|
+
Returns
|
|
101
|
+
-------
|
|
102
|
+
dict
|
|
103
|
+
``{indptr_key: indptr, indices_key: indices}``
|
|
104
|
+
"""
|
|
105
|
+
indptr, indices = build_neighbor_csr(
|
|
106
|
+
positions,
|
|
107
|
+
cutoff,
|
|
108
|
+
box,
|
|
109
|
+
exclude_self=exclude_self,
|
|
110
|
+
)
|
|
111
|
+
return {
|
|
112
|
+
indptr_key: jnp.array(indptr),
|
|
113
|
+
indices_key: jnp.array(indices),
|
|
114
|
+
}
|
|
115
|
+
|
|
116
|
+
|
|
117
|
+
def pad_neighbor_csr(
|
|
118
|
+
indptr: np.ndarray,
|
|
119
|
+
indices: np.ndarray,
|
|
120
|
+
max_nnz: int,
|
|
121
|
+
*,
|
|
122
|
+
fill_index: int = 0,
|
|
123
|
+
) -> Tuple[np.ndarray, np.ndarray]:
|
|
124
|
+
"""Pad a CSR neighbor list to a fixed ``max_nnz``.
|
|
125
|
+
|
|
126
|
+
JAX JIT recompiles when array shapes change. Padding the indices
|
|
127
|
+
array to a fixed length avoids recompilation across simulation
|
|
128
|
+
chunks with fluctuating neighbor counts.
|
|
129
|
+
|
|
130
|
+
Excess entries are filled with ``fill_index`` (default 0). Because
|
|
131
|
+
``indptr`` keeps the true lengths, the dispatcher will only iterate
|
|
132
|
+
over the real neighbours — the padded entries are never evaluated.
|
|
133
|
+
|
|
134
|
+
.. note::
|
|
135
|
+
This only pads ``indices``. ``indptr`` is left unchanged (always
|
|
136
|
+
``N + 1`` long). If the actual nnz exceeds ``max_nnz``, a
|
|
137
|
+
``ValueError`` is raised.
|
|
138
|
+
|
|
139
|
+
Parameters
|
|
140
|
+
----------
|
|
141
|
+
indptr, indices
|
|
142
|
+
As returned by :func:`build_neighbor_csr`.
|
|
143
|
+
max_nnz : int
|
|
144
|
+
Target length for ``indices``.
|
|
145
|
+
fill_index : int
|
|
146
|
+
Index used to fill padded entries.
|
|
147
|
+
|
|
148
|
+
Returns
|
|
149
|
+
-------
|
|
150
|
+
indptr, indices_padded
|
|
151
|
+
Same ``indptr``, padded ``indices`` of length ``max_nnz``.
|
|
152
|
+
"""
|
|
153
|
+
nnz = len(indices)
|
|
154
|
+
if nnz > max_nnz:
|
|
155
|
+
raise ValueError(
|
|
156
|
+
f"Actual nnz ({nnz}) exceeds max_nnz ({max_nnz}). Increase max_nnz or enlarge the cutoff safety margin."
|
|
157
|
+
)
|
|
158
|
+
if nnz == max_nnz:
|
|
159
|
+
return indptr, indices
|
|
160
|
+
padded = np.full(max_nnz, fill_index, dtype=indices.dtype)
|
|
161
|
+
padded[:nnz] = indices
|
|
162
|
+
return indptr, padded
|