PFASGroups 3.2.2__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,977 @@
1
+ import os
2
+ import json
3
+ import math
4
+ import warnings
5
+ import numpy as np
6
+ import networkx as nx
7
+ from rdkit import Chem
8
+ from .core import add_componentSmarts, mol_to_nx
9
+
10
+ # ── BDE helpers (self-contained, no dependency on molecular_quantum_graph) ────
11
+
12
+ def _bde_keys_to_int(x):
13
+ return {int(k) if isinstance(k, str) and k.isdigit() else k: v
14
+ for k, v in x.items()}
15
+
16
+
17
+ def _load_bde_dict_local():
18
+ """Load diatomic BDE dict (kcal/mol) from PFASGroups' own data folder."""
19
+ _path = os.path.join(os.path.dirname(__file__), 'data', 'diatomic_bonds_dict.json')
20
+ try:
21
+ with open(_path) as fh:
22
+ raw = json.load(fh, object_hook=_bde_keys_to_int)
23
+ return raw
24
+ except Exception as exc:
25
+ warnings.warn(f'PFASGroups: could not load BDE dict ({exc}); '
26
+ 'falling back to uniform resistance.')
27
+ return None
28
+
29
+
30
+ # ── Hard-coded bond-order scaling model ──────────────────────────────────────
31
+ # Best model selected by Psi4 B3LYP/6-31G* calibration on 134 diatomic
32
+ # molecules (see molecular_quantum_graph/bde_computation/bond_order_calibration/
33
+ # for the full analysis). Model: poly2 f(n) = 1 + a*(n-1) + b*(n-1)^2
34
+ _BOND_ORDER_MODEL_NAME = 'poly2'
35
+ _BOND_ORDER_MODEL_PARAMS = {'a': 1.2650122708517233, 'b': -0.3142013833397031}
36
+
37
+
38
+ def _bond_order_factor(n: float, model_name: str, params: dict) -> float:
39
+ """Evaluate the bond-order scaling factor f(n), f(1)=1."""
40
+ x = n - 1.0
41
+ if model_name == 'linear':
42
+ val = 1.0 + params.get('alpha', 0.3) * x
43
+ elif model_name == 'power':
44
+ val = n ** params.get('beta', 0.6)
45
+ elif model_name == 'log':
46
+ val = 1.0 + params.get('a', 1.0) * math.log(max(n, 1e-10))
47
+ elif model_name == 'poly2':
48
+ val = 1.0 + params.get('a', 0.3) * x + params.get('b', 0.0) * x ** 2
49
+ elif model_name == 'poly3':
50
+ val = (1.0
51
+ + params.get('a', 0.3) * x
52
+ + params.get('b', 0.0) * x ** 2
53
+ + params.get('c', 0.0) * x ** 3)
54
+ else:
55
+ val = 1.0 + 0.3 * x
56
+ return max(val, 1e-6)
57
+
58
+
59
+ class _BDEScheme:
60
+ """Lightweight BDE weighting scheme bundled with PFASGroups."""
61
+
62
+ def __init__(self):
63
+ self.bde_dict = _load_bde_dict_local()
64
+ self._model_name = _BOND_ORDER_MODEL_NAME
65
+ self._model_params = _BOND_ORDER_MODEL_PARAMS
66
+ # C-C single-bond BDE as normalisation reference (kcal/mol)
67
+ self.ref_bde = (
68
+ self.bde_dict.get(6, {}).get(6, 83.1)
69
+ if self.bde_dict else 83.1
70
+ )
71
+
72
+ def conductance(self, z1: int, z2: int, bond_order: float = 1.0) -> float:
73
+ """Return BDE conductance = BDE(z1,z2,order) / ref_bde.
74
+
75
+ Higher BDE → stronger bond → higher conductance → shorter resistance path.
76
+ Returns 1.0 (uniform) if BDE data is unavailable.
77
+ """
78
+ if self.bde_dict is None:
79
+ return 1.0
80
+ try:
81
+ base = self.bde_dict[z1][z2]
82
+ except KeyError:
83
+ try:
84
+ base = self.bde_dict[z2][z1]
85
+ except KeyError:
86
+ b1 = self.bde_dict.get(z1, {}).get(6, 80.0)
87
+ b2 = self.bde_dict.get(z2, {}).get(6, 80.0)
88
+ base = (b1 + b2) / 2.0
89
+ bde = base * _bond_order_factor(bond_order, self._model_name, self._model_params)
90
+ return bde / self.ref_bde
91
+
92
+
93
+ # Module-level singleton — loaded once per process
94
+ _BDE_SCHEME: '_BDEScheme | None' = None
95
+
96
+
97
+ def _get_bde_scheme() -> _BDEScheme:
98
+ global _BDE_SCHEME
99
+ if _BDE_SCHEME is None:
100
+ _BDE_SCHEME = _BDEScheme()
101
+ return _BDE_SCHEME
102
+
103
+ class ComponentsSolver:
104
+ """Class to hold components information with comprehensive graph metrics."""
105
+ @add_componentSmarts()
106
+ def __init__(self, mol, **kwargs):
107
+ self.componentSmartss = kwargs.get('componentSmartss')
108
+ self.mol = mol
109
+ self.mol_size = mol.GetNumAtoms() # Total atoms in molecule for fraction calculation
110
+ self.total_carbons = sum(1 for atom in mol.GetAtoms() if atom.GetSymbol() == 'C') # Total carbon atoms
111
+
112
+ # Per-halogen density metrics (F, Cl, Br, I)
113
+ # Naming convention:
114
+ # total_{sym}s : total count of that halogen (total_fluorines, total_chlorines, …)
115
+ # per{X}ination_density: halogen count / heavy-atom count (perfluorination_density, …)
116
+ # c{sym}2_count : carbons bearing ≥2 of that halogen (cf2_count, ccl2_count, …)
117
+ # c{sym}2_density : c{sym}2_count / total_carbons
118
+ _HALOGEN_META = [
119
+ ('F', 'fluorines', 'perfluorination_density', 'cf2_count', 'cf2_density'),
120
+ ('Cl', 'chlorines', 'perchlorination_density', 'ccl2_count', 'ccl2_density'),
121
+ ('Br', 'bromines', 'perbromination_density', 'cbr2_count', 'cbr2_density'),
122
+ ('I', 'iodines', 'periodination_density', 'ci2_count', 'ci2_density'),
123
+ ]
124
+ for halogen_sym, total_attr, per_attr, cx2_count_attr, cx2_density_attr in _HALOGEN_META:
125
+ total_hal = sum(1 for atom in mol.GetAtoms() if atom.GetSymbol() == halogen_sym)
126
+ setattr(self, 'total_' + total_attr, total_hal)
127
+ setattr(self, per_attr, total_hal / self.mol_size if self.mol_size > 0 else 0.0)
128
+ cx2 = sum(
129
+ 1 for atom in mol.GetAtoms()
130
+ if atom.GetSymbol() == 'C' and
131
+ sum(1 for nb in atom.GetNeighbors() if nb.GetSymbol() == halogen_sym) >= 2
132
+ )
133
+ setattr(self, cx2_count_attr, cx2)
134
+ setattr(self, cx2_density_attr, cx2 / self.total_carbons if self.total_carbons > 0 else 0.0)
135
+ self.G = mol_to_nx(mol)
136
+ self.bde_scheme = _get_bde_scheme()
137
+ self.limit_effective_graph_resistance = kwargs.get('limit_effective_graph_resistance', None)
138
+ self.skip_component_metrics = not kwargs.get('compute_component_metrics', True)
139
+ self.total_branching = self._compute_total_branching()
140
+ self.components = self.get_fluorinated_subgraph()
141
+ self.extended_components = {k:{0:v} for k,v in self.components.items()}
142
+ # Mapping from (pathType, max_dist, extended_component_index) -> original_component_index
143
+ self.component_to_original_index = {}
144
+ self.levels = {0}
145
+ # Cache for component metrics
146
+ self._component_metrics_cache = {}
147
+ # Precompute full component sizes (including all attached atoms: H, F, Cl, Br, I)
148
+ self.component_full_sizes = {}
149
+ for path_type, components_list in self.components.items():
150
+ if path_type not in self.component_full_sizes:
151
+ self.component_full_sizes[path_type] = {}
152
+ for i, comp in enumerate(components_list):
153
+ full_comp = self.get_full_component_atoms(comp)
154
+ self.component_full_sizes[path_type][i] = len(full_comp)
155
+ # Compute and store metrics for all components on creation (if enabled)
156
+ if not self.skip_component_metrics:
157
+ self._precompute_component_metrics()
158
+ # Initialize mapping for level 0 (original components)
159
+ self._init_component_mapping()
160
+
161
+ def __enter__(self):
162
+ return self
163
+ def __exit__(self, exc_type, exc_value, traceback):
164
+ self.components = None
165
+ self.extended_components = None
166
+ self.mol = None
167
+ self.G = None
168
+ self._component_metrics_cache = None
169
+
170
+ def __len__(self):
171
+ return len(self.components)
172
+
173
+ def get(self, pathType, max_dist=0, default = []):
174
+ if max_dist not in self.levels:
175
+ self.extend_components(max_dist)
176
+ return self.extended_components.get(pathType, {}).get(max_dist, default)
177
+
178
+ def max_size(self):
179
+ return max([len(x) for x in self.components]) if len(self.components)>0 else 0
180
+
181
+ def sizes(self):
182
+ return [len(x) for x in self.components]
183
+
184
+ def _connected_components(self, subset):
185
+ """Find connected components in a molecule."""
186
+ G = self.G.subgraph(subset)
187
+ components = list(nx.connected_components(G))
188
+ return components
189
+
190
+ @staticmethod
191
+ def _extract_component_smarts(entry):
192
+ if entry is None:
193
+ return None
194
+ if isinstance(entry, (list, tuple)):
195
+ return entry[0] if len(entry) > 0 else None
196
+ if isinstance(entry, dict):
197
+ return entry.get('smarts', entry.get('component', entry.get('chain')))
198
+ return entry
199
+
200
+ def get_fluorinated_subgraph(self, **kwargs):
201
+ """Get the fluorinated indices by connected components of a molecule based on path SMARTS."""
202
+ subsets = {}
203
+ for pathName, d in self.componentSmartss.items():
204
+ path_smarts = self._extract_component_smarts(d)
205
+ if path_smarts is None:
206
+ subsets[pathName] = []
207
+ continue
208
+ if isinstance(path_smarts, str):
209
+ path_smarts = Chem.MolFromSmarts(path_smarts)
210
+ if path_smarts is None:
211
+ subsets[pathName] = []
212
+ continue
213
+ path_smarts.UpdatePropertyCache()
214
+ Chem.GetSymmSSSR(path_smarts)
215
+ path_smarts.GetRingInfo().NumRings()
216
+ matches = self.mol.GetSubstructMatches(path_smarts)
217
+ subset = [y for x in matches for y in x]
218
+ if len(subset)==0:
219
+ subsets[pathName]= []
220
+ continue
221
+ components = self._connected_components(subset)
222
+ subsets[pathName] = components
223
+ return subsets
224
+
225
+ def _compute_total_branching(self):
226
+ non_hf_atoms = [
227
+ atom.GetIdx()
228
+ for atom in self.mol.GetAtoms()
229
+ if atom.GetSymbol() not in ['H', 'F','Cl','Br','I']
230
+ ]
231
+ if not non_hf_atoms:
232
+ return 0.0
233
+ return calculate_branching(self.mol, non_hf_atoms)
234
+
235
+ def get_full_component_atoms(self, component):
236
+ """Get all atoms in a component including those attached (H, F, halogens).
237
+
238
+ Parameters
239
+ ----------
240
+ component : set
241
+ Set of atom indices representing the carbon backbone
242
+
243
+ Returns
244
+ -------
245
+ set
246
+ Set of all atom indices including the component and all directly attached atoms
247
+ """
248
+ full_component = set(component)
249
+ # Add all neighbors of component atoms that are not already in the component
250
+ # This includes H, F, and other halogens attached to the carbon backbone
251
+ for atom_idx in component:
252
+ for neighbor_idx in self.G.neighbors(atom_idx):
253
+ if self.mol.GetAtomWithIdx(neighbor_idx).GetSymbol() in ['H', 'F', 'Cl', 'Br', 'I']:
254
+ full_component.add(neighbor_idx)
255
+ return full_component
256
+
257
+ def get_total_components_fraction(self, matched_components_list):
258
+ """Calculate the fraction of carbon atoms in the molecule covered by the union of all components.
259
+
260
+ Parameters
261
+ ----------
262
+ matched_components_list : list of dict
263
+ List of matched component dictionaries, each with 'component' and 'smarts_matches' keys
264
+
265
+ Returns
266
+ -------
267
+ float
268
+ Fraction of carbon atoms covered by the union of all components (0.0 to 1.0)
269
+ """
270
+ if len(matched_components_list) == 0 or self.total_carbons == 0:
271
+ return 0.0
272
+
273
+ # Union all carbon atoms from components (augmented components already include
274
+ # SMARTS-match and linker atoms, so just iterate over 'component' atom sets).
275
+ union_carbon_atoms = set()
276
+ for comp_dict in matched_components_list:
277
+ component = set(comp_dict.get('component', []))
278
+ smarts_matches = comp_dict.get('smarts_matches')
279
+
280
+ # Add carbon atoms from component
281
+ for atom_idx in component:
282
+ if self.mol.GetAtomWithIdx(atom_idx).GetSymbol() == 'C':
283
+ union_carbon_atoms.add(atom_idx)
284
+
285
+ # smarts_matches are intentionally excluded: the augmented component already
286
+ # contains those atoms via get_augmented_component. Adding them separately
287
+ # would double-count for OECD groups and overcount for telomers.
288
+
289
+ # Total coverage = union of C atoms from all augmented components.
290
+ # smarts_extra_atoms is intentionally not added here: the augmented components
291
+ # already incorporate SMARTS-match and linker atoms, so no additive correction
292
+ # is needed and it would push telomers above 1.0.
293
+ total_fraction = len(union_carbon_atoms) / self.total_carbons
294
+ return total_fraction
295
+
296
+ def _precompute_component_metrics(self):
297
+ """Precompute metrics for all initial components."""
298
+ for path_type, components_list in self.components.items():
299
+ for comp in components_list:
300
+ # This will populate the cache
301
+ self.compute_component_metrics(comp)
302
+
303
+ def _init_component_mapping(self):
304
+ """Initialize mapping for original components (max_dist=0)."""
305
+ for pathType in self.components.keys():
306
+ for i in range(len(self.components[pathType])):
307
+ self.component_to_original_index[(pathType, 0, i)] = i
308
+
309
+ def extend_components(self, max_dist):
310
+ """Extend a component in a graph by a maximum distance.
311
+ This is used to match functional groups that are not directly connected to the component. Different components that overlap are not merged.
312
+ """
313
+ if max_dist>0:
314
+ for pathType, components in self.components.items():
315
+ extended_components = []
316
+ for i, component in enumerate(components):
317
+ extended = component.copy()
318
+ for node in component:
319
+ lengths = nx.single_source_shortest_path_length(self.G, node, cutoff=max_dist)
320
+ extended.update([n for n,d in lengths.items() if d<=max_dist])
321
+ extended_components.append(extended)
322
+ # Map this extended component back to its original component
323
+ self.component_to_original_index[(pathType, max_dist, i)] = i
324
+ self.extended_components.setdefault(pathType, {})[max_dist] = extended_components
325
+ self.levels.add(max_dist)
326
+ def shortest_path_to_component(self, atom, component):
327
+ """Get shortest paths from SMARTS matches to original component within extended component.
328
+
329
+ Parameters
330
+ ----------
331
+ atom_index : int
332
+ component : set
333
+
334
+ Returns
335
+ -------
336
+ dict
337
+ Mapping from SMARTS match atom index to shortest path list to original component
338
+ """
339
+ path = nx.shortest_path(self.G, atom, list(component)[0])
340
+ path = set(path).difference(component)
341
+ if len(path)==0:
342
+ return None, []
343
+ return [x for x in path if x!=atom]
344
+ def get_augmented_component(self, pathType, max_dist, component_index, smarts_matches, linker_smarts=None):
345
+ """Get original component augmented with connecting atoms to SMARTS matches.
346
+
347
+ Parameters
348
+ ----------
349
+ pathType : str
350
+ Type of component path (e.g., 'Perfluoroalkyl')
351
+ max_dist : int
352
+ Distance used for extension
353
+ component_index : int
354
+ Index of the component in the extended components list
355
+ smarts_matches : set
356
+ Set of atom indices that matched the SMARTS pattern
357
+ linker_smarts : Chem.Mol, optional
358
+ Compiled SMARTS pattern for validating linker atoms.
359
+ If provided, only paths where intermediate atoms match this pattern are accepted.
360
+
361
+ Returns
362
+ -------
363
+ set
364
+ Original component augmented with shortest path atoms connecting SMARTS matches
365
+ """
366
+ if max_dist == 0:
367
+ # No augmentation needed, return original component
368
+ return self.components[pathType][component_index]
369
+
370
+ # Get original and extended components
371
+ orig_index = self.component_to_original_index.get((pathType, max_dist, component_index), component_index)
372
+ orig_comp = self.components[pathType][orig_index]
373
+ ext_comp = self.extended_components[pathType][max_dist][component_index]
374
+
375
+ # Start with original component
376
+ augmented = set(orig_comp)
377
+ linker_matches = set([y for x in self.mol.GetSubstructMatches(linker_smarts) for y in x]) if linker_smarts is not None else []
378
+ # Add shortest paths from SMARTS matches to original component
379
+ for smarts_atom in smarts_matches:
380
+ if smarts_atom in ext_comp and smarts_atom not in orig_comp:
381
+ try:
382
+ linker_atoms = self.shortest_path_to_component(smarts_atom, orig_comp)
383
+
384
+
385
+ # Validate linker atoms only if there are intermediate atoms
386
+ # Direct connections (no linker) are accepted without validation
387
+ if linker_smarts is not None and len(linker_atoms) > 0:
388
+ # Validate the intermediate linker atoms
389
+ if not all(atom in linker_matches for atom in linker_atoms):
390
+ continue # Skip this SMARTS atom if linker validation fails
391
+ except nx.NetworkXNoPath:
392
+ continue # Skip this SMARTS atom if no path exists
393
+ # Add the complete path: SMARTS atom + linker atoms + component border atom
394
+ augmented.update([smarts_atom] + linker_atoms)
395
+ elif smarts_atom in orig_comp:
396
+ # SMARTS atom already in original component
397
+ augmented.add(smarts_atom)
398
+
399
+ # Verify that augmented component still contains the SMARTS matches
400
+ # Count how many SMARTS atoms ended up in the augmented component
401
+ smarts_in_augmented = sum(1 for atom in smarts_matches if atom in augmented)
402
+
403
+ # If no SMARTS atoms made it into the augmented component, reject it
404
+ # This happens when all SMARTS atoms failed linker validation
405
+ if smarts_in_augmented == 0 and len(smarts_matches) > 0:
406
+ return []
407
+
408
+ return augmented
409
+
410
+ def _kirchhoff_index_pinv(self, subG, uniform: bool = False) -> float:
411
+ """Kirchhoff index via Laplacian pseudoinverse. Fast for n < 30."""
412
+ nodes = list(subG.nodes())
413
+ n = len(nodes)
414
+ idx = {node: i for i, node in enumerate(nodes)}
415
+
416
+ L = np.zeros((n, n), dtype=float)
417
+ for u, v, data in subG.edges(data=True):
418
+ if uniform:
419
+ c = 1.0
420
+ else:
421
+ bond_order = data.get('order', 1.0)
422
+ z_u = subG.nodes[u].get('element', 6)
423
+ z_v = subG.nodes[v].get('element', 6)
424
+ c = self.bde_scheme.conductance(z_u, z_v, bond_order)
425
+ i, j = idx[u], idx[v]
426
+ L[i, i] += c
427
+ L[j, j] += c
428
+ L[i, j] -= c
429
+ L[j, i] -= c
430
+
431
+ Lp = np.linalg.pinv(L)
432
+ diag = np.diag(Lp)
433
+ kirchhoff = 0.0
434
+ for i in range(n):
435
+ for j in range(i + 1, n):
436
+ kirchhoff += max(diag[i] + diag[j] - 2.0 * Lp[i, j], 0.0)
437
+ return kirchhoff
438
+
439
+ def _kirchhoff_index(self, subG, uniform: bool = False) -> float:
440
+ """Kirchhoff (effective graph resistance) index via Laplacian eigenspectrum.
441
+
442
+ Equivalent to nx.effective_graph_resistance (Theorem 2.2, Ellens et al.
443
+ 2011) but without the internal G.copy(). Differences from nx:
444
+ - Uses np.linalg.eigvalsh (symmetric-aware, returns sorted reals)
445
+ instead of np.linalg.eigvals; results are identical for symmetric L.
446
+ - Builds the weighted Laplacian directly from edge data rather than
447
+ relying on graph weight attributes, enabling the BDE-weighted variant.
448
+
449
+ Parameters
450
+ ----------
451
+ subG : networkx.Graph
452
+ Subgraph to operate on (original component or 1-hop expanded).
453
+ uniform : bool
454
+ True → all edge conductances = 1 (topological / unweighted).
455
+ False → BDE-calibrated conductances (bond-strength weighted).
456
+
457
+ Returns
458
+ -------
459
+ float
460
+ Sum of all pairwise effective resistance distances.
461
+ """
462
+ n = subG.number_of_nodes()
463
+ nodes = list(subG.nodes())
464
+ idx = {node: i for i, node in enumerate(nodes)}
465
+
466
+ L = np.zeros((n, n), dtype=float)
467
+ for u, v, data in subG.edges(data=True):
468
+ if uniform:
469
+ c = 1.0
470
+ else:
471
+ bond_order = data.get('order', 1.0)
472
+ z_u = subG.nodes[u].get('element', 6)
473
+ z_v = subG.nodes[v].get('element', 6)
474
+ c = self.bde_scheme.conductance(z_u, z_v, bond_order)
475
+ i, j = idx[u], idx[v]
476
+ L[i, i] += c
477
+ L[j, j] += c
478
+ L[i, j] -= c
479
+ L[j, i] -= c
480
+
481
+ # Eigenvalues only; skip the zero eigenvalue (index 0)
482
+ mu = np.sort(np.linalg.eigvalsh(L))
483
+ return float(np.sum(1.0 / mu[1:]) * n)
484
+
485
+ def compute_component_metrics(self, component):
486
+ """Compute comprehensive graph metrics for a component.
487
+
488
+ Parameters
489
+ ----------
490
+ component : set or frozenset
491
+ Set of atom indices in the component
492
+
493
+ Returns
494
+ -------
495
+ dict
496
+ Dictionary with graph metrics including:
497
+ - diameter: maximum eccentricity
498
+ - radius: minimum eccentricity
499
+ - eccentricity_values: dict mapping node to its eccentricity
500
+ - centre: nodes with minimum eccentricity
501
+ - periphery: nodes with maximum eccentricity
502
+ - barycentre: nodes minimizing sum of distances
503
+ - effective_graph_resistance: BDE-weighted Kirchhoff index
504
+ - _rdist: internal BDE-weighted pairwise resistance distance dict
505
+ """
506
+ # If metrics computation is disabled, return minimal metrics (size only)
507
+ if self.skip_component_metrics:
508
+ return {
509
+ 'size': len(component),
510
+ 'diameter': float('nan'),
511
+ 'radius': float('nan'),
512
+ 'eccentricity_values': {},
513
+ 'centre': [],
514
+ 'periphery': [],
515
+ 'barycentre': [],
516
+ 'effective_graph_resistance': float('nan'),
517
+ 'effective_graph_resistance_BDE': float('nan'),
518
+ }
519
+
520
+ cache_key = frozenset(component)
521
+ if cache_key in self._component_metrics_cache:
522
+ return self._component_metrics_cache[cache_key]
523
+
524
+ if len(component) <= 1:
525
+ metrics = {
526
+ 'diameter': 0,
527
+ 'radius': 0,
528
+ 'eccentricity_values': {list(component)[0]: 0} if len(component) == 1 else {},
529
+ 'centre': list(component),
530
+ 'periphery': list(component),
531
+ 'barycentre': list(component),
532
+ 'effective_graph_resistance': 0.0,
533
+ 'effective_graph_resistance_BDE': 0.0,
534
+ }
535
+ self._component_metrics_cache[cache_key] = metrics
536
+ return metrics
537
+
538
+ # Create subgraph for this component
539
+ subG = self.G.subgraph(component)
540
+
541
+ # Check if connected
542
+ if not nx.is_connected(subG):
543
+ metrics = {
544
+ 'diameter': float('inf'),
545
+ 'radius': 0,
546
+ 'eccentricity_values': {},
547
+ 'centre': [],
548
+ 'periphery': [],
549
+ 'barycentre': [],
550
+ 'effective_graph_resistance': float('inf'),
551
+ 'effective_graph_resistance_BDE': float('inf'),
552
+ }
553
+ self._component_metrics_cache[cache_key] = metrics
554
+ return metrics
555
+
556
+ try:
557
+ # Compute eccentricity for each node
558
+ eccentricity_values = nx.eccentricity(subG)
559
+
560
+ # Diameter and radius
561
+ diameter = nx.diameter(subG)
562
+ radius = nx.radius(subG)
563
+
564
+ # Centre and periphery
565
+ centre = nx.center(subG)
566
+ periphery = nx.periphery(subG)
567
+
568
+ # Barycentre: nodes that minimize total distance to all other nodes
569
+ # total_distances = {}
570
+ # for node in subG.nodes():
571
+ # lengths = nx.single_source_shortest_path_length(subG, node)
572
+ # total_distances[node] = sum(lengths.values())
573
+
574
+ # min_total_dist = min(total_distances.values())
575
+ # barycentre = [node for node, dist in total_distances.items() if dist == min_total_dist]
576
+ barycentre = nx.barycenter(subG)
577
+
578
+ # Effective graph resistance (Kirchhoff index) — two variants:
579
+ # uniform: topological (edge weights = 1), original C-skeleton component
580
+ # BDE: bond-strength-weighted, component expanded 1 hop to include F/H/Cl/Br…
581
+ try:
582
+ should_compute_resistance = (
583
+ self.limit_effective_graph_resistance is None or
584
+ (self.limit_effective_graph_resistance > 0
585
+ and len(component) < self.limit_effective_graph_resistance)
586
+ )
587
+
588
+ if should_compute_resistance:
589
+ effective_graph_resistance = self._kirchhoff_index(subG, uniform=True)
590
+ # 1-hop expansion: add all neighbours of every component node
591
+ expanded = set(component)
592
+ for node in list(component):
593
+ expanded.update(self.G.neighbors(node))
594
+ subG_exp = self.G.subgraph(expanded)
595
+ effective_graph_resistance_BDE = self._kirchhoff_index(subG_exp, uniform=False)
596
+ else:
597
+ effective_graph_resistance = float('nan')
598
+ effective_graph_resistance_BDE = float('nan')
599
+ except Exception:
600
+ effective_graph_resistance = float('nan')
601
+ effective_graph_resistance_BDE = float('nan')
602
+
603
+ metrics = {
604
+ 'diameter': diameter,
605
+ 'radius': radius,
606
+ 'eccentricity_values': eccentricity_values,
607
+ 'centre': centre,
608
+ 'periphery': periphery,
609
+ 'barycentre': barycentre,
610
+ 'effective_graph_resistance': effective_graph_resistance,
611
+ 'effective_graph_resistance_BDE': effective_graph_resistance_BDE,
612
+ }
613
+
614
+ except Exception as e:
615
+ # Fallback for any computation errors
616
+ metrics = {
617
+ 'diameter': float('nan'),
618
+ 'radius': float('nan'),
619
+ 'eccentricity_values': {},
620
+ 'centre': [],
621
+ 'periphery': [],
622
+ 'barycentre': [],
623
+ 'effective_graph_resistance': float('nan'),
624
+ 'effective_graph_resistance_BDE': float('nan'),
625
+ }
626
+
627
+ self._component_metrics_cache[cache_key] = metrics
628
+ return metrics
629
+
630
+ def compute_smarts_component_metrics(self, component, smarts_matches):
631
+ """Compute metrics relating SMARTS matches to component structural features.
632
+
633
+ Parameters
634
+ ----------
635
+ component : set
636
+ Set of atom indices in the component
637
+ smarts_matches : set
638
+ Set of atom indices matching the SMARTS pattern
639
+
640
+ Returns
641
+ -------
642
+ dict
643
+ Dictionary with SMARTS-specific metrics, or None if no smarts_matches
644
+ """
645
+ if smarts_matches is None or len(smarts_matches) == 0:
646
+ return None
647
+
648
+ comp_metrics = self.compute_component_metrics(component)
649
+ smarts_in_comp = smarts_matches.intersection(component)
650
+
651
+ if len(smarts_in_comp) == 0 or len(component) <= 1:
652
+ return {
653
+ 'min_dist_to_barycentre': 0,
654
+ 'min_dist_to_centre': 0,
655
+ 'max_dist_to_periphery': 0,
656
+ }
657
+
658
+ subG = self.G.subgraph(component)
659
+
660
+ min_dist_to_barycentre = float('inf')
661
+ min_dist_to_centre = float('inf')
662
+ max_dist_to_periphery = 0
663
+
664
+ try:
665
+ for smarts_node in smarts_in_comp:
666
+ if smarts_node not in subG:
667
+ continue
668
+
669
+ for bc_node in comp_metrics['barycentre']:
670
+ try:
671
+ dist = nx.shortest_path_length(subG, smarts_node, bc_node)
672
+ min_dist_to_barycentre = min(min_dist_to_barycentre, dist)
673
+ except Exception:
674
+ pass
675
+
676
+ for centre_node in comp_metrics['centre']:
677
+ try:
678
+ dist = nx.shortest_path_length(subG, smarts_node, centre_node)
679
+ min_dist_to_centre = min(min_dist_to_centre, dist)
680
+ except Exception:
681
+ pass
682
+
683
+ for periph_node in comp_metrics['periphery']:
684
+ try:
685
+ dist = nx.shortest_path_length(subG, smarts_node, periph_node)
686
+ max_dist_to_periphery = max(max_dist_to_periphery, dist)
687
+ except Exception:
688
+ pass
689
+
690
+ if min_dist_to_barycentre == float('inf'):
691
+ min_dist_to_barycentre = 0
692
+ if min_dist_to_centre == float('inf'):
693
+ min_dist_to_centre = 0
694
+
695
+ except Exception:
696
+ min_dist_to_barycentre = 0
697
+ min_dist_to_centre = 0
698
+ max_dist_to_periphery = 0
699
+
700
+ return {
701
+ 'min_dist_to_barycentre': min_dist_to_barycentre,
702
+ 'min_dist_to_centre': min_dist_to_centre,
703
+ 'max_dist_to_periphery': max_dist_to_periphery,
704
+ }
705
+
706
+ def get_matched_component_dict(self, component, smarts_matches=None, smarts_type='unknown', pfas_group=None, comp_id = None):
707
+ """Get a dictionary with all metrics for a matched component.
708
+
709
+ Parameters
710
+ ----------
711
+ component : set
712
+ Set of atom indices in the component
713
+ smarts_matches : set or None
714
+ Set of atom indices matching the SMARTS pattern (None if no SMARTS)
715
+ smarts_type : str
716
+ Type identifier for the SMARTS pattern
717
+ pfas_group : HalogenGroup or None
718
+ HalogenGroup object with precomputed SMARTS atom counts
719
+
720
+ Returns
721
+ -------
722
+ dict
723
+ Complete dictionary with all component metrics
724
+ """
725
+ # Basic branching metric
726
+ smarts_set = smarts_matches if smarts_matches is not None else set()
727
+ basic_metrics = calculate_component_metrics(self.mol, self.G, component, smarts_set)
728
+
729
+ # Comprehensive graph metrics (cached)
730
+ comp_metrics = self.compute_component_metrics(component)
731
+
732
+ # SMARTS-specific metrics (computed on the fly, None if no smarts)
733
+ smarts_metrics = self.compute_smarts_component_metrics(component, smarts_matches)
734
+
735
+ # Get precomputed SMARTS extra atoms count if pfas_group is available
736
+ smarts_extra_atoms = 0
737
+ if pfas_group is not None and smarts_matches is not None and len(smarts_matches) > 0:
738
+ if pfas_group.smarts_extra_atoms is not None:
739
+ # Sum the extra atoms from all SMARTS patterns
740
+ # For groups with multiple matches, we count each match
741
+ smarts_extra_atoms = pfas_group.component_specific_extra_atoms[comp_id] if comp_id is not None else sum(pfas_group.smarts_extra_atoms) * len(smarts_matches)
742
+
743
+ # Calculate mean and median eccentricity from eccentricity_values
744
+ eccentricity_values = comp_metrics.get('eccentricity_values', {})
745
+ if len(eccentricity_values) > 0:
746
+ ecc_list = list(eccentricity_values.values())
747
+ mean_eccentricity = sum(ecc_list) / len(ecc_list)
748
+ sorted_ecc = sorted(ecc_list)
749
+ n = len(sorted_ecc)
750
+ if n % 2 == 0:
751
+ median_eccentricity = (sorted_ecc[n//2 - 1] + sorted_ecc[n//2]) / 2.0
752
+ else:
753
+ median_eccentricity = sorted_ecc[n//2]
754
+ else:
755
+ mean_eccentricity = 0.0
756
+ median_eccentricity = 0.0
757
+
758
+ # Calculate component fraction based on carbon atoms only
759
+ # Count carbon atoms in component
760
+ component_carbons = sum(1 for atom_idx in component if self.mol.GetAtomWithIdx(atom_idx).GetSymbol() == 'C')
761
+
762
+ # Component-level halogen density metrics (F, Cl, Br, I)
763
+ component_carbons_indices = [idx for idx in component
764
+ if self.mol.GetAtomWithIdx(idx).GetSymbol() == 'C']
765
+ _n_comp_c = len(component_carbons_indices)
766
+ _COMP_HALOGEN_META = [
767
+ ('F', 'component_f_count', 'component_perfluorination_density', 'component_cf2_count', 'component_cf2_density'),
768
+ ('Cl', 'component_cl_count', 'component_perchlorination_density', 'component_ccl2_count', 'component_ccl2_density'),
769
+ ('Br', 'component_br_count', 'component_perbromination_density', 'component_cbr2_count', 'component_cbr2_density'),
770
+ ('I', 'component_i_count', 'component_periodination_density', 'component_ci2_count', 'component_ci2_density'),
771
+ ]
772
+ _comp_hal_vals = {}
773
+ for halogen_sym, cnt_key, per_key, cx2_cnt_key, cx2_den_key in _COMP_HALOGEN_META:
774
+ hal_count = sum(
775
+ 1 for atom_idx in component_carbons_indices
776
+ for nb in self.mol.GetAtomWithIdx(atom_idx).GetNeighbors()
777
+ if nb.GetSymbol() == halogen_sym
778
+ )
779
+ cx2_count = sum(
780
+ 1 for atom_idx in component_carbons_indices
781
+ if sum(1 for nb in self.mol.GetAtomWithIdx(atom_idx).GetNeighbors()
782
+ if nb.GetSymbol() == halogen_sym) >= 2
783
+ )
784
+ _comp_hal_vals[cnt_key] = hal_count
785
+ _comp_hal_vals[per_key] = hal_count / _n_comp_c if _n_comp_c > 0 else 0.0
786
+ _comp_hal_vals[cx2_cnt_key] = cx2_count
787
+ _comp_hal_vals[cx2_den_key] = cx2_count / _n_comp_c if _n_comp_c > 0 else 0.0
788
+ # Convenience aliases for backward compatibility
789
+ component_f_count = _comp_hal_vals['component_f_count']
790
+ component_cf2_count = _comp_hal_vals['component_cf2_count']
791
+
792
+ # component_fraction: ratio of carbon atoms in this component to all carbon atoms in the
793
+ # molecule, counting only C atoms actually present in the (augmented) component atom set.
794
+ # The augmented component already includes linker atoms and SMARTS-match atoms added by
795
+ # get_augmented_component, so no further additive correction is needed.
796
+ # Excluded: O, F, and other heteroatoms that may appear in telomer linker paths.
797
+ component_fraction = component_carbons / self.total_carbons if self.total_carbons > 0 else 0.0
798
+
799
+ # --- n_spacer (telomer CH₂ linker chain length) -------------------------
800
+ # Only meaningful for groups that have a linker_smarts (telomers).
801
+ # Formula: the augmented component includes the orig pfluorinated component
802
+ # PLUS the CH₂ linker atoms PLUS the SMARTS-match atom (the functional group
803
+ # C adjacent to the linker). So n_spacer = |augmented - orig_comp| + 1.
804
+ # For n=1: the SMARTS-match atom IS already in orig_comp, so ΔΔ=0 → spacer=1.
805
+ n_spacer = 0
806
+ if pfas_group is not None and getattr(pfas_group, 'linker_smarts', None) is not None and comp_id is not None:
807
+ max_dist_for_spacer = getattr(pfas_group, 'max_dist_from_comp', 0)
808
+ orig_idx = self.component_to_original_index.get(
809
+ (smarts_type, max_dist_for_spacer, comp_id), comp_id
810
+ )
811
+ if smarts_type in self.components and orig_idx < len(self.components[smarts_type]):
812
+ orig_comp_set = set(self.components[smarts_type][orig_idx])
813
+ n_spacer = len(set(component) - orig_comp_set) + 1
814
+
815
+ # --- ring_size (smallest ring overlapping with component) ---------------
816
+ ring_size = 0
817
+ component_ring_atoms = [
818
+ idx for idx in component if self.mol.GetAtomWithIdx(idx).IsInRing()
819
+ ]
820
+ if component_ring_atoms:
821
+ ring_info = self.mol.GetRingInfo()
822
+ comp_set = set(component)
823
+ for ring in sorted(ring_info.AtomRings(), key=len):
824
+ if comp_set & set(ring):
825
+ ring_size = len(ring)
826
+ break
827
+
828
+ result = {
829
+ 'component': sorted(list(component)),
830
+ 'size': len(component),
831
+ 'component_fraction': component_fraction, # C atoms in augmented component / total C atoms in molecule (always ≤ 1.0)
832
+ 'smarts_matches': sorted(list(smarts_matches)) if smarts_matches is not None else None, # Store for union calculation
833
+ 'smarts_extra_atoms': smarts_extra_atoms, # Extra carbons from functional group
834
+ 'SMARTS': smarts_type,
835
+ # Telomer spacer and cyclic ring metrics
836
+ 'n_spacer': n_spacer,
837
+ 'ring_size': ring_size,
838
+ # Basic metrics
839
+ 'branching': basic_metrics['branching'],
840
+ 'branching_ratio_to_molecule': basic_metrics['branching'] / self.total_branching if self.total_branching > 0 else 0.0,
841
+ 'total_branching': self.total_branching,
842
+ 'smarts_centrality': basic_metrics['smarts_centrality'],
843
+ # Graph structure metrics
844
+ 'diameter': comp_metrics.get('diameter', float('nan')),
845
+ 'radius': comp_metrics.get('radius', float('nan')),
846
+ 'effective_graph_resistance': comp_metrics.get('effective_graph_resistance', float('nan')),
847
+ 'effective_graph_resistance_BDE': comp_metrics.get('effective_graph_resistance_BDE', float('nan')),
848
+ 'eccentricity_values': comp_metrics.get('eccentricity_values', {}),
849
+ 'mean_eccentricity': mean_eccentricity,
850
+ 'median_eccentricity': median_eccentricity,
851
+ 'centre': comp_metrics.get('centre', []),
852
+ 'periphery': comp_metrics.get('periphery', []),
853
+ 'barycentre': comp_metrics.get('barycentre', []),
854
+ # Distance metrics (with defaults)
855
+ 'min_dist_to_barycentre': 0,
856
+ 'min_dist_to_centre': 0,
857
+ 'max_dist_to_periphery': 0,
858
+ # Halogen density metrics — component-level (F, Cl, Br, I)
859
+ **_comp_hal_vals,
860
+ # Global molecule-level halogen density (for context)
861
+ 'molecule_perfluorination_density': self.perfluorination_density,
862
+ 'molecule_cf2_density': self.cf2_density,
863
+ 'molecule_perchlorination_density': self.perchlorination_density,
864
+ 'molecule_ccl2_density': self.ccl2_density,
865
+ 'molecule_perbromination_density': self.perbromination_density,
866
+ 'molecule_cbr2_density': self.cbr2_density,
867
+ 'molecule_periodination_density': self.periodination_density,
868
+ 'molecule_ci2_density': self.ci2_density,
869
+ }
870
+
871
+ # Override distance metrics with SMARTS-specific values if available
872
+ if smarts_metrics is not None:
873
+ result.update({
874
+ 'min_dist_to_barycentre': smarts_metrics.get('min_dist_to_barycentre', 0),
875
+ 'min_dist_to_centre': smarts_metrics.get('min_dist_to_centre', 0),
876
+ 'max_dist_to_periphery': smarts_metrics.get('max_dist_to_periphery', 0),
877
+ })
878
+
879
+ return result
880
+
881
+
882
+ def calculate_branching(mol, subset=None):
883
+ """
884
+ Calculate branching of carbon atoms in a molecules
885
+ Measure of branching vs linearity
886
+ For linear chains: branching → 1.0
887
+ For highly branched: branching → 0.0
888
+
889
+ 1. extract carbon atoms and their connectivity
890
+ 2. Count branch points (degree > 2 in the carbon subgraph)
891
+ 3. Normalize by component size
892
+ """
893
+ # Count branch points (degree > 2 in the carbon subgraph)
894
+ if isinstance(mol, Chem.Mol):
895
+ try:
896
+ carbon_nodes = [atom.GetIdx() for atom in mol.GetAtoms() if atom.GetSymbol() == 'C' and (subset is None or atom.GetIdx() in subset)]
897
+ # Only count C-C bonds so functional-group heteroatoms (e.g. -COOH oxygens)
898
+ # do not create false branch points on terminal carbons.
899
+ branch_points = sum(
900
+ max(0, sum(1 for nb in mol.GetAtomWithIdx(node).GetNeighbors() if nb.GetSymbol() == 'C') - 2) for node in carbon_nodes
901
+ )
902
+ except:
903
+ return 0.0
904
+ elif isinstance(mol, nx.Graph):
905
+ try:
906
+ subG = mol.subgraph(subset) if subset is not None else mol
907
+ # Work on carbon nodes only; count only C-C edges to avoid
908
+ # heteroatom bonds (e.g. C=O in carboxyl groups) inflating degree.
909
+ carbon_nodes = [n for n in subG.nodes() if subG.nodes[n].get('symbol') == 'C']
910
+ branch_points = sum(
911
+ max(0, sum(1 for nb in subG.neighbors(node) if subG.nodes[nb].get('symbol') == 'C') - 2) for node in carbon_nodes
912
+ )
913
+ except:
914
+ return 0.0
915
+ else:
916
+ raise TypeError(f"Unsupported molecule type for branching calculation: {type(mol)}")
917
+ # Normalize by component size
918
+ branching = 1 - 2 * branch_points / max(1, len(carbon_nodes))
919
+ branching = max(0.0, min(1.0, branching)) # Clamp to [0, 1]
920
+ return branching
921
+
922
+ def calculate_component_metrics(mol, G, component, smarts_matches):
923
+ """Calculate branching and centrality metrics for a component.
924
+
925
+ Parameters
926
+ ----------
927
+ G : networkx.Graph
928
+ Molecular graph
929
+ component : set
930
+ Set of atom indices in the component
931
+ smarts_matches : set
932
+ Set of atom indices matching the SMARTS pattern
933
+
934
+ Returns
935
+ -------
936
+ dict
937
+ Dictionary with 'branching' (float) and 'smarts_centrality' (float)
938
+ """
939
+ if len(component) <= 1:
940
+ return {'branching': 0.0, 'smarts_centrality': 1.0}
941
+
942
+ # Create subgraph for this component
943
+ subG = G.subgraph(component)
944
+
945
+ # Calculate branching: fraction of nodes with degree > 2
946
+ branching = calculate_branching(mol, component)
947
+
948
+ # Calculate SMARTS centrality: how central the matched atoms are
949
+ smarts_in_component = smarts_matches.intersection(component)
950
+ if len(smarts_in_component) == 0:
951
+ smarts_centrality = 0.0
952
+ else:
953
+ try:
954
+ # Calculate average shortest path distance from SMARTS matches to all other nodes
955
+ total_distance = 0
956
+ count = 0
957
+ for smarts_node in smarts_in_component:
958
+ if smarts_node in subG:
959
+ lengths = nx.single_source_shortest_path_length(subG, smarts_node)
960
+ for node, dist in lengths.items():
961
+ if node != smarts_node:
962
+ total_distance += dist
963
+ count += 1
964
+
965
+ if count > 0:
966
+ avg_distance = total_distance / count
967
+ # Calculate maximum possible average distance (for peripheral node)
968
+ # For a linear chain of n nodes, max avg distance is ~n/3
969
+ max_possible_distance = len(component) / 3.0
970
+ # Centrality: 1.0 = central, 0.0 = peripheral
971
+ smarts_centrality = 1.0 - min(1.0, avg_distance / max(1.0, max_possible_distance))
972
+ else:
973
+ smarts_centrality = 1.0
974
+ except:
975
+ smarts_centrality = 0.5
976
+
977
+ return {'branching': branching, 'smarts_centrality': smarts_centrality}