mdinterface 1.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.
@@ -0,0 +1,10 @@
1
+ """
2
+ mdinterface is a package to setup molecular dynamics systems.
3
+ """
4
+
5
+ __version__ = '1.0.0'
6
+ __date__ = '16 Dec. 2024'
7
+ __author__ = 'Fabrice Roncoroni'
8
+ __all__ = ['SimulationBox']
9
+
10
+ from .simulationbox import SimulationBox
@@ -0,0 +1,8 @@
1
+ # core/__init__.py
2
+
3
+ """
4
+ core: Core functionalities including related components.
5
+ """
6
+
7
+ from .specie import *
8
+ from .topology import *
@@ -0,0 +1,507 @@
1
+ #!/usr/bin/env python3
2
+ # -*- coding: utf-8 -*-
3
+ """
4
+ Created on Fri Apr 19 13:58:43 2024
5
+
6
+ @author: roncofaber
7
+ """
8
+
9
+ import ase
10
+ import ase.build
11
+ import ase.visualize
12
+ from ase.data.colors import jmol_colors
13
+ import MDAnalysis as mda
14
+
15
+ from mdinterface.utils.auxiliary import as_list, find_smallest_missing, atoms_to_indexes
16
+
17
+ from mdinterface.core.topology import Atom
18
+ import mdinterface.utils.auxiliary as aux
19
+ from mdinterface.io.read import read_lammps_data_file
20
+ import mdinterface.utils.map as pmap
21
+
22
+ import copy
23
+ import numpy as np
24
+ import networkx as nx
25
+
26
+ import matplotlib.pyplot as plt
27
+
28
+ #%%
29
+
30
+ class Specie(object):
31
+
32
+ def __init__(self, atoms=None, charges=None, atom_types=None, bonds=None,
33
+ angles=None, dihedrals=None, impropers=None, lj={}, cutoff=1.0,
34
+ name=None, lammps_data=None, fix_missing=False, chg_scaling=1.0):
35
+
36
+ # if file provided, read it
37
+ if lammps_data is not None:
38
+ atoms, atom_types, bonds, angles, dihedrals, impropers = read_lammps_data_file(lammps_data)
39
+ else:
40
+ # read/update atoms atoms
41
+ atoms, atom_types = self._read_atoms(atoms, charges, chg_scaling=chg_scaling)
42
+
43
+ # assign name
44
+ if name is None:
45
+ name = atoms.get_chemical_formula()
46
+ if len(name) > 4:
47
+ print("ATTENTION: resname for Specie could be misleading")
48
+ self.resname = name
49
+
50
+ # set up atoms and generate graph of specie
51
+ self.set_atoms(atoms, cutoff)
52
+
53
+ # read atom_types from LJ
54
+ if atom_types is None:
55
+ atom_types = self._atom_types_from_lj(lj)
56
+
57
+ # set up internal topology attributes
58
+ self._setup_topology(atom_types, bonds, angles, dihedrals, impropers,
59
+ fix_missing=fix_missing)
60
+
61
+ # initialize topology info
62
+ self._update_topology()
63
+
64
+ return
65
+
66
+ # read atoms to return ase.Atoms
67
+ @staticmethod
68
+ def _read_atoms(atoms, charges, chg_scaling=1.0):
69
+
70
+ # initialize atoms obj
71
+ if isinstance(atoms, str):
72
+ try:
73
+ atoms = ase.io.read(atoms)
74
+ except:
75
+ atoms = ase.build.molecule(atoms)
76
+ elif isinstance(atoms, ase.Atoms):
77
+ atoms = atoms.copy()
78
+
79
+ # assign charges
80
+ if charges is not None:
81
+ charges = as_list(charges)
82
+ if len(charges) == 1:
83
+ charges = len(atoms)*charges
84
+ else:
85
+ charges = atoms.get_initial_charges()
86
+
87
+ # rescale charges if needed
88
+ charges = chg_scaling*np.asarray(charges)
89
+ atoms.set_initial_charges(charges)
90
+
91
+ # see if stype is already present
92
+ if "stype" in atoms.arrays:
93
+ stype = atoms.arrays["stype"]
94
+ else:
95
+ stype = None
96
+
97
+ return atoms, stype
98
+
99
+ def _setup_topology(self, atoms, bonds, angles, dihedrals, impropers,
100
+ fix_missing=False):
101
+
102
+ # map list of inputs
103
+ atoms_list, atom_map, atom_ids = pmap.map_atoms(as_list(atoms))
104
+ bonds_list, bond_map, bond_ids = pmap.map_bonds(as_list(bonds))
105
+ angles_list, angle_map, angle_ids = pmap.map_angles(as_list(angles))
106
+ dihedrals_list, dihedral_map, dihedral_ids = pmap.map_dihedrals(as_list(dihedrals))
107
+ impropers_list, improper_map, improper_ids = pmap.map_impropers(as_list(impropers))
108
+
109
+ self._btype = bonds_list
110
+ self._atype = angles_list
111
+ self._dtype = dihedrals_list
112
+ self._itype = impropers_list
113
+ self._stype = atoms_list
114
+
115
+ self._smap = atom_map
116
+ self._bmap = bond_map
117
+ self._amap = angle_map
118
+ self._dmap = dihedral_map
119
+ self._imap = improper_map
120
+
121
+ self._sids = atom_ids
122
+ self._bids = bond_ids
123
+ self._aids = angle_ids
124
+ self._dids = dihedral_ids
125
+ self._iids = improper_ids
126
+
127
+ if fix_missing: #ugly but seems to work
128
+
129
+ mss_bnd = pmap.generate_missing_interactions(self, "bonds")
130
+ mss_ang = pmap.generate_missing_interactions(self, "angles")
131
+ mss_dih = pmap.generate_missing_interactions(self, "dihedrals")
132
+ mss_imp = pmap.generate_missing_interactions(self, "impropers")
133
+
134
+ # map list of inputs
135
+ # atoms_list, atom_map, atom_ids = pmap.map_atoms(as_list(atoms))
136
+ bonds_list, bond_map, bond_ids = pmap.map_bonds(as_list(bonds) + mss_bnd)
137
+ angles_list, angle_map, angle_ids = pmap.map_angles(as_list(angles) + mss_ang)
138
+ dihedrals_list, dihedral_map, dihedral_ids = pmap.map_dihedrals(as_list(dihedrals) + mss_dih)
139
+ impropers_list, improper_map, improper_ids = pmap.map_impropers(as_list(impropers) + mss_imp)
140
+
141
+ self._btype = bonds_list
142
+ self._atype = angles_list
143
+ self._dtype = dihedrals_list
144
+ self._itype = impropers_list
145
+ self._stype = atoms_list
146
+
147
+ self._smap = atom_map
148
+ self._bmap = bond_map
149
+ self._amap = angle_map
150
+ self._dmap = dihedral_map
151
+ self._imap = improper_map
152
+
153
+ self._sids = atom_ids
154
+ self._bids = bond_ids
155
+ self._aids = angle_ids
156
+ self._dids = dihedral_ids
157
+ self._iids = improper_ids
158
+
159
+ return
160
+
161
+ # function to setup atom types
162
+ def _atom_types_from_lj(self, lj):
163
+
164
+ # use function to retrieve IDs
165
+ atom_type_ids, types_map = aux.find_atom_types(self.atoms, max_depth=1)
166
+
167
+ atom_types = []
168
+ for atom_id in atom_type_ids:
169
+
170
+ atom_symbol = types_map[atom_id][0]
171
+ atom_neighs = "".join(types_map[atom_id][1])
172
+
173
+ label = "{}_{}".format(atom_symbol, atom_neighs)
174
+ # types_map[atom_type] = label
175
+
176
+ if label in lj:
177
+ eps, sig = lj[label]
178
+ elif atom_symbol in lj:
179
+ eps, sig = lj[atom_symbol]
180
+ else:
181
+ eps, sig = None,None
182
+
183
+ atom = Atom(atom_symbol, label=label, eps=eps, sig=sig)
184
+ atom_types.append(atom)
185
+
186
+ return atom_types
187
+
188
+ def set_atoms(self, atoms, cutoff=1.0):
189
+
190
+ self._atoms = atoms
191
+
192
+ self._graph = aux.molecule_to_graph(atoms, cutoff_scale=cutoff)
193
+
194
+ return
195
+
196
+ def _update_topology(self):
197
+ formula = self.atoms.get_chemical_formula()
198
+
199
+ for attributes in ["_btype", "_atype", "_dtype", "_itype", "_stype"]:
200
+ attr_type = []
201
+ for attr in self.__getattribute__(attributes):
202
+ attr.set_formula(formula)
203
+ attr.set_resname(self.resname)
204
+ if attr.id is None or attr.id in attr_type:
205
+ idx = find_smallest_missing(attr_type, start=1)
206
+ attr.set_id(idx)
207
+ else:
208
+ idx = attr.id
209
+ attr_type.append(idx)
210
+
211
+ return
212
+
213
+
214
+ def add_topology(self, attribute):
215
+ # Determine the type of the attribute
216
+ attribute_type = attribute.__class__.__name__
217
+
218
+ # Get the corresponding attribute list
219
+ if attribute_type == "Bond":
220
+ attribute_list = self._btype
221
+ elif attribute_type == "Angle":
222
+ attribute_list = self._atype
223
+ elif attribute_type == "Dihedral":
224
+ attribute_list = self._dtype
225
+ elif attribute_type == "Improper":
226
+ attribute_list = self._itype
227
+ elif attribute_type == "Atom":
228
+ attribute_list = self._stype
229
+ else:
230
+ raise ValueError("Invalid topology attribute type.")
231
+
232
+ # Check if the attribute already exists based on symbols
233
+ attribute_symbols = attribute.symbols
234
+ for existing_attribute in attribute_list:
235
+ if existing_attribute.symbols == attribute_symbols:
236
+ print(f"{attribute_type} with symbols {attribute_symbols} already exists and will not be added again.")
237
+ return
238
+
239
+ # Add the attribute to the list
240
+ attribute_list.append(attribute)
241
+
242
+ # Reinitialize topology to ensure consistency
243
+ self._update_topology()
244
+
245
+ return
246
+
247
+
248
+ # covnert to mda.Universe
249
+ def to_universe(self, charges=True, layered=False):
250
+
251
+ # empty top object
252
+ top = mda.core.topology.Topology(n_atoms=len(self.atoms))
253
+
254
+ # empty universe
255
+ uni = mda.Universe(top, self.atoms.get_positions())
256
+
257
+ # add some stuff
258
+ uni.add_TopologyAttr("masses", self.atoms.get_masses())
259
+ uni.add_TopologyAttr("resnames", [self.resname])
260
+
261
+ # generate type
262
+ types, indexes = self.get_atom_types(return_index=True)
263
+ uni.add_TopologyAttr("types", types)
264
+ uni.add_TopologyAttr("type_index", indexes)
265
+
266
+ # populate with bonds and angles
267
+ for att in ["bonds", "angles", "dihedrals", "impropers"]:
268
+ attribute, types = self.__getattribute__(att)
269
+ att_types = self._type2id(att, types)
270
+ uni._add_topology_objects(att, attribute, types=att_types)
271
+
272
+ # add charges
273
+ if charges:
274
+ uni.add_TopologyAttr("charges", self.atoms.get_initial_charges())
275
+
276
+ # layer it nicely
277
+ if layered:
278
+ layer_idxs, __ = ase.geometry.get_layers(self.atoms, [0,0,1],
279
+ tolerance=0.01)
280
+ groups = []
281
+ for idx in np.unique(layer_idxs):
282
+ idxs = np.where(layer_idxs == idx)[0]
283
+
284
+ groups.append(uni.atoms[idxs])
285
+
286
+ uni = mda.Merge(*groups)
287
+
288
+ # uni.residues[0].resids = layer_idxs
289
+
290
+ # if has cell info, pass them along
291
+ if self.atoms.get_cell():
292
+ uni.dimensions = self.atoms.cell.cellpar()
293
+
294
+ return uni
295
+
296
+ def repeat(self, rep, make_cubic=False):
297
+
298
+ atoms = self.atoms.repeat(rep)
299
+
300
+ if make_cubic: #TODO: this can be dangerous if the cell
301
+
302
+ xsize = [1,0,0]@atoms.cell@[1,0,0]
303
+ ysize = [0,1,0]@atoms.cell@[0,1,0]
304
+ zsize = [0,0,1]@atoms.cell@[0,0,1]
305
+
306
+ atoms.set_cell([xsize, ysize, zsize, 90, 90, 90])
307
+ atoms.wrap()
308
+
309
+ self.set_atoms(atoms)
310
+
311
+ return
312
+
313
+ @property
314
+ def atoms(self):
315
+ return self._atoms
316
+
317
+ @property
318
+ def graph(self):
319
+ return self._graph
320
+
321
+ def _find_interactions(self, path_length, tag_map, impropers=False):
322
+
323
+ if not impropers:
324
+ paths = aux.find_unique_paths_of_length(self.graph, path_length)
325
+ else:
326
+ paths = aux.find_improper_idxs(self.graph)
327
+
328
+ interaction_list = []
329
+ interaction_type = []
330
+
331
+ for indices in paths:
332
+ atoms = [self._sids[idx] for idx in indices]
333
+ symbols = [atom.split("_")[0] for atom in atoms]
334
+
335
+ if tuple(atoms) in tag_map:
336
+ interaction_list.append(indices)
337
+ interaction_type.append(tag_map[tuple(atoms)])
338
+ elif tuple(atoms[::-1]) in tag_map:
339
+ interaction_list.append(indices[::-1])
340
+ interaction_type.append(tag_map[tuple(atoms[::-1])])
341
+ elif tuple(symbols) in tag_map:
342
+ interaction_list.append(indices)
343
+ interaction_type.append(tag_map[tuple(symbols)])
344
+ elif tuple(symbols[::-1]) in tag_map:
345
+ interaction_list.append(indices[::-1])
346
+ interaction_type.append(tag_map[tuple(symbols[::-1])])
347
+
348
+ interaction_list = np.array(interaction_list, dtype=int)
349
+ interaction_type = np.array(interaction_type, dtype=int)
350
+
351
+ return [interaction_list.tolist(), interaction_type.tolist()]
352
+
353
+ @property
354
+ def bonds(self):
355
+ return self._find_interactions(1, self._bmap)
356
+
357
+ @property
358
+ def angles(self):
359
+ return self._find_interactions(2, self._amap)
360
+
361
+ @property
362
+ def dihedrals(self):
363
+ return self._find_interactions(3, self._dmap)
364
+
365
+ @property
366
+ def impropers(self):
367
+ return self._find_interactions(3, self._imap, impropers=True)
368
+
369
+ @property
370
+ def _sids(self):
371
+ return self.atoms.arrays["sids"]
372
+
373
+ @_sids.setter
374
+ def _sids(self, atom_ids):
375
+ self.atoms.arrays["sids"] = atom_ids
376
+
377
+ def copy(self):
378
+ return copy.deepcopy(self)
379
+
380
+ def get_atom_types(self, return_index=False):
381
+
382
+ type_indexes = np.array([self._smap[ii] for ii in self._sids])
383
+ atom_types = np.array([self._stype[ii].extended_label for ii in type_indexes])
384
+
385
+ if return_index:
386
+ return atom_types, type_indexes
387
+
388
+ return atom_types
389
+
390
+ def estimate_sphere_radius(self):
391
+ """
392
+ Estimate the radius of the sphere containing the given points.
393
+
394
+ Parameters:
395
+ points (numpy.ndarray): A 2D array of shape (n, 3) where n is the number of points.
396
+
397
+ Returns:
398
+ float: The estimated radius of the sphere.
399
+ """
400
+
401
+ points = self.atoms.get_positions()
402
+
403
+ # Calculate the centroid of the points
404
+ centroid = np.mean(points, axis=0)
405
+
406
+ # Calculate the distances from the centroid to each point
407
+ distances = np.linalg.norm(points - centroid, axis=1)
408
+
409
+ # The radius of the sphere is the maximum distance from the centroid to any point
410
+ radius = np.max(distances)
411
+
412
+ return radius
413
+
414
+
415
+ def view(self):
416
+ ase.visualize.view(self.atoms)
417
+ return
418
+
419
+ def _type2id(self, attribute, types):
420
+
421
+ # Get the corresponding attribute list
422
+ if attribute == "bonds":
423
+ attribute_list = self._btype
424
+ elif attribute == "angles":
425
+ attribute_list = self._atype
426
+ elif attribute == "dihedrals":
427
+ attribute_list = self._dtype
428
+ elif attribute == "impropers":
429
+ attribute_list = self._itype
430
+ elif attribute == "atoms":
431
+ attribute_list = self._stype
432
+ else:
433
+ raise ValueError("Invalid topology attribute type.")
434
+
435
+ return [attribute_list[idx].id for idx in types]
436
+
437
+ def suggest_missing_interactions(self, stype="all"):
438
+
439
+
440
+ # Check for missing interactions
441
+ missing_bonds = pmap.find_missing_bonds(self)
442
+ missing_angles = pmap.find_missing_angles(self)
443
+ missing_dihedrals = pmap.find_missing_dihedrals(self)
444
+ missing_impropers = []#pmap.find_missing_impropers(self) #TODO change
445
+
446
+ suggestions = {
447
+ "bonds": missing_bonds,
448
+ "angles": missing_angles,
449
+ "dihedrals": missing_dihedrals,
450
+ "impropers": missing_impropers
451
+ }
452
+
453
+ if stype == "all":
454
+ return suggestions
455
+
456
+ return suggestions[stype]
457
+
458
+ def __repr__(self):
459
+ return f"{self.__class__.__name__}({self.resname})"
460
+
461
+ def find_relevant_distances(self, Nmax, Nmin=0, centers=None, Ninv=0):
462
+
463
+ # get list of relevant nodes
464
+ if centers is None:
465
+ relevant_nodes = self.graph.nodes()
466
+ else:
467
+ relevant_nodes = aux.as_list(centers)
468
+
469
+ # Set to store unique pairs
470
+ unique_pairs = set()
471
+
472
+ # Iterate over all nodes in the graph
473
+ for node1 in relevant_nodes:
474
+ # Get the shortest path lengths from node node to all other reachable nodes
475
+ shortest_paths = nx.single_source_shortest_path_length(self.graph, node1)
476
+
477
+ longest_path = max([dist for _, dist in shortest_paths.items()])
478
+
479
+ # Collect pairs where the distance is within N but above Nmin
480
+ for node2, distance in shortest_paths.items():
481
+ if distance <= Nmax and distance > Nmin:
482
+ # Use tuple (min(node, n), max(node, n)) to avoid duplicates
483
+ pair = (min(node1, node2), max(node1, node2))
484
+ unique_pairs.add(pair)
485
+ elif Ninv > longest_path - distance:
486
+ pair = (min(node1, node2), max(node1, node2))
487
+ unique_pairs.add(pair)
488
+
489
+ # Convert set to list
490
+ unique_pairs_list = list(unique_pairs)
491
+ unique_pairs_list.sort()
492
+
493
+ return np.array(unique_pairs_list, dtype=int)
494
+
495
+ def plot_graph(self, **kwargs):
496
+
497
+ colors = [jmol_colors[a.number] for a in self.atoms]
498
+
499
+ fig, ax = plt.subplots()
500
+ nx.draw(self.graph, with_labels=True, node_color=colors,
501
+ node_size=1000, edge_color='black', linewidths=2, font_size=15,
502
+ edgecolors="black", ax=ax, width=2)
503
+
504
+
505
+ plt.show()
506
+
507
+ return