extract-clusters-step 2026.9.17__py2.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.
- extract_clusters_step/__init__.py +37 -0
- extract_clusters_step/_version.py +21 -0
- extract_clusters_step/cluster_sampling.py +638 -0
- extract_clusters_step/data/properties.csv +8 -0
- extract_clusters_step/data/references.bib +40 -0
- extract_clusters_step/extract_clusters.py +508 -0
- extract_clusters_step/extract_clusters_parameters.py +280 -0
- extract_clusters_step/extract_clusters_step.py +100 -0
- extract_clusters_step/metadata.py +12 -0
- extract_clusters_step/tk_extract_clusters.py +345 -0
- extract_clusters_step-2026.9.17.dist-info/METADATA +142 -0
- extract_clusters_step-2026.9.17.dist-info/RECORD +18 -0
- extract_clusters_step-2026.9.17.dist-info/WHEEL +6 -0
- extract_clusters_step-2026.9.17.dist-info/entry_points.txt +5 -0
- extract_clusters_step-2026.9.17.dist-info/licenses/AUTHORS.rst +5 -0
- extract_clusters_step-2026.9.17.dist-info/licenses/LICENSE +31 -0
- extract_clusters_step-2026.9.17.dist-info/top_level.txt +1 -0
- extract_clusters_step-2026.9.17.dist-info/zip-safe +1 -0
|
@@ -0,0 +1,37 @@
|
|
|
1
|
+
# -*- coding: utf-8 -*-
|
|
2
|
+
|
|
3
|
+
"""
|
|
4
|
+
extract_clusters_step
|
|
5
|
+
A SEAMM plug-in for extracting molecular clusters from a periodic cell of molecules
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
# Bring up the classes so that they appear to be directly in
|
|
9
|
+
# the extract_clusters_step package.
|
|
10
|
+
|
|
11
|
+
from extract_clusters_step.extract_clusters import ExtractClusters # noqa: F401, E501
|
|
12
|
+
from extract_clusters_step.extract_clusters_parameters import ( # noqa: F401
|
|
13
|
+
ExtractClustersParameters,
|
|
14
|
+
)
|
|
15
|
+
from extract_clusters_step.extract_clusters_step import ( # noqa: F401
|
|
16
|
+
ExtractClustersStep,
|
|
17
|
+
)
|
|
18
|
+
from extract_clusters_step.tk_extract_clusters import ( # noqa: F401
|
|
19
|
+
TkExtractClusters,
|
|
20
|
+
)
|
|
21
|
+
from extract_clusters_step.metadata import metadata # noqa: F401
|
|
22
|
+
from extract_clusters_step.cluster_sampling import ( # noqa: F401
|
|
23
|
+
PROPERTY_TAG,
|
|
24
|
+
classify_motif,
|
|
25
|
+
cluster_summary,
|
|
26
|
+
extract_nmers,
|
|
27
|
+
)
|
|
28
|
+
|
|
29
|
+
# Handle versioneer
|
|
30
|
+
from ._version import get_versions
|
|
31
|
+
|
|
32
|
+
__author__ = "Paul Saxe"
|
|
33
|
+
__email__ = "psaxe@molssi.org"
|
|
34
|
+
versions = get_versions()
|
|
35
|
+
__version__ = versions["version"]
|
|
36
|
+
__git_revision__ = versions["full-revisionid"]
|
|
37
|
+
del get_versions, versions
|
|
@@ -0,0 +1,21 @@
|
|
|
1
|
+
|
|
2
|
+
# This file was generated by 'versioneer.py' (0.18) from
|
|
3
|
+
# revision-control system data, or from the parent directory name of an
|
|
4
|
+
# unpacked source archive. Distribution tarballs contain a pre-generated copy
|
|
5
|
+
# of this file.
|
|
6
|
+
|
|
7
|
+
import json
|
|
8
|
+
|
|
9
|
+
version_json = '''
|
|
10
|
+
{
|
|
11
|
+
"date": "2026-09-17T16:12:40-0400",
|
|
12
|
+
"dirty": false,
|
|
13
|
+
"error": null,
|
|
14
|
+
"full-revisionid": "4e45e4ac9aade4248cb088e555c6d08dfe70c735",
|
|
15
|
+
"version": "2026.9.17"
|
|
16
|
+
}
|
|
17
|
+
''' # END VERSION_JSON
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
def get_versions():
|
|
21
|
+
return json.loads(version_json)
|
|
@@ -0,0 +1,638 @@
|
|
|
1
|
+
# -*- coding: utf-8 -*-
|
|
2
|
+
|
|
3
|
+
"""Sampling of n-molecule clusters from a (periodic) configuration.
|
|
4
|
+
|
|
5
|
+
The pure functions here have no dependence on the SEAMM framework, only on
|
|
6
|
+
numpy/scipy and the duck-typed ``molsystem`` configuration passed in, so they
|
|
7
|
+
can be unit-tested directly and reused outside a flowchart.
|
|
8
|
+
|
|
9
|
+
The approach:
|
|
10
|
+
|
|
11
|
+
1. Find the molecules (by bonds) and build a molecular **contact graph**:
|
|
12
|
+
molecule A is adjacent to B if any *contact atom* of A is within ``cutoff``
|
|
13
|
+
(minimum image) of any contact atom of B.
|
|
14
|
+
2. Sample n-mers as **connected subgraphs** of that graph by random frontier
|
|
15
|
+
growth from a random seed molecule. Every member interacts with at least one
|
|
16
|
+
other, and chains, rings and stars all occur -- unlike nearest-neighbour
|
|
17
|
+
selection, which only ever yields the most compact cluster.
|
|
18
|
+
3. Optionally **stratify** the accepted clusters to be flat in a spread
|
|
19
|
+
coordinate (radius of gyration or largest centroid separation of the
|
|
20
|
+
molecules), with bin edges taken as quantiles of a pilot sample, and
|
|
21
|
+
optionally balance jointly over contact-graph topology (motif).
|
|
22
|
+
4. Emit each cluster as an **unwrapped, centred, non-periodic** configuration
|
|
23
|
+
with the molecules intact, a unique name carrying provenance, and optional
|
|
24
|
+
per-configuration properties.
|
|
25
|
+
"""
|
|
26
|
+
|
|
27
|
+
import math
|
|
28
|
+
from collections import Counter
|
|
29
|
+
import logging
|
|
30
|
+
|
|
31
|
+
import numpy as np
|
|
32
|
+
|
|
33
|
+
logger = logging.getLogger(__name__)
|
|
34
|
+
|
|
35
|
+
PROPERTY_TAG = "#ExtractClusters#scan"
|
|
36
|
+
|
|
37
|
+
SPREAD_METRICS = ("rg", "dmax")
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
def classify_motif(n, n_edges, degrees):
|
|
41
|
+
"""Name the topology of the induced contact graph of an n-mer.
|
|
42
|
+
|
|
43
|
+
Only n <= 4 get descriptive names; larger clusters are labelled by edge
|
|
44
|
+
count so they can still be balanced on topology.
|
|
45
|
+
|
|
46
|
+
Parameters
|
|
47
|
+
----------
|
|
48
|
+
n : int
|
|
49
|
+
Number of molecules.
|
|
50
|
+
n_edges : int
|
|
51
|
+
Number of contact edges within the cluster.
|
|
52
|
+
degrees : [int]
|
|
53
|
+
The degree of each molecule within the cluster.
|
|
54
|
+
|
|
55
|
+
Returns
|
|
56
|
+
-------
|
|
57
|
+
str
|
|
58
|
+
"""
|
|
59
|
+
maxdeg = max(degrees) if degrees else 0
|
|
60
|
+
if n == 2:
|
|
61
|
+
return "dimer"
|
|
62
|
+
if n == 3:
|
|
63
|
+
return {2: "chain", 3: "ring"}.get(n_edges, f"e{n_edges}")
|
|
64
|
+
if n == 4:
|
|
65
|
+
if n_edges == 3:
|
|
66
|
+
return "star" if maxdeg == 3 else "chain"
|
|
67
|
+
if n_edges == 4:
|
|
68
|
+
return "ring" if maxdeg == 2 else "paw" # triangle + tail
|
|
69
|
+
if n_edges == 5:
|
|
70
|
+
return "diamond"
|
|
71
|
+
if n_edges == 6:
|
|
72
|
+
return "K4"
|
|
73
|
+
return f"e{n_edges}"
|
|
74
|
+
|
|
75
|
+
|
|
76
|
+
def molecule_neighbour_graph(
|
|
77
|
+
frac_atoms, atom_mol, contact_mask, T, cutoff, periodic, cell_parameters
|
|
78
|
+
):
|
|
79
|
+
"""Molecule-level adjacency from atom contacts.
|
|
80
|
+
|
|
81
|
+
A~B if any contact atom of A is within ``cutoff`` (Å, minimum image) of any
|
|
82
|
+
contact atom of B.
|
|
83
|
+
|
|
84
|
+
Parameters
|
|
85
|
+
----------
|
|
86
|
+
frac_atoms : (N,3) ndarray
|
|
87
|
+
Fractional coordinates (Cartesian if not periodic).
|
|
88
|
+
atom_mol : (N,) ndarray of int
|
|
89
|
+
Molecule index of each atom.
|
|
90
|
+
contact_mask : (N,) ndarray of bool
|
|
91
|
+
Which atoms define contact.
|
|
92
|
+
T : (3,3) ndarray
|
|
93
|
+
Fractional->Cartesian transform (``xyz = uvw @ T``); identity if not
|
|
94
|
+
periodic.
|
|
95
|
+
cutoff : float
|
|
96
|
+
Contact distance, Å.
|
|
97
|
+
periodic : bool
|
|
98
|
+
cell_parameters : (a, b, c, alpha, beta, gamma) or None
|
|
99
|
+
|
|
100
|
+
Returns
|
|
101
|
+
-------
|
|
102
|
+
dict[int, set[int]]
|
|
103
|
+
Adjacency lists keyed by molecule index. Molecules with no contacts
|
|
104
|
+
are absent.
|
|
105
|
+
"""
|
|
106
|
+
from scipy.spatial import cKDTree
|
|
107
|
+
|
|
108
|
+
idx = np.nonzero(contact_mask)[0]
|
|
109
|
+
mol_of = atom_mol[idx]
|
|
110
|
+
|
|
111
|
+
if not periodic:
|
|
112
|
+
xyz = frac_atoms[idx]
|
|
113
|
+
pairs = cKDTree(xyz).query_pairs(cutoff, output_type="ndarray")
|
|
114
|
+
else:
|
|
115
|
+
a, b, c, alpha, beta, gamma = cell_parameters
|
|
116
|
+
orthorhombic = all(abs(x - 90.0) < 1e-6 for x in (alpha, beta, gamma))
|
|
117
|
+
if orthorhombic:
|
|
118
|
+
# Fast path: periodic KD-tree on wrapped Cartesian coordinates.
|
|
119
|
+
L = np.array([a, b, c])
|
|
120
|
+
xyz = (frac_atoms[idx] % 1.0) * L
|
|
121
|
+
xyz = np.minimum(xyz, np.nextafter(L, 0)) # guard x == L
|
|
122
|
+
pairs = cKDTree(xyz, boxsize=L).query_pairs(cutoff, output_type="ndarray")
|
|
123
|
+
else:
|
|
124
|
+
# General cell: exact minimum-image on fractional differences,
|
|
125
|
+
# chunked so memory stays bounded. O(M^2) in contact atoms.
|
|
126
|
+
f = frac_atoms[idx]
|
|
127
|
+
M = len(f)
|
|
128
|
+
out = []
|
|
129
|
+
chunk = max(1, int(2e7 // max(M, 1)))
|
|
130
|
+
for s in range(0, M, chunk):
|
|
131
|
+
d = f[s : s + chunk, None, :] - f[None, :, :]
|
|
132
|
+
d -= np.round(d)
|
|
133
|
+
r = np.linalg.norm(d @ T, axis=2)
|
|
134
|
+
ii, jj = np.nonzero(r < cutoff)
|
|
135
|
+
ii += s
|
|
136
|
+
keep = ii < jj
|
|
137
|
+
out.append(np.stack([ii[keep], jj[keep]], 1))
|
|
138
|
+
pairs = np.concatenate(out) if out else np.zeros((0, 2), int)
|
|
139
|
+
|
|
140
|
+
adj = {}
|
|
141
|
+
if len(pairs):
|
|
142
|
+
mi, mj = mol_of[pairs[:, 0]], mol_of[pairs[:, 1]]
|
|
143
|
+
keep = mi != mj
|
|
144
|
+
for i, j in zip(mi[keep], mj[keep]):
|
|
145
|
+
adj.setdefault(int(i), set()).add(int(j))
|
|
146
|
+
adj.setdefault(int(j), set()).add(int(i))
|
|
147
|
+
return adj
|
|
148
|
+
|
|
149
|
+
|
|
150
|
+
def grow_connected(adj, seed, n, rng):
|
|
151
|
+
"""Grow a connected n-molecule subgraph from ``seed``.
|
|
152
|
+
|
|
153
|
+
Random frontier expansion: at each step a molecule adjacent to the current
|
|
154
|
+
cluster is chosen uniformly and attached through a random one of its
|
|
155
|
+
neighbours already inside. Uniform choice from the frontier is a cheap
|
|
156
|
+
approximation to uniform connected-subgraph sampling; it mildly favours
|
|
157
|
+
high-degree regions.
|
|
158
|
+
|
|
159
|
+
Returns
|
|
160
|
+
-------
|
|
161
|
+
(order, parent) or None
|
|
162
|
+
``order`` is the list of molecule indices in growth order; ``parent``
|
|
163
|
+
maps each to the molecule it attached through (None for the seed).
|
|
164
|
+
None if the connected component is smaller than n.
|
|
165
|
+
"""
|
|
166
|
+
order = [seed]
|
|
167
|
+
parent = {seed: None}
|
|
168
|
+
inside = {seed}
|
|
169
|
+
frontier = set(adj.get(seed, ()))
|
|
170
|
+
while len(order) < n:
|
|
171
|
+
if not frontier:
|
|
172
|
+
return None
|
|
173
|
+
cands = sorted(frontier)
|
|
174
|
+
c = int(cands[rng.integers(len(cands))])
|
|
175
|
+
anchors = sorted(adj[c] & inside)
|
|
176
|
+
parent[c] = int(anchors[rng.integers(len(anchors))])
|
|
177
|
+
order.append(c)
|
|
178
|
+
inside.add(c)
|
|
179
|
+
frontier |= adj.get(c, set())
|
|
180
|
+
frontier -= inside
|
|
181
|
+
return order, parent
|
|
182
|
+
|
|
183
|
+
|
|
184
|
+
def unwrap_shifts(order, parent, frac_centroids, periodic):
|
|
185
|
+
"""Integer lattice shifts making each molecule contiguous with its parent.
|
|
186
|
+
|
|
187
|
+
Each molecule is shifted to the image nearest the neighbour it attached
|
|
188
|
+
through, so the whole cluster is contiguous in space.
|
|
189
|
+
"""
|
|
190
|
+
shifts = {}
|
|
191
|
+
for m in order:
|
|
192
|
+
if parent[m] is None or not periodic:
|
|
193
|
+
shifts[m] = np.zeros(3)
|
|
194
|
+
else:
|
|
195
|
+
p = parent[m]
|
|
196
|
+
u_p = frac_centroids[p] + shifts[p]
|
|
197
|
+
shifts[m] = np.round(u_p - frac_centroids[m])
|
|
198
|
+
return shifts
|
|
199
|
+
|
|
200
|
+
|
|
201
|
+
def _spreads(cen):
|
|
202
|
+
"""Radius of gyration and largest pairwise distance of a set of points."""
|
|
203
|
+
cen0 = cen.mean(axis=0)
|
|
204
|
+
rg = float(np.sqrt(((cen - cen0) ** 2).sum(axis=1).mean()))
|
|
205
|
+
n = len(cen)
|
|
206
|
+
dmax = float(
|
|
207
|
+
max(np.linalg.norm(cen[i] - cen[j]) for i in range(n) for j in range(i))
|
|
208
|
+
)
|
|
209
|
+
return rg, dmax
|
|
210
|
+
|
|
211
|
+
|
|
212
|
+
def pilot_spreads(
|
|
213
|
+
adj, n, frac_centroids, T, periodic, spread_metric, n_mol, rng, n_pilot
|
|
214
|
+
):
|
|
215
|
+
"""Spread values of an unstratified pilot sample, used to place bin edges."""
|
|
216
|
+
out = []
|
|
217
|
+
tries = 0
|
|
218
|
+
while len(out) < n_pilot and tries < 20 * n_pilot:
|
|
219
|
+
tries += 1
|
|
220
|
+
g = grow_connected(adj, int(rng.integers(n_mol)), n, rng)
|
|
221
|
+
if g is None:
|
|
222
|
+
continue
|
|
223
|
+
order, parent = g
|
|
224
|
+
shifts = unwrap_shifts(order, parent, frac_centroids, periodic)
|
|
225
|
+
cen = np.array([(frac_centroids[m] + shifts[m]) @ T for m in order])
|
|
226
|
+
rg, dmax = _spreads(cen)
|
|
227
|
+
out.append(rg if spread_metric == "rg" else dmax)
|
|
228
|
+
return np.array(out)
|
|
229
|
+
|
|
230
|
+
|
|
231
|
+
def bond_index_pairs(configuration):
|
|
232
|
+
"""The bonds of a configuration as 0-based atom-index pairs, with orders.
|
|
233
|
+
|
|
234
|
+
For a configuration without symmetry operators (the normal case, and every
|
|
235
|
+
non-periodic one) the pairs come straight from the bond table, since
|
|
236
|
+
``symmetry.bond_atoms`` is only populated for periodic configurations. For
|
|
237
|
+
a symmetric crystal the symmetry-expanded list is used.
|
|
238
|
+
|
|
239
|
+
Returns
|
|
240
|
+
-------
|
|
241
|
+
(pairs, orders)
|
|
242
|
+
``pairs`` is a list of (i, j) index tuples; ``orders`` the matching
|
|
243
|
+
list of bond orders, or None if it could not be matched up.
|
|
244
|
+
"""
|
|
245
|
+
atoms = configuration.atoms
|
|
246
|
+
bonds = configuration.bonds
|
|
247
|
+
if configuration.symmetry.n_symops == 1:
|
|
248
|
+
index = {aid: k for k, aid in enumerate(atoms.ids)}
|
|
249
|
+
Is = bonds.get_column_data("i")
|
|
250
|
+
Js = bonds.get_column_data("j")
|
|
251
|
+
pairs = [(index[i], index[j]) for i, j in zip(Is, Js)]
|
|
252
|
+
orders = list(bonds.get_column_data("bondorder"))
|
|
253
|
+
else:
|
|
254
|
+
pairs = [(int(i), int(j)) for i, j in configuration.symmetry.bond_atoms]
|
|
255
|
+
orders = list(bonds.bondorders)
|
|
256
|
+
if len(orders) != len(pairs):
|
|
257
|
+
orders = None
|
|
258
|
+
return pairs, orders
|
|
259
|
+
|
|
260
|
+
|
|
261
|
+
def make_put_property(conf, tag=PROPERTY_TAG):
|
|
262
|
+
"""Return a closure storing ``name+tag`` properties on ``conf``.
|
|
263
|
+
|
|
264
|
+
Defines the property first if the database lacks it (the properties are
|
|
265
|
+
normally pre-registered from ``data/properties.csv``).
|
|
266
|
+
"""
|
|
267
|
+
props = conf.properties
|
|
268
|
+
|
|
269
|
+
def _put(name, value, units=None, _type="float"):
|
|
270
|
+
full = f"{name}{tag}"
|
|
271
|
+
if not props.exists(full):
|
|
272
|
+
props.add(full, _type, units=units, noerror=True)
|
|
273
|
+
props.put(full, value)
|
|
274
|
+
|
|
275
|
+
return _put
|
|
276
|
+
|
|
277
|
+
|
|
278
|
+
def extract_nmers(
|
|
279
|
+
configuration,
|
|
280
|
+
n,
|
|
281
|
+
n_samples,
|
|
282
|
+
*,
|
|
283
|
+
cutoff=3.5,
|
|
284
|
+
contact_elements=None,
|
|
285
|
+
spread_metric="rg",
|
|
286
|
+
spread_bins=None,
|
|
287
|
+
balance_motifs=False,
|
|
288
|
+
system=None,
|
|
289
|
+
system_name="clusters",
|
|
290
|
+
name_prefix="",
|
|
291
|
+
max_attempts=None,
|
|
292
|
+
rng=None,
|
|
293
|
+
store_properties=True,
|
|
294
|
+
seen=None,
|
|
295
|
+
):
|
|
296
|
+
"""Extract n-molecule clusters from a (periodic) configuration.
|
|
297
|
+
|
|
298
|
+
The clusters are returned as unwrapped, non-periodic configurations
|
|
299
|
+
suitable for cluster QM.
|
|
300
|
+
|
|
301
|
+
Selection is by *connected subgraph* of a molecular contact graph, not by
|
|
302
|
+
nearest neighbours: every molecule in a returned n-mer is within
|
|
303
|
+
``cutoff`` of at least one other, but chains, rings and stars all occur,
|
|
304
|
+
and seeds are random so a single frame yields many distinct clusters.
|
|
305
|
+
|
|
306
|
+
Parameters
|
|
307
|
+
----------
|
|
308
|
+
configuration : molsystem _Configuration
|
|
309
|
+
Source configuration (periodic or not). Molecules are found by bonds.
|
|
310
|
+
n : int
|
|
311
|
+
Molecules per cluster (>= 2).
|
|
312
|
+
n_samples : int
|
|
313
|
+
Target number of clusters (fewer if the pool or quotas run out).
|
|
314
|
+
cutoff : float
|
|
315
|
+
Contact distance (Å) defining the neighbour graph. 3.5 with
|
|
316
|
+
``contact_elements=["O"]`` is the usual H-bond criterion for water.
|
|
317
|
+
contact_elements : list[str] or None
|
|
318
|
+
Elements whose atoms define contact; None uses all atoms.
|
|
319
|
+
spread_metric : "rg" | "dmax"
|
|
320
|
+
Compactness coordinate used for stratification: radius of gyration of
|
|
321
|
+
the molecular centroids, or the largest centroid-centroid distance.
|
|
322
|
+
spread_bins : int, sequence of float, or None
|
|
323
|
+
Bin edges on the spread coordinate. An int asks for that many
|
|
324
|
+
equal-quantile bins, with edges taken from an unstratified pilot
|
|
325
|
+
sample of the same frame -- the usual choice, since sensible edges
|
|
326
|
+
depend on n and on the system. If given, acceptance is quota-
|
|
327
|
+
limited so bins fill roughly evenly (flat-in-spread, analogous to the
|
|
328
|
+
dimer builder's flat-in-energy selection). Samples outside the edges
|
|
329
|
+
are rejected. None disables stratification.
|
|
330
|
+
balance_motifs : bool
|
|
331
|
+
Also balance across contact-graph topology (chain/ring/star/...),
|
|
332
|
+
jointly with the spread bins.
|
|
333
|
+
system : molsystem _System or None
|
|
334
|
+
Destination system. If None, ``system_name`` is looked up in the
|
|
335
|
+
configuration's database and created if needed.
|
|
336
|
+
system_name : str
|
|
337
|
+
Destination system name, used when ``system`` is None.
|
|
338
|
+
name_prefix : str
|
|
339
|
+
Prefix (e.g. a frame label) for the configuration names, which are
|
|
340
|
+
``f"{name_prefix}{seed}_{m1-m2-...}"`` and unique within a frame.
|
|
341
|
+
max_attempts : int or None
|
|
342
|
+
Sampling attempts before giving up (default 50 * n_samples).
|
|
343
|
+
rng : int, numpy Generator or None
|
|
344
|
+
Seed / generator for reproducibility.
|
|
345
|
+
store_properties : bool
|
|
346
|
+
Store size, spread, motif, edge count, and source molecule list as
|
|
347
|
+
``*#ExtractClusters#scan`` properties on each configuration.
|
|
348
|
+
seen : set or None
|
|
349
|
+
Molecule sets already emitted (as frozensets); updated in place. Pass
|
|
350
|
+
the same set across calls on the same frame to avoid duplicates.
|
|
351
|
+
|
|
352
|
+
Returns
|
|
353
|
+
-------
|
|
354
|
+
(configurations, records, info)
|
|
355
|
+
New configurations; one metadata dict per cluster with keys
|
|
356
|
+
``name, n, molecules, seed, order, rg, dmax, n_edges, degrees, motif,
|
|
357
|
+
bin``; and an ``info`` dict with ``edges`` (the bin edges used, or
|
|
358
|
+
None), ``attempts``, ``per_key`` quota, ``n_molecules`` and
|
|
359
|
+
``n_contacts`` of the frame's graph, and ``warning`` (a message or
|
|
360
|
+
None) if fewer than requested were produced.
|
|
361
|
+
"""
|
|
362
|
+
if n < 2:
|
|
363
|
+
raise ValueError("The cluster size must be >= 2")
|
|
364
|
+
if spread_metric not in SPREAD_METRICS:
|
|
365
|
+
raise ValueError(f"spread_metric must be one of {SPREAD_METRICS}")
|
|
366
|
+
if n_samples < 1:
|
|
367
|
+
raise ValueError("The number of clusters must be >= 1")
|
|
368
|
+
rng = np.random.default_rng(rng)
|
|
369
|
+
if max_attempts is None:
|
|
370
|
+
max_attempts = 50 * n_samples
|
|
371
|
+
if seen is None:
|
|
372
|
+
seen = set()
|
|
373
|
+
|
|
374
|
+
periodic = configuration.periodicity != 0
|
|
375
|
+
molecules = [np.asarray(m) for m in configuration.find_molecules(as_indices=True)]
|
|
376
|
+
n_mol = len(molecules)
|
|
377
|
+
if n_mol < n:
|
|
378
|
+
raise ValueError(
|
|
379
|
+
f"The configuration has only {n_mol} molecules; cannot make {n}-mers"
|
|
380
|
+
)
|
|
381
|
+
atom_mol = np.empty(configuration.n_atoms, dtype=int)
|
|
382
|
+
for k, m in enumerate(molecules):
|
|
383
|
+
atom_mol[m] = k
|
|
384
|
+
|
|
385
|
+
if periodic:
|
|
386
|
+
# Fractional, molecules kept whole across the boundary.
|
|
387
|
+
frac = configuration.atoms.get_coordinates(
|
|
388
|
+
fractionals=True, in_cell="molecule", as_array=True
|
|
389
|
+
)
|
|
390
|
+
T = configuration.cell.to_cartesians_transform(as_array=True)
|
|
391
|
+
cell_parameters = configuration.cell.parameters
|
|
392
|
+
else:
|
|
393
|
+
frac = configuration.atoms.get_coordinates(fractionals=False, as_array=True)
|
|
394
|
+
T = np.eye(3)
|
|
395
|
+
cell_parameters = None
|
|
396
|
+
|
|
397
|
+
frac_centroids = np.array([frac[m].mean(axis=0) for m in molecules])
|
|
398
|
+
|
|
399
|
+
symbols = np.array(configuration.atoms.symbols)
|
|
400
|
+
if contact_elements is None:
|
|
401
|
+
contact_mask = np.ones(len(symbols), dtype=bool)
|
|
402
|
+
else:
|
|
403
|
+
contact_mask = np.isin(symbols, list(contact_elements))
|
|
404
|
+
if not contact_mask.any():
|
|
405
|
+
raise ValueError(
|
|
406
|
+
f"No atoms match the contact elements {list(contact_elements)}"
|
|
407
|
+
)
|
|
408
|
+
|
|
409
|
+
adj = molecule_neighbour_graph(
|
|
410
|
+
frac, atom_mol, contact_mask, T, cutoff, periodic, cell_parameters
|
|
411
|
+
)
|
|
412
|
+
if not adj:
|
|
413
|
+
raise ValueError(f"No molecular contacts within {cutoff} Å")
|
|
414
|
+
n_contacts = sum(len(v) for v in adj.values()) // 2
|
|
415
|
+
|
|
416
|
+
# Stratification bins -----------------------------------------------------
|
|
417
|
+
if isinstance(spread_bins, (int, np.integer)):
|
|
418
|
+
# Choose edges from the data: equal-quantile bins of the spread
|
|
419
|
+
# coordinate over an unstratified pilot sample.
|
|
420
|
+
n_bins_requested = int(spread_bins)
|
|
421
|
+
pilot = pilot_spreads(
|
|
422
|
+
adj,
|
|
423
|
+
n,
|
|
424
|
+
frac_centroids,
|
|
425
|
+
T,
|
|
426
|
+
periodic,
|
|
427
|
+
spread_metric,
|
|
428
|
+
n_mol,
|
|
429
|
+
rng,
|
|
430
|
+
n_pilot=max(200, 20 * n_bins_requested),
|
|
431
|
+
)
|
|
432
|
+
if len(pilot) < 2 * n_bins_requested:
|
|
433
|
+
raise ValueError(
|
|
434
|
+
f"Too few {n}-mers found ({len(pilot)}) to choose "
|
|
435
|
+
f"{n_bins_requested} spread bins"
|
|
436
|
+
)
|
|
437
|
+
edges = np.quantile(pilot, np.linspace(0.0, 1.0, n_bins_requested + 1))
|
|
438
|
+
edges[0] -= 1e-9
|
|
439
|
+
edges[-1] += 1e-9
|
|
440
|
+
elif spread_bins is None:
|
|
441
|
+
edges = None
|
|
442
|
+
else:
|
|
443
|
+
edges = np.asarray(spread_bins, dtype=float)
|
|
444
|
+
if edges.ndim != 1 or len(edges) < 2:
|
|
445
|
+
raise ValueError("spread_bins must give at least two bin edges")
|
|
446
|
+
if np.any(np.diff(edges) <= 0):
|
|
447
|
+
raise ValueError("spread_bins edges must be strictly increasing")
|
|
448
|
+
n_bins = 1 if edges is None else len(edges) - 1
|
|
449
|
+
stratified = edges is not None or balance_motifs
|
|
450
|
+
|
|
451
|
+
atnos = np.asarray(configuration.atoms.atomic_numbers)
|
|
452
|
+
bond_pairs, bond_orders = bond_index_pairs(configuration)
|
|
453
|
+
bond_indices = {}
|
|
454
|
+
for k, (i, j) in enumerate(bond_pairs):
|
|
455
|
+
bond_indices.setdefault(int(i), []).append(k)
|
|
456
|
+
|
|
457
|
+
if system is None:
|
|
458
|
+
system_db = configuration.system_db
|
|
459
|
+
systems = system_db.get_systems(system_name)
|
|
460
|
+
system = systems[0] if systems else system_db.create_system(name=system_name)
|
|
461
|
+
|
|
462
|
+
# Quotas: each (motif?, bin) key gets an equal share of n_samples over the
|
|
463
|
+
# keys seen so far -- the motif set is not known in advance.
|
|
464
|
+
per_key = None
|
|
465
|
+
configurations = []
|
|
466
|
+
records = []
|
|
467
|
+
counts = Counter()
|
|
468
|
+
attempts = 0
|
|
469
|
+
while len(configurations) < n_samples and attempts < max_attempts:
|
|
470
|
+
attempts += 1
|
|
471
|
+
seed = int(rng.integers(n_mol))
|
|
472
|
+
grown = grow_connected(adj, seed, n, rng)
|
|
473
|
+
if grown is None:
|
|
474
|
+
continue
|
|
475
|
+
order, parent = grown
|
|
476
|
+
key = frozenset(order)
|
|
477
|
+
if key in seen:
|
|
478
|
+
continue
|
|
479
|
+
|
|
480
|
+
shifts = unwrap_shifts(order, parent, frac_centroids, periodic)
|
|
481
|
+
cen = np.array([(frac_centroids[m] + shifts[m]) @ T for m in order])
|
|
482
|
+
rg, dmax = _spreads(cen)
|
|
483
|
+
spread = rg if spread_metric == "rg" else dmax
|
|
484
|
+
|
|
485
|
+
inside = set(order)
|
|
486
|
+
sub_edges = [
|
|
487
|
+
(i, j) for i in order for j in adj.get(i, ()) if j in inside and i < j
|
|
488
|
+
]
|
|
489
|
+
n_edges = len(sub_edges)
|
|
490
|
+
deg = Counter()
|
|
491
|
+
for i, j in sub_edges:
|
|
492
|
+
deg[i] += 1
|
|
493
|
+
deg[j] += 1
|
|
494
|
+
degrees = sorted((deg[m] for m in order), reverse=True)
|
|
495
|
+
motif = classify_motif(n, n_edges, degrees)
|
|
496
|
+
|
|
497
|
+
if edges is not None:
|
|
498
|
+
b = int(np.digitize(spread, edges)) - 1
|
|
499
|
+
if b < 0 or b >= n_bins:
|
|
500
|
+
continue # outside requested range
|
|
501
|
+
else:
|
|
502
|
+
b = 0
|
|
503
|
+
qkey = (motif, b) if balance_motifs else (b,)
|
|
504
|
+
if stratified:
|
|
505
|
+
# equal share per key, over bins x (motifs seen so far)
|
|
506
|
+
if balance_motifs:
|
|
507
|
+
n_motifs = max(1, len({k[0] for k in counts} | {motif}))
|
|
508
|
+
else:
|
|
509
|
+
n_motifs = 1
|
|
510
|
+
per_key = math.ceil(n_samples / (n_bins * n_motifs))
|
|
511
|
+
if counts[qkey] >= per_key:
|
|
512
|
+
continue
|
|
513
|
+
|
|
514
|
+
seen.add(key)
|
|
515
|
+
counts[qkey] += 1
|
|
516
|
+
|
|
517
|
+
# Build the cluster ---------------------------------------------------
|
|
518
|
+
atom_indices = np.concatenate([molecules[m] for m in order])
|
|
519
|
+
xyz = np.concatenate([(frac[molecules[m]] + shifts[m]) @ T for m in order])
|
|
520
|
+
xyz -= xyz.mean(axis=0) # centre the cluster at the origin
|
|
521
|
+
to_new = {int(old): new for new, old in enumerate(atom_indices)}
|
|
522
|
+
Is, Js, orders = [], [], []
|
|
523
|
+
for i in atom_indices:
|
|
524
|
+
for k in bond_indices.get(int(i), ()):
|
|
525
|
+
i0, j0 = bond_pairs[k]
|
|
526
|
+
if j0 in to_new:
|
|
527
|
+
Is.append(to_new[int(i0)])
|
|
528
|
+
Js.append(to_new[int(j0)])
|
|
529
|
+
if bond_orders is not None:
|
|
530
|
+
orders.append(bond_orders[k])
|
|
531
|
+
|
|
532
|
+
name = f"{name_prefix}{seed}_" + "-".join(str(m) for m in sorted(order))
|
|
533
|
+
conf = system.create_configuration(
|
|
534
|
+
name=name, periodicity=0, coordinate_system="Cartesian", make_current=False
|
|
535
|
+
)
|
|
536
|
+
ids = conf.atoms.append(
|
|
537
|
+
atno=atnos[atom_indices].tolist(),
|
|
538
|
+
x=xyz[:, 0].tolist(),
|
|
539
|
+
y=xyz[:, 1].tolist(),
|
|
540
|
+
z=xyz[:, 2].tolist(),
|
|
541
|
+
)
|
|
542
|
+
if Is:
|
|
543
|
+
kwargs = {}
|
|
544
|
+
if bond_orders is not None:
|
|
545
|
+
kwargs["bondorder"] = orders
|
|
546
|
+
conf.bonds.append(i=[ids[k] for k in Is], j=[ids[k] for k in Js], **kwargs)
|
|
547
|
+
|
|
548
|
+
rec = dict(
|
|
549
|
+
name=name,
|
|
550
|
+
n=n,
|
|
551
|
+
molecules=sorted(order),
|
|
552
|
+
seed=seed,
|
|
553
|
+
order=list(order),
|
|
554
|
+
rg=rg,
|
|
555
|
+
dmax=dmax,
|
|
556
|
+
n_edges=n_edges,
|
|
557
|
+
degrees=degrees,
|
|
558
|
+
motif=motif,
|
|
559
|
+
bin=b,
|
|
560
|
+
)
|
|
561
|
+
if store_properties:
|
|
562
|
+
_put = make_put_property(conf)
|
|
563
|
+
_put("cluster size", n, None, "int")
|
|
564
|
+
_put("spread rg", rg, "Å")
|
|
565
|
+
_put("spread dmax", dmax, "Å")
|
|
566
|
+
_put("n edges", n_edges, None, "int")
|
|
567
|
+
_put("motif", motif, None, "str")
|
|
568
|
+
_put("spread bin", b, None, "int")
|
|
569
|
+
_put(
|
|
570
|
+
"source molecules", "-".join(str(m) for m in sorted(order)), None, "str"
|
|
571
|
+
)
|
|
572
|
+
configurations.append(conf)
|
|
573
|
+
records.append(rec)
|
|
574
|
+
|
|
575
|
+
warning = None
|
|
576
|
+
if len(configurations) < n_samples:
|
|
577
|
+
warning = (
|
|
578
|
+
f"Produced {len(configurations)} of {n_samples} requested {n}-mers "
|
|
579
|
+
f"after {attempts} attempts"
|
|
580
|
+
)
|
|
581
|
+
if stratified:
|
|
582
|
+
warning += (
|
|
583
|
+
f" (per-key quota {per_key}). Quotas for (motif, bin) cells that "
|
|
584
|
+
"never occur cannot fill -- e.g. rings only exist in the compact "
|
|
585
|
+
"bins -- so a short fill here usually reflects the physics rather "
|
|
586
|
+
"than a sampling problem. Check the motif x bin table and revisit "
|
|
587
|
+
"the stratification or motif balancing if needed."
|
|
588
|
+
)
|
|
589
|
+
else:
|
|
590
|
+
warning += (
|
|
591
|
+
". The frame may not contain that many distinct connected "
|
|
592
|
+
f"{n}-mers within the contact cutoff."
|
|
593
|
+
)
|
|
594
|
+
|
|
595
|
+
info = dict(
|
|
596
|
+
edges=None if edges is None else [float(e) for e in edges],
|
|
597
|
+
attempts=attempts,
|
|
598
|
+
per_key=per_key,
|
|
599
|
+
n_molecules=n_mol,
|
|
600
|
+
n_contacts=n_contacts,
|
|
601
|
+
warning=warning,
|
|
602
|
+
)
|
|
603
|
+
return configurations, records, info
|
|
604
|
+
|
|
605
|
+
|
|
606
|
+
def cluster_summary(records, edges=None):
|
|
607
|
+
"""Tabulate how a set of extracted n-mers is distributed over motif and
|
|
608
|
+
spread bin -- the quick check that the set is actually balanced.
|
|
609
|
+
|
|
610
|
+
Parameters
|
|
611
|
+
----------
|
|
612
|
+
records : [dict]
|
|
613
|
+
The records from :func:`extract_nmers`.
|
|
614
|
+
edges : [float] or None
|
|
615
|
+
The bin edges, printed as a final line if given.
|
|
616
|
+
|
|
617
|
+
Returns
|
|
618
|
+
-------
|
|
619
|
+
str
|
|
620
|
+
"""
|
|
621
|
+
tab = Counter((r["motif"], r["bin"]) for r in records)
|
|
622
|
+
motifs = sorted({m for m, _ in tab})
|
|
623
|
+
bins = sorted({b for _, b in tab})
|
|
624
|
+
w = max(8, max((len(m) for m in motifs), default=8))
|
|
625
|
+
header = " " * w + "".join(f"{('bin%d' % b):>8}" for b in bins) + f"{'total':>8}"
|
|
626
|
+
lines = [header]
|
|
627
|
+
for m in motifs:
|
|
628
|
+
row = [tab[(m, b)] for b in bins]
|
|
629
|
+
lines.append(f"{m:<{w}}" + "".join(f"{c:>8}" for c in row) + f"{sum(row):>8}")
|
|
630
|
+
tot = [sum(tab[(m, b)] for m in motifs) for b in bins]
|
|
631
|
+
lines.append(f"{'total':<{w}}" + "".join(f"{c:>8}" for c in tot) + f"{sum(tot):>8}")
|
|
632
|
+
if edges is not None:
|
|
633
|
+
e = list(edges)
|
|
634
|
+
lines.append(
|
|
635
|
+
"bins (Å): "
|
|
636
|
+
+ ", ".join(f"{e[i]:.2f}-{e[i + 1]:.2f}" for i in range(len(e) - 1))
|
|
637
|
+
)
|
|
638
|
+
return "\n".join(lines)
|