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.
- HalogenGroups/__init__.py +246 -0
- PFASGroups/ComponentsSolverModel.py +977 -0
- PFASGroups/HalogenGroupModel.py +810 -0
- PFASGroups/PFASDefinitionModel.py +393 -0
- PFASGroups/PFASEmbeddings.py +3315 -0
- PFASGroups/__init__.py +21 -0
- PFASGroups/cli.py +618 -0
- PFASGroups/core.py +415 -0
- PFASGroups/data/Halogen_groups_smarts.json +9024 -0
- PFASGroups/data/PFAS_definitions_smarts.json +170 -0
- PFASGroups/data/component_smarts.json +4 -0
- PFASGroups/data/component_smarts_halogens.json +142 -0
- PFASGroups/data/diatomic_bonds_dict.json +12802 -0
- PFASGroups/draw_mols.py +374 -0
- PFASGroups/embeddings.py +150 -0
- PFASGroups/fragmentation.py +548 -0
- PFASGroups/generate_homologues.py +256 -0
- PFASGroups/generate_mol.py +656 -0
- PFASGroups/generate_paper_figures.py +266 -0
- PFASGroups/getter.py +111 -0
- PFASGroups/homologue_series.py +473 -0
- PFASGroups/parser.py +942 -0
- PFASGroups/prioritise.py +439 -0
- pfasgroups-3.2.2.dist-info/METADATA +724 -0
- pfasgroups-3.2.2.dist-info/RECORD +28 -0
- pfasgroups-3.2.2.dist-info/WHEEL +5 -0
- pfasgroups-3.2.2.dist-info/entry_points.txt +3 -0
- pfasgroups-3.2.2.dist-info/top_level.txt +2 -0
|
@@ -0,0 +1,3315 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from dataclasses import dataclass
|
|
4
|
+
from typing import TYPE_CHECKING, Any, Dict, Iterable, Iterator, List, Optional, Sequence, Tuple, Union
|
|
5
|
+
import os
|
|
6
|
+
import re
|
|
7
|
+
from pathlib import Path
|
|
8
|
+
from io import BytesIO
|
|
9
|
+
|
|
10
|
+
if TYPE_CHECKING:
|
|
11
|
+
try:
|
|
12
|
+
import sqlalchemy
|
|
13
|
+
except ImportError:
|
|
14
|
+
pass
|
|
15
|
+
|
|
16
|
+
from rdkit import Chem
|
|
17
|
+
from rdkit.Chem import Draw
|
|
18
|
+
from PIL import Image
|
|
19
|
+
|
|
20
|
+
import pandas as pd
|
|
21
|
+
import numpy as np
|
|
22
|
+
|
|
23
|
+
from typing import Callable
|
|
24
|
+
|
|
25
|
+
def _load_palette() -> List[str]:
|
|
26
|
+
"""Load hex colours from color_scheme.yaml (stdlib only, no pyyaml needed)."""
|
|
27
|
+
_defaults = ["#E15D0B", "#306DBA", "#9D206C", "#51127C"]
|
|
28
|
+
try:
|
|
29
|
+
_p = Path(__file__).parent / "data" / "color_scheme.yaml"
|
|
30
|
+
_colors = re.findall(r'"(#[0-9A-Fa-f]{6})"', _p.read_text())
|
|
31
|
+
if len(_colors) >= 4:
|
|
32
|
+
return _colors[:4]
|
|
33
|
+
except Exception:
|
|
34
|
+
pass
|
|
35
|
+
return _defaults
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
_PALETTE = _load_palette()
|
|
39
|
+
# C0=orange, C1=blue (FG table), C2=magenta (metrics table), C3=dark-purple
|
|
40
|
+
_C0, _C1, _C2, _C3 = _PALETTE
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
def _hex_to_rgb_float(h: str) -> Tuple[float, float, float]:
|
|
44
|
+
"""Convert a '#RRGGBB' hex string to an RGB float triple in [0, 1]."""
|
|
45
|
+
h = h.lstrip('#')
|
|
46
|
+
return tuple(int(h[i:i+2], 16) / 255.0 for i in (0, 2, 4)) # type: ignore
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
def _lighter(h: str, factor: float = 0.82) -> str:
|
|
50
|
+
"""Return a lighter hex colour by blending *h* with white."""
|
|
51
|
+
r, g, b = _hex_to_rgb_float(h)
|
|
52
|
+
r2 = int((r + (1 - r) * factor) * 255)
|
|
53
|
+
g2 = int((g + (1 - g) * factor) * 255)
|
|
54
|
+
b2 = int((b + (1 - b) * factor) * 255)
|
|
55
|
+
return f'#{r2:02X}{g2:02X}{b2:02X}'
|
|
56
|
+
|
|
57
|
+
# ---------------------------------------------------------------------------
|
|
58
|
+
# Sentinel for "argument not supplied" (distinct from None)
|
|
59
|
+
# ---------------------------------------------------------------------------
|
|
60
|
+
_UNSET = object()
|
|
61
|
+
|
|
62
|
+
# ---------------------------------------------------------------------------
|
|
63
|
+
# Embedding helpers (used by PFASEmbedding.to_array / PFASEmbeddingSet.to_array)
|
|
64
|
+
# ---------------------------------------------------------------------------
|
|
65
|
+
|
|
66
|
+
def _finite(v) -> bool:
|
|
67
|
+
"""True when v is a real, finite number (not NaN/Inf)."""
|
|
68
|
+
try:
|
|
69
|
+
return v == v and abs(v) != float('inf')
|
|
70
|
+
except TypeError:
|
|
71
|
+
return False
|
|
72
|
+
|
|
73
|
+
|
|
74
|
+
def _agg(comps: list, metric: str, aggregation: str) -> float:
|
|
75
|
+
"""Aggregate *metric* over component dicts using mean or median."""
|
|
76
|
+
vals = [c[metric] for c in comps if c.get(metric) is not None and _finite(c[metric])]
|
|
77
|
+
if not vals:
|
|
78
|
+
return 0.0
|
|
79
|
+
if aggregation == 'median':
|
|
80
|
+
return float(np.median(vals))
|
|
81
|
+
return float(np.mean(vals))
|
|
82
|
+
|
|
83
|
+
|
|
84
|
+
def _encode_count(match, mode: str) -> float:
|
|
85
|
+
"""Scalar encoding for one matched group given count *mode*."""
|
|
86
|
+
if match is None or not match.get('components'):
|
|
87
|
+
return 0.0
|
|
88
|
+
return _encode_count_comps(match['components'], mode)
|
|
89
|
+
|
|
90
|
+
|
|
91
|
+
def _encode_count_comps(comps: list, mode: str) -> float:
|
|
92
|
+
"""Like :func:`_encode_count` but accepts a pre-filtered component list."""
|
|
93
|
+
if not comps:
|
|
94
|
+
return 0.0
|
|
95
|
+
if mode == 'binary':
|
|
96
|
+
return 1.0
|
|
97
|
+
if mode == 'count':
|
|
98
|
+
return float(len(comps))
|
|
99
|
+
if mode == 'max_component':
|
|
100
|
+
return float(max(c.get('size', 0) for c in comps))
|
|
101
|
+
if mode == 'total_component':
|
|
102
|
+
return float(sum(c.get('size', 0) for c in comps))
|
|
103
|
+
return 0.0
|
|
104
|
+
|
|
105
|
+
|
|
106
|
+
def _mol_metric(all_comps: list, metric: str) -> float:
|
|
107
|
+
"""Molecule-wide scalar for *metric* over all matched components."""
|
|
108
|
+
if metric == 'n_components':
|
|
109
|
+
return float(len(all_comps))
|
|
110
|
+
if metric == 'total_size':
|
|
111
|
+
return float(sum(c.get('size', 0) or 0 for c in all_comps))
|
|
112
|
+
if metric == 'mean_size':
|
|
113
|
+
sizes = [c.get('size', 0) or 0 for c in all_comps]
|
|
114
|
+
return float(np.mean(sizes)) if sizes else 0.0
|
|
115
|
+
if metric == 'max_size':
|
|
116
|
+
sizes = [c.get('size', 0) or 0 for c in all_comps]
|
|
117
|
+
return float(max(sizes)) if sizes else 0.0
|
|
118
|
+
if metric == 'mean_branching':
|
|
119
|
+
return _agg(all_comps, 'branching', 'mean')
|
|
120
|
+
if metric == 'max_branching':
|
|
121
|
+
vals = [c.get('branching') for c in all_comps if c.get('branching') is not None]
|
|
122
|
+
return float(max(vals)) if vals else 0.0
|
|
123
|
+
if metric == 'mean_eccentricity':
|
|
124
|
+
return _agg(all_comps, 'mean_eccentricity', 'mean')
|
|
125
|
+
if metric == 'max_diameter':
|
|
126
|
+
vals = [c.get('diameter') for c in all_comps if c.get('diameter') is not None]
|
|
127
|
+
return float(max(vals)) if vals else 0.0
|
|
128
|
+
if metric == 'mean_component_fraction':
|
|
129
|
+
return _agg(all_comps, 'component_fraction', 'mean')
|
|
130
|
+
if metric == 'max_component_fraction':
|
|
131
|
+
vals = [c.get('component_fraction') for c in all_comps if c.get('component_fraction') is not None]
|
|
132
|
+
return float(max(vals)) if vals else 0.0
|
|
133
|
+
return 0.0
|
|
134
|
+
|
|
135
|
+
|
|
136
|
+
def _select_groups(group_selection, selected_group_ids, pfas_groups):
|
|
137
|
+
"""Translate group_selection / selected_group_ids to a list of group objects."""
|
|
138
|
+
id_to_group = {g.id: g for g in pfas_groups}
|
|
139
|
+
|
|
140
|
+
if selected_group_ids is not None:
|
|
141
|
+
return [id_to_group[gid] for gid in selected_group_ids if gid in id_to_group]
|
|
142
|
+
|
|
143
|
+
if group_selection is None or group_selection == 'all':
|
|
144
|
+
return list(pfas_groups)
|
|
145
|
+
if group_selection == 'oecd':
|
|
146
|
+
return [id_to_group[gid] for gid in range(1, 29) if gid in id_to_group]
|
|
147
|
+
if group_selection == 'generic':
|
|
148
|
+
return [id_to_group[gid] for gid in range(29, 77) if gid in id_to_group]
|
|
149
|
+
if group_selection == 'telomers':
|
|
150
|
+
return [id_to_group[gid] for gid in range(77, 119) if gid in id_to_group]
|
|
151
|
+
if group_selection == 'generic+telomers':
|
|
152
|
+
ids = list(range(29, 119))
|
|
153
|
+
return [id_to_group[gid] for gid in ids if gid in id_to_group]
|
|
154
|
+
|
|
155
|
+
from .getter import get_HalogenGroups
|
|
156
|
+
raw = get_HalogenGroups()
|
|
157
|
+
compute_raw = [g for g in raw if g.get('compute', True)]
|
|
158
|
+
matching_ids = {
|
|
159
|
+
g['id'] for g in compute_raw
|
|
160
|
+
if g.get('test', {}).get('category', 'other') == group_selection
|
|
161
|
+
}
|
|
162
|
+
if matching_ids:
|
|
163
|
+
return [id_to_group[gid] for gid in matching_ids if gid in id_to_group]
|
|
164
|
+
|
|
165
|
+
raise ValueError(
|
|
166
|
+
f"Unknown group_selection: {group_selection!r}. "
|
|
167
|
+
f"Choose from: 'all', 'oecd', 'generic', 'telomers', 'generic+telomers'"
|
|
168
|
+
)
|
|
169
|
+
|
|
170
|
+
Colour = Tuple[float, float, float]
|
|
171
|
+
Color = Colour # US alias
|
|
172
|
+
|
|
173
|
+
# ---------------------------------------------------------------------------
|
|
174
|
+
# ANSI terminal styling helpers
|
|
175
|
+
# ---------------------------------------------------------------------------
|
|
176
|
+
|
|
177
|
+
_ANSI_RESET = "\033[0m"
|
|
178
|
+
_ANSI_BOLD = "\033[1m"
|
|
179
|
+
|
|
180
|
+
_ANSI_HALOGEN: Dict[str, str] = {
|
|
181
|
+
"F": "\033[96m", # bright cyan
|
|
182
|
+
"Cl": "\033[92m", # bright green
|
|
183
|
+
"Br": "\033[93m", # bright yellow
|
|
184
|
+
"I": "\033[95m", # bright magenta
|
|
185
|
+
}
|
|
186
|
+
_ANSI_FORM: Dict[str, str] = {
|
|
187
|
+
"alkyl": "\033[34m", # blue
|
|
188
|
+
"cyclic": "\033[35m", # magenta
|
|
189
|
+
}
|
|
190
|
+
_ANSI_SAT: Dict[str, str] = {
|
|
191
|
+
"per": "\033[31m", # red
|
|
192
|
+
"poly": "\033[33m", # yellow
|
|
193
|
+
}
|
|
194
|
+
|
|
195
|
+
def _ansi(text: str, *codes: str) -> str:
|
|
196
|
+
"""Wrap *text* with ANSI escape codes and reset afterwards."""
|
|
197
|
+
return "".join(codes) + text + _ANSI_RESET
|
|
198
|
+
|
|
199
|
+
|
|
200
|
+
# ---------------------------------------------------------------------------
|
|
201
|
+
# Molecule-highlight colour palettes (RGB float triples, 0–1)
|
|
202
|
+
# ---------------------------------------------------------------------------
|
|
203
|
+
|
|
204
|
+
# Highlight colours by halogen element — mapped to the project colour palette
|
|
205
|
+
_HALOGEN_COLOURS: Dict[str, Colour] = {
|
|
206
|
+
"F": _hex_to_rgb_float(_C1), # blue (most common PFAS element)
|
|
207
|
+
"Cl": _hex_to_rgb_float(_C0), # orange
|
|
208
|
+
"Br": _hex_to_rgb_float(_C2), # magenta
|
|
209
|
+
"I": _hex_to_rgb_float(_C3), # dark purple
|
|
210
|
+
}
|
|
211
|
+
_HALOGEN_COLORS = _HALOGEN_COLOURS # US alias
|
|
212
|
+
_HALOGEN_COLOUR_DEFAULT: Colour = (0.75, 0.75, 0.75) # grey for unknown
|
|
213
|
+
_HALOGEN_COLOR_DEFAULT = _HALOGEN_COLOUR_DEFAULT # US alias
|
|
214
|
+
|
|
215
|
+
# Tint modifiers for form (shift hue slightly)
|
|
216
|
+
_FORM_TINT: Dict[str, Tuple[float, float, float]] = {
|
|
217
|
+
"alkyl": (0.0, 0.0, 0.0), # no change
|
|
218
|
+
"cyclic": (0.12, -0.08, 0.05), # warm shift
|
|
219
|
+
}
|
|
220
|
+
|
|
221
|
+
# Brightness modifiers for saturation
|
|
222
|
+
_SAT_BRIGHTNESS: Dict[str, float] = {
|
|
223
|
+
"per": 1.0, # full saturation → bright
|
|
224
|
+
"poly": 0.65, # partial saturation → dimmer
|
|
225
|
+
}
|
|
226
|
+
|
|
227
|
+
def _component_colour(halogen: Optional[str], form: Optional[str], saturation: Optional[str]) -> Colour:
|
|
228
|
+
"""Return an RGB highlight colour encoding halogen, form and saturation."""
|
|
229
|
+
base = _HALOGEN_COLOURS.get(halogen or "", _HALOGEN_COLOUR_DEFAULT)
|
|
230
|
+
tint = _FORM_TINT.get(form or "", (0.0, 0.0, 0.0))
|
|
231
|
+
brightness = _SAT_BRIGHTNESS.get(saturation or "", 0.85)
|
|
232
|
+
r = min(1.0, max(0.0, (base[0] + tint[0]) * brightness))
|
|
233
|
+
g = min(1.0, max(0.0, (base[1] + tint[1]) * brightness))
|
|
234
|
+
b = min(1.0, max(0.0, (base[2] + tint[2]) * brightness))
|
|
235
|
+
return (r, g, b)
|
|
236
|
+
_component_color = _component_colour # US alias
|
|
237
|
+
|
|
238
|
+
|
|
239
|
+
# Simple color palette to distinguish PFAS groups in highlight plots (legacy fallback)
|
|
240
|
+
_GROUP_COLOURS: List[Colour] = [
|
|
241
|
+
_hex_to_rgb_float(_C0), # orange
|
|
242
|
+
_hex_to_rgb_float(_C1), # blue
|
|
243
|
+
_hex_to_rgb_float(_C2), # magenta
|
|
244
|
+
_hex_to_rgb_float(_C3), # dark purple
|
|
245
|
+
(0.00, 0.70, 0.70), # teal (extra)
|
|
246
|
+
(0.60, 0.60, 0.60), # grey (extra)
|
|
247
|
+
]
|
|
248
|
+
_GROUP_COLORS = _GROUP_COLOURS # US alias
|
|
249
|
+
|
|
250
|
+
# ---------------------------------------------------------------------------
|
|
251
|
+
# Component-SMARTS metadata cache
|
|
252
|
+
# ---------------------------------------------------------------------------
|
|
253
|
+
|
|
254
|
+
_COMPONENT_META_CACHE: Optional[Dict[str, Dict[str, Optional[str]]]] = None
|
|
255
|
+
|
|
256
|
+
|
|
257
|
+
def _get_component_meta() -> Dict[str, Dict[str, Optional[str]]]:
|
|
258
|
+
"""Return a dict mapping component SMARTS name → {halogen, form, saturation}.
|
|
259
|
+
|
|
260
|
+
Built from the same preprocessed component dictionary used by the parser,
|
|
261
|
+
so names are guaranteed to match the SMARTS labels stored in match results.
|
|
262
|
+
"""
|
|
263
|
+
global _COMPONENT_META_CACHE
|
|
264
|
+
if _COMPONENT_META_CACHE is not None:
|
|
265
|
+
return _COMPONENT_META_CACHE
|
|
266
|
+
try:
|
|
267
|
+
from .core import get_componentSMARTSs
|
|
268
|
+
raw = get_componentSMARTSs()
|
|
269
|
+
_COMPONENT_META_CACHE = {
|
|
270
|
+
name: {
|
|
271
|
+
"halogen": info.get("halogen"),
|
|
272
|
+
"form": info.get("form"),
|
|
273
|
+
"saturation": info.get("saturation"),
|
|
274
|
+
}
|
|
275
|
+
for name, info in raw.items()
|
|
276
|
+
}
|
|
277
|
+
except Exception:
|
|
278
|
+
_COMPONENT_META_CACHE = {}
|
|
279
|
+
return _COMPONENT_META_CACHE
|
|
280
|
+
|
|
281
|
+
|
|
282
|
+
# ---------------------------------------------------------------------------
|
|
283
|
+
# Group-info cache and classification helpers
|
|
284
|
+
# ---------------------------------------------------------------------------
|
|
285
|
+
|
|
286
|
+
_GROUP_INFO_CACHE: Optional[Dict[int, Dict[str, str]]] = None
|
|
287
|
+
|
|
288
|
+
|
|
289
|
+
def _get_group_info() -> Dict[int, Dict[str, str]]:
|
|
290
|
+
"""Return a dict mapping group_id → {name, category}.
|
|
291
|
+
|
|
292
|
+
Category is one of ``'OECD'``, ``'generic'``, ``'telomer'``, or ``'other'``.
|
|
293
|
+
"""
|
|
294
|
+
global _GROUP_INFO_CACHE
|
|
295
|
+
if _GROUP_INFO_CACHE is not None:
|
|
296
|
+
return _GROUP_INFO_CACHE
|
|
297
|
+
try:
|
|
298
|
+
from .getter import get_HalogenGroups
|
|
299
|
+
raw = get_HalogenGroups()
|
|
300
|
+
_GROUP_INFO_CACHE = {
|
|
301
|
+
g["id"]: {
|
|
302
|
+
"name": g.get("name", ""),
|
|
303
|
+
"category": g.get("test", {}).get("category", "other"),
|
|
304
|
+
}
|
|
305
|
+
for g in raw
|
|
306
|
+
}
|
|
307
|
+
except Exception:
|
|
308
|
+
_GROUP_INFO_CACHE = {}
|
|
309
|
+
return _GROUP_INFO_CACHE
|
|
310
|
+
|
|
311
|
+
|
|
312
|
+
# Groups always excluded from the non-OECD category label (perhalogenated /
|
|
313
|
+
# polyhalogenated alkyl catch-alls that add no structural specificity).
|
|
314
|
+
_CLASSIFY_EXCLUDED_IDS: frozenset = frozenset({51, 52})
|
|
315
|
+
|
|
316
|
+
# Name subsumption for non-OECD classification: when a more-specific group
|
|
317
|
+
# name (key) is present, every name in its list is suppressed.
|
|
318
|
+
# Keys and values must match exactly the ``group_name`` values stored in
|
|
319
|
+
# match results (i.e. the ``name`` field from the raw group JSON).
|
|
320
|
+
_SUBSUMES: Dict[str, List[str]] = {
|
|
321
|
+
"sulfonamide": ["amine"],
|
|
322
|
+
"amide": ["amine"],
|
|
323
|
+
"Telomer sulfonamide": ["amine"],
|
|
324
|
+
"phosphonamide": ["amine"],
|
|
325
|
+
}
|
|
326
|
+
|
|
327
|
+
|
|
328
|
+
def _grid_images(imgs: Sequence[Image.Image], buffer: int = 4, ncols: int = 3) -> Tuple[Image.Image, int, int]:
|
|
329
|
+
"""Arrange PIL images in a simple grid layout.
|
|
330
|
+
|
|
331
|
+
Parameters
|
|
332
|
+
----------
|
|
333
|
+
imgs : sequence of PIL.Image
|
|
334
|
+
Images to arrange.
|
|
335
|
+
buffer : int, default 4
|
|
336
|
+
Spacing between images in pixels.
|
|
337
|
+
ncols : int, default 3
|
|
338
|
+
Number of columns.
|
|
339
|
+
"""
|
|
340
|
+
if not imgs:
|
|
341
|
+
raise ValueError("No images provided to _grid_images")
|
|
342
|
+
|
|
343
|
+
# Compute per-row layout
|
|
344
|
+
rows: List[List[Image.Image]] = [list(imgs[i : i + ncols]) for i in range(0, len(imgs), ncols)]
|
|
345
|
+
max_width = 0
|
|
346
|
+
total_height = 0
|
|
347
|
+
row_heights: List[int] = []
|
|
348
|
+
|
|
349
|
+
for row in rows:
|
|
350
|
+
row_width = sum(im.width for im in row) + buffer * (len(row) - 1 if len(row) > 0 else 0)
|
|
351
|
+
max_width = max(max_width, row_width)
|
|
352
|
+
h = max((im.height for im in row), default=0)
|
|
353
|
+
row_heights.append(h)
|
|
354
|
+
total_height += h + buffer
|
|
355
|
+
|
|
356
|
+
if rows:
|
|
357
|
+
total_height -= buffer # no buffer after last row
|
|
358
|
+
|
|
359
|
+
canvas = Image.new("RGBA", (max_width, total_height), (255, 255, 255, 0))
|
|
360
|
+
|
|
361
|
+
y = 0
|
|
362
|
+
for row, h in zip(rows, row_heights):
|
|
363
|
+
x = 0
|
|
364
|
+
for im in row:
|
|
365
|
+
canvas.paste(im, (x, y))
|
|
366
|
+
x += im.width + buffer
|
|
367
|
+
y += h + buffer
|
|
368
|
+
|
|
369
|
+
return canvas, max_width, total_height
|
|
370
|
+
|
|
371
|
+
|
|
372
|
+
def _mol_image_with_table(
|
|
373
|
+
mol_img: Image.Image,
|
|
374
|
+
entries: List[Tuple[str, str, str, str, str]],
|
|
375
|
+
comp_metrics: Optional[Tuple[str, ...]] = None,
|
|
376
|
+
mol_label: str = "",
|
|
377
|
+
halogen_label: str = "",
|
|
378
|
+
font_size: int = 9,
|
|
379
|
+
) -> Image.Image:
|
|
380
|
+
"""Composite a molecule image (PIL) with two formatted matplotlib tables.
|
|
381
|
+
|
|
382
|
+
Parameters
|
|
383
|
+
----------
|
|
384
|
+
mol_img : PIL Image
|
|
385
|
+
The molecule drawing produced by RDKit (no legend text).
|
|
386
|
+
entries : list of tuples
|
|
387
|
+
Each tuple: (group_name, smarts_label, dist_center, dist_periphery)
|
|
388
|
+
FG-specific metrics per matched component row.
|
|
389
|
+
comp_metrics : tuple of str or None
|
|
390
|
+
Component-wide row: (size, branching, eccentricity, diameter, radius,
|
|
391
|
+
eff_graph_resistance, bde_eff_graph_resistance, chain_pct).
|
|
392
|
+
When provided, rendered in a separate second table below the FG table.
|
|
393
|
+
mol_label : str
|
|
394
|
+
Optional header shown above the table (e.g. "mol#1").
|
|
395
|
+
halogen_label : str
|
|
396
|
+
Optional halogen element symbol (e.g. "F") shown as a badge in the
|
|
397
|
+
top-right corner of the molecule image.
|
|
398
|
+
font_size : int
|
|
399
|
+
Font size for table body text.
|
|
400
|
+
|
|
401
|
+
Returns
|
|
402
|
+
-------
|
|
403
|
+
PIL Image
|
|
404
|
+
Combined molecule + tables image.
|
|
405
|
+
"""
|
|
406
|
+
import matplotlib
|
|
407
|
+
matplotlib.use('Agg')
|
|
408
|
+
import matplotlib.pyplot as plt
|
|
409
|
+
import io as _io
|
|
410
|
+
|
|
411
|
+
dpi = 96
|
|
412
|
+
mol_w, mol_h = mol_img.size
|
|
413
|
+
fig_w_in = mol_w / dpi
|
|
414
|
+
|
|
415
|
+
# Row height in inches
|
|
416
|
+
row_h_in = (font_size + 5) / 72.0
|
|
417
|
+
header_h_in = row_h_in * 1.4
|
|
418
|
+
n_rows = max(1, len(entries))
|
|
419
|
+
label_h_in = (font_size + 4) / 72.0 * 1.3 if mol_label else 0.0
|
|
420
|
+
# Second table height: header + 1 data row (if comp_metrics provided)
|
|
421
|
+
metrics_tbl_h_in = (header_h_in + row_h_in + 0.04) if comp_metrics else 0.0
|
|
422
|
+
table_h_in = label_h_in + header_h_in + n_rows * row_h_in + 0.06 + metrics_tbl_h_in + 0.06
|
|
423
|
+
mol_h_in = mol_h / dpi
|
|
424
|
+
fig_h_in = mol_h_in + table_h_in
|
|
425
|
+
|
|
426
|
+
fig = plt.figure(figsize=(fig_w_in, fig_h_in), dpi=dpi)
|
|
427
|
+
|
|
428
|
+
mol_ratio = mol_h_in / fig_h_in
|
|
429
|
+
ax_mol = fig.add_axes([0, 1 - mol_ratio, 1, mol_ratio])
|
|
430
|
+
ax_mol.imshow(mol_img)
|
|
431
|
+
ax_mol.axis('off')
|
|
432
|
+
if halogen_label:
|
|
433
|
+
ax_mol.text(
|
|
434
|
+
0.98, 0.98, halogen_label,
|
|
435
|
+
ha='right', va='top',
|
|
436
|
+
fontsize=font_size + 1, fontweight='bold',
|
|
437
|
+
color='white',
|
|
438
|
+
bbox=dict(boxstyle='round,pad=0.2', facecolor=_C0, edgecolor='none', alpha=0.85),
|
|
439
|
+
transform=ax_mol.transAxes,
|
|
440
|
+
)
|
|
441
|
+
|
|
442
|
+
tbl_ratio = 1.0 - mol_ratio
|
|
443
|
+
|
|
444
|
+
# --- Layout: split tbl_ratio into label, table1, gap, table2 (bottom-up) ---
|
|
445
|
+
bottom_pad = 0.03 / tbl_ratio # small padding at very bottom (in axes coords)
|
|
446
|
+
metrics_frac = (metrics_tbl_h_in / table_h_in) if comp_metrics else 0.0
|
|
447
|
+
gap_frac = (0.06 / table_h_in) if comp_metrics else 0.0
|
|
448
|
+
fg_tbl_frac = (header_h_in + n_rows * row_h_in + 0.06) / table_h_in
|
|
449
|
+
label_frac = label_h_in / table_h_in
|
|
450
|
+
|
|
451
|
+
ax_tbl = fig.add_axes([0.01, 0, 0.98, tbl_ratio])
|
|
452
|
+
ax_tbl.axis('off')
|
|
453
|
+
ax_tbl.set_xlim(0, 1)
|
|
454
|
+
ax_tbl.set_ylim(0, 1)
|
|
455
|
+
|
|
456
|
+
# y positions (axes coords, from bottom=0 to top=1)
|
|
457
|
+
metrics_bottom = bottom_pad
|
|
458
|
+
metrics_top = metrics_bottom + metrics_frac
|
|
459
|
+
gap_top = metrics_top + gap_frac
|
|
460
|
+
fg_tbl_bottom = gap_top
|
|
461
|
+
fg_tbl_top = fg_tbl_bottom + fg_tbl_frac
|
|
462
|
+
label_bottom = fg_tbl_top
|
|
463
|
+
|
|
464
|
+
if mol_label:
|
|
465
|
+
label_center_y = label_bottom + label_frac / 2.0
|
|
466
|
+
ax_tbl.text(
|
|
467
|
+
0.5, label_center_y, mol_label,
|
|
468
|
+
ha='center', va='center',
|
|
469
|
+
fontsize=font_size + 1, fontweight='bold',
|
|
470
|
+
transform=ax_tbl.transAxes,
|
|
471
|
+
)
|
|
472
|
+
|
|
473
|
+
# --- FG-specific table (table 1) ---
|
|
474
|
+
col_labels_fg = ["Group", "SMARTS", "Dist.Ctr", "Dist.Per"]
|
|
475
|
+
col_widths_fg = [0.35, 0.30, 0.175, 0.175]
|
|
476
|
+
|
|
477
|
+
cell_text_fg = [
|
|
478
|
+
[grp, sls or "\u2014", dct, dpr]
|
|
479
|
+
for grp, sls, dct, dpr, *_ in entries
|
|
480
|
+
] or [["\u2014", "\u2014", "\u2014", "\u2014"]]
|
|
481
|
+
|
|
482
|
+
tbl = ax_tbl.table(
|
|
483
|
+
cellText=cell_text_fg,
|
|
484
|
+
colLabels=col_labels_fg,
|
|
485
|
+
colWidths=col_widths_fg,
|
|
486
|
+
loc='upper center',
|
|
487
|
+
cellLoc='left',
|
|
488
|
+
bbox=[0, fg_tbl_bottom, 1, fg_tbl_frac],
|
|
489
|
+
)
|
|
490
|
+
tbl.auto_set_font_size(False)
|
|
491
|
+
tbl.set_fontsize(font_size)
|
|
492
|
+
|
|
493
|
+
for j in range(len(col_labels_fg)):
|
|
494
|
+
cell = tbl[0, j]
|
|
495
|
+
cell.set_facecolor(_C1)
|
|
496
|
+
cell.set_text_props(color='white', fontweight='bold')
|
|
497
|
+
cell.set_edgecolor(_C1)
|
|
498
|
+
|
|
499
|
+
for i in range(len(cell_text_fg)):
|
|
500
|
+
for j in range(len(col_labels_fg)):
|
|
501
|
+
cell = tbl[i + 1, j]
|
|
502
|
+
cell.set_facecolor(_lighter(_C1) if i % 2 == 0 else 'white')
|
|
503
|
+
cell.set_edgecolor(_lighter(_C1, factor=0.55))
|
|
504
|
+
|
|
505
|
+
# --- Component-wide metrics table (table 2) ---
|
|
506
|
+
if comp_metrics:
|
|
507
|
+
col_labels_m = ["Size", "Branch.", "Ecc.", "ø", "Radius", "Eff.Res.", "% C"]
|
|
508
|
+
col_widths_m = [0.11, 0.13, 0.13, 0.10, 0.12, 0.19, 0.22]
|
|
509
|
+
cell_text_m = [list(comp_metrics)]
|
|
510
|
+
|
|
511
|
+
tbl2 = ax_tbl.table(
|
|
512
|
+
cellText=cell_text_m,
|
|
513
|
+
colLabels=col_labels_m,
|
|
514
|
+
colWidths=col_widths_m,
|
|
515
|
+
loc='upper center',
|
|
516
|
+
cellLoc='center',
|
|
517
|
+
bbox=[0, metrics_bottom, 1, metrics_frac],
|
|
518
|
+
)
|
|
519
|
+
tbl2.auto_set_font_size(False)
|
|
520
|
+
tbl2.set_fontsize(font_size)
|
|
521
|
+
|
|
522
|
+
for j in range(len(col_labels_m)):
|
|
523
|
+
cell = tbl2[0, j]
|
|
524
|
+
cell.set_facecolor(_C2)
|
|
525
|
+
cell.set_text_props(color='white', fontweight='bold')
|
|
526
|
+
cell.set_edgecolor(_C2)
|
|
527
|
+
|
|
528
|
+
for j in range(len(col_labels_m)):
|
|
529
|
+
cell = tbl2[1, j]
|
|
530
|
+
cell.set_facecolor(_lighter(_C2))
|
|
531
|
+
cell.set_edgecolor(_lighter(_C2, factor=0.55))
|
|
532
|
+
|
|
533
|
+
fig.patch.set_facecolor('white')
|
|
534
|
+
buf = _io.BytesIO()
|
|
535
|
+
fig.savefig(buf, format='png', dpi=dpi, bbox_inches='tight', pad_inches=0.04,
|
|
536
|
+
facecolor='white')
|
|
537
|
+
plt.close(fig)
|
|
538
|
+
buf.seek(0)
|
|
539
|
+
return Image.open(buf).copy()
|
|
540
|
+
|
|
541
|
+
|
|
542
|
+
@dataclass
|
|
543
|
+
class ComponentView:
|
|
544
|
+
"""Lightweight wrapper around a matched component dict.
|
|
545
|
+
|
|
546
|
+
This does not change the underlying structure; it only provides
|
|
547
|
+
nicer attribute-style access where useful.
|
|
548
|
+
"""
|
|
549
|
+
|
|
550
|
+
data: Dict[str, Any]
|
|
551
|
+
|
|
552
|
+
@property
|
|
553
|
+
def atoms(self) -> List[int]:
|
|
554
|
+
return self.data.get("component", [])
|
|
555
|
+
|
|
556
|
+
@property
|
|
557
|
+
def smarts_label(self) -> Optional[str]:
|
|
558
|
+
return self.data.get("SMARTS")
|
|
559
|
+
|
|
560
|
+
@property
|
|
561
|
+
def size(self) -> int:
|
|
562
|
+
"""Number of carbon atoms in the component (falls back to atom-index count)."""
|
|
563
|
+
return self.data.get("size", len(self.atoms))
|
|
564
|
+
|
|
565
|
+
@property
|
|
566
|
+
def branching(self) -> Optional[float]:
|
|
567
|
+
"""Branching metric: 1.0 = linear, 0.0 = fully branched."""
|
|
568
|
+
return self.data.get("branching")
|
|
569
|
+
|
|
570
|
+
@property
|
|
571
|
+
def mean_eccentricity(self) -> Optional[float]:
|
|
572
|
+
"""Mean graph eccentricity across nodes in the component."""
|
|
573
|
+
return self.data.get("mean_eccentricity")
|
|
574
|
+
|
|
575
|
+
@property
|
|
576
|
+
def min_dist_to_centre(self) -> Optional[int]:
|
|
577
|
+
"""Minimum graph distance from any SMARTS match atom to the component centre."""
|
|
578
|
+
return self.data.get("min_dist_to_centre")
|
|
579
|
+
|
|
580
|
+
@property
|
|
581
|
+
def min_dist_to_center(self) -> Optional[int]: # US alias
|
|
582
|
+
return self.min_dist_to_centre
|
|
583
|
+
|
|
584
|
+
@property
|
|
585
|
+
def min_dist_to_barycentre(self) -> Optional[int]:
|
|
586
|
+
"""Minimum graph distance from any SMARTS match atom to the component barycentre."""
|
|
587
|
+
return self.data.get("min_dist_to_barycentre")
|
|
588
|
+
|
|
589
|
+
@property
|
|
590
|
+
def min_dist_to_barycenter(self) -> Optional[int]: # US alias
|
|
591
|
+
return self.min_dist_to_barycentre
|
|
592
|
+
|
|
593
|
+
@property
|
|
594
|
+
def max_dist_to_periphery(self) -> Optional[int]:
|
|
595
|
+
"""Maximum graph distance from any SMARTS match atom to the component periphery."""
|
|
596
|
+
return self.data.get("max_dist_to_periphery")
|
|
597
|
+
|
|
598
|
+
@property
|
|
599
|
+
def component_fraction(self) -> Optional[float]:
|
|
600
|
+
"""Fraction of total molecular carbon atoms that belong to this component.
|
|
601
|
+
|
|
602
|
+
Computed as (# C atoms in the augmented component) / (total # C atoms in molecule).
|
|
603
|
+
Oxygen, fluorine, and other heteroatoms in the augmented component are excluded
|
|
604
|
+
from both numerator and denominator, so the value is always in [0, 1].
|
|
605
|
+
"""
|
|
606
|
+
return self.data.get("component_fraction")
|
|
607
|
+
|
|
608
|
+
@property
|
|
609
|
+
def diameter(self) -> Optional[float]:
|
|
610
|
+
"""Graph diameter of the component (longest shortest path)."""
|
|
611
|
+
return self.data.get("diameter")
|
|
612
|
+
|
|
613
|
+
@property
|
|
614
|
+
def radius(self) -> Optional[float]:
|
|
615
|
+
"""Graph radius of the component (minimum eccentricity)."""
|
|
616
|
+
return self.data.get("radius")
|
|
617
|
+
|
|
618
|
+
@property
|
|
619
|
+
def effective_graph_resistance(self) -> Optional[float]:
|
|
620
|
+
"""Effective graph resistance (Kirchhoff index) of the component."""
|
|
621
|
+
return self.data.get("effective_graph_resistance")
|
|
622
|
+
|
|
623
|
+
@property
|
|
624
|
+
def effective_graph_resistance_BDE(self) -> Optional[float]:
|
|
625
|
+
"""BDE-weighted effective graph resistance of the component."""
|
|
626
|
+
return self.data.get("effective_graph_resistance_BDE")
|
|
627
|
+
|
|
628
|
+
|
|
629
|
+
class MatchView(dict):
|
|
630
|
+
"""Wrapper for a single match dict (PFAS group or definition).
|
|
631
|
+
|
|
632
|
+
Behaves as a normal dict but provides helpers for component access.
|
|
633
|
+
"""
|
|
634
|
+
|
|
635
|
+
@property
|
|
636
|
+
def is_group(self) -> bool:
|
|
637
|
+
return self.get("type") == "HalogenGroup"
|
|
638
|
+
|
|
639
|
+
@property
|
|
640
|
+
def is_definition(self) -> bool:
|
|
641
|
+
return self.get("type") == "PFASdefinition"
|
|
642
|
+
|
|
643
|
+
@property
|
|
644
|
+
def group_id(self) -> Optional[int]:
|
|
645
|
+
return self.get("id") if self.is_group else None
|
|
646
|
+
|
|
647
|
+
@property
|
|
648
|
+
def group_name(self) -> Optional[str]:
|
|
649
|
+
return self.get("group_name") if self.is_group else None
|
|
650
|
+
|
|
651
|
+
@property
|
|
652
|
+
def components(self) -> List[ComponentView]:
|
|
653
|
+
return [ComponentView(c) for c in self.get("components", [])]
|
|
654
|
+
|
|
655
|
+
|
|
656
|
+
class EmbeddingArray(np.ndarray):
|
|
657
|
+
"""A numpy array subclass that carries molecule identity metadata.
|
|
658
|
+
|
|
659
|
+
Returned by :meth:`PFASEmbedding.to_array` (1-D) and
|
|
660
|
+
:meth:`PFASEmbeddingSet.to_array` (2-D). All standard numpy operations
|
|
661
|
+
work unchanged; the extra attributes allow callers to trace each row back
|
|
662
|
+
to its source molecule.
|
|
663
|
+
|
|
664
|
+
Attributes
|
|
665
|
+
----------
|
|
666
|
+
smiles : str or list of str
|
|
667
|
+
SMILES string(s) for the molecule(s).
|
|
668
|
+
inchi : str or list of str
|
|
669
|
+
InChI string(s).
|
|
670
|
+
inchikey : str or list of str
|
|
671
|
+
InChIKey(s).
|
|
672
|
+
source : PFASEmbedding or PFASEmbeddingSet
|
|
673
|
+
Reference to the originating result object.
|
|
674
|
+
"""
|
|
675
|
+
|
|
676
|
+
def __new__(cls, array, smiles='', inchi='', inchikey='', source=None):
|
|
677
|
+
obj = np.asarray(array).view(cls)
|
|
678
|
+
obj._emb_smiles = smiles
|
|
679
|
+
obj._emb_inchi = inchi
|
|
680
|
+
obj._emb_inchikey = inchikey
|
|
681
|
+
obj._emb_source = source
|
|
682
|
+
return obj
|
|
683
|
+
|
|
684
|
+
def __array_finalize__(self, obj):
|
|
685
|
+
if obj is None:
|
|
686
|
+
return
|
|
687
|
+
self._emb_smiles = getattr(obj, '_emb_smiles', '')
|
|
688
|
+
self._emb_inchi = getattr(obj, '_emb_inchi', '')
|
|
689
|
+
self._emb_inchikey = getattr(obj, '_emb_inchikey', '')
|
|
690
|
+
self._emb_source = getattr(obj, '_emb_source', None)
|
|
691
|
+
|
|
692
|
+
@property
|
|
693
|
+
def smiles(self):
|
|
694
|
+
return self._emb_smiles
|
|
695
|
+
|
|
696
|
+
@property
|
|
697
|
+
def inchi(self):
|
|
698
|
+
return self._emb_inchi
|
|
699
|
+
|
|
700
|
+
@property
|
|
701
|
+
def inchikey(self):
|
|
702
|
+
return self._emb_inchikey
|
|
703
|
+
|
|
704
|
+
@property
|
|
705
|
+
def source(self):
|
|
706
|
+
return self._emb_source
|
|
707
|
+
|
|
708
|
+
|
|
709
|
+
class PFASEmbedding(dict):
|
|
710
|
+
"""Single-molecule PFAS result and embedding generator.
|
|
711
|
+
|
|
712
|
+
Subclasses :class:`dict` so all existing code accessing ``result['smiles']``,
|
|
713
|
+
``result['matches']``, etc. continues to work unchanged. Call
|
|
714
|
+
:meth:`to_array` to produce a numeric embedding vector from the stored
|
|
715
|
+
parsed data.
|
|
716
|
+
"""
|
|
717
|
+
|
|
718
|
+
@property
|
|
719
|
+
def smiles(self) -> str:
|
|
720
|
+
return self.get("smiles", "")
|
|
721
|
+
|
|
722
|
+
@property
|
|
723
|
+
def mol_with_h(self):
|
|
724
|
+
"""Get the molecule with explicit hydrogens used for component detection.
|
|
725
|
+
|
|
726
|
+
Atom indices in the molblock are guaranteed to match those stored in
|
|
727
|
+
matched component dicts. The older SMILES path is kept only for
|
|
728
|
+
backward compatibility with cached results (MolToSmiles reorders atoms).
|
|
729
|
+
"""
|
|
730
|
+
molblock = self.get("molblock_with_h")
|
|
731
|
+
if molblock:
|
|
732
|
+
return Chem.MolFromMolBlock(molblock, removeHs=False)
|
|
733
|
+
# Backward compatibility: SMILES path (atom ordering may differ)
|
|
734
|
+
smiles_h = self.get("smiles_with_h")
|
|
735
|
+
if smiles_h:
|
|
736
|
+
return Chem.MolFromSmiles(smiles_h)
|
|
737
|
+
return self.get("mol_with_h")
|
|
738
|
+
|
|
739
|
+
@property
|
|
740
|
+
def matches(self) -> List[MatchView]:
|
|
741
|
+
return [MatchView(m) for m in self.get("matches", [])]
|
|
742
|
+
|
|
743
|
+
def iter_group_matches(self, group_id: Optional[int] = None, group_name: Optional[str] = None) -> Iterator[MatchView]:
|
|
744
|
+
"""Iterate over PFAS group matches, optionally filtered by id/name."""
|
|
745
|
+
|
|
746
|
+
for m in self.matches:
|
|
747
|
+
if not m.is_group:
|
|
748
|
+
continue
|
|
749
|
+
if group_id is not None and m.group_id != group_id:
|
|
750
|
+
continue
|
|
751
|
+
if group_name is not None and m.group_name != group_name:
|
|
752
|
+
continue
|
|
753
|
+
yield m
|
|
754
|
+
|
|
755
|
+
def collect_component_atoms(self, group_id: Optional[int] = None, group_name: Optional[str] = None) -> List[int]:
|
|
756
|
+
"""Return all atom indices that belong to matching components."""
|
|
757
|
+
|
|
758
|
+
atoms: List[int] = []
|
|
759
|
+
for m in self.iter_group_matches(group_id=group_id, group_name=group_name):
|
|
760
|
+
for comp in m.components:
|
|
761
|
+
atoms.extend(comp.atoms)
|
|
762
|
+
# Deduplicate but keep order stable
|
|
763
|
+
seen = set()
|
|
764
|
+
deduped: List[int] = []
|
|
765
|
+
for idx in atoms:
|
|
766
|
+
if idx not in seen:
|
|
767
|
+
seen.add(idx)
|
|
768
|
+
deduped.append(idx)
|
|
769
|
+
return deduped
|
|
770
|
+
|
|
771
|
+
def summarise(self) -> str:
|
|
772
|
+
"""Return a coloured text summary of this molecule's results.
|
|
773
|
+
|
|
774
|
+
The summary includes:
|
|
775
|
+
- SMILES representation
|
|
776
|
+
- counts of PFAS group and definition matches
|
|
777
|
+
- total number of components across all group matches
|
|
778
|
+
- list of matched PFAS groups (colour-coded by halogen)
|
|
779
|
+
"""
|
|
780
|
+
meta = _get_component_meta()
|
|
781
|
+
|
|
782
|
+
total_group_matches = 0
|
|
783
|
+
total_definition_matches = 0
|
|
784
|
+
total_components = 0
|
|
785
|
+
group_counts: Dict[str, int] = {}
|
|
786
|
+
|
|
787
|
+
for m in self.matches:
|
|
788
|
+
if m.is_group:
|
|
789
|
+
total_group_matches += 1
|
|
790
|
+
total_components += len(m.components)
|
|
791
|
+
name = m.group_name or str(m.get("match_id", ""))
|
|
792
|
+
group_counts[name] = group_counts.get(name, 0) + 1
|
|
793
|
+
elif m.is_definition:
|
|
794
|
+
total_definition_matches += 1
|
|
795
|
+
|
|
796
|
+
lines: List[str] = []
|
|
797
|
+
lines.append(_ansi("PFASEmbedding summary", _ANSI_BOLD))
|
|
798
|
+
lines.append(f"- SMILES: {self.smiles}")
|
|
799
|
+
lines.append(f"- PFAS group matches: {total_group_matches}")
|
|
800
|
+
lines.append(f"- PFAS definition matches: {total_definition_matches}")
|
|
801
|
+
lines.append(f"- Total components: {total_components}")
|
|
802
|
+
|
|
803
|
+
if group_counts:
|
|
804
|
+
lines.append("- Matched PFAS groups:")
|
|
805
|
+
for name, count in sorted(group_counts.items(), key=lambda kv: kv[1], reverse=True):
|
|
806
|
+
# Determine halogen from first matching component of this group
|
|
807
|
+
halogen: Optional[str] = None
|
|
808
|
+
for m in self.matches:
|
|
809
|
+
if m.is_group and (m.group_name or "") == name:
|
|
810
|
+
for comp in m.components:
|
|
811
|
+
sl = comp.smarts_label
|
|
812
|
+
sl_str = str(sl) if isinstance(sl, list) else sl
|
|
813
|
+
info = meta.get(sl_str or "", {})
|
|
814
|
+
halogen = info.get("halogen")
|
|
815
|
+
break
|
|
816
|
+
break
|
|
817
|
+
hal_code = _ANSI_HALOGEN.get(halogen or "", "")
|
|
818
|
+
lines.append(f" * {hal_code}{_ANSI_BOLD}{name}{_ANSI_RESET}: {count} match(es)")
|
|
819
|
+
|
|
820
|
+
return "\n".join(lines)
|
|
821
|
+
|
|
822
|
+
def __str__(self) -> str:
|
|
823
|
+
return self.summarise()
|
|
824
|
+
|
|
825
|
+
def summary(self) -> None:
|
|
826
|
+
"""Print a detailed coloured summary of matched groups and components.
|
|
827
|
+
|
|
828
|
+
Each component is shown with graph metrics: size (carbon count),
|
|
829
|
+
branching (1.0 = linear, 0.0 = highly branched) and mean eccentricity.
|
|
830
|
+
"""
|
|
831
|
+
|
|
832
|
+
print("=" * 80)
|
|
833
|
+
print(f"{_ANSI_BOLD}MOLECULE:{_ANSI_RESET} {self.smiles}")
|
|
834
|
+
print("=" * 80)
|
|
835
|
+
|
|
836
|
+
if not self.matches:
|
|
837
|
+
print("No PFAS groups matched.")
|
|
838
|
+
return
|
|
839
|
+
|
|
840
|
+
# Collect group information
|
|
841
|
+
groups_info: Dict[Tuple[Optional[int], str], List[ComponentView]] = {}
|
|
842
|
+
|
|
843
|
+
for m in self.matches:
|
|
844
|
+
if not m.is_group:
|
|
845
|
+
continue
|
|
846
|
+
key = (m.group_id, m.group_name or "Unknown")
|
|
847
|
+
if key not in groups_info:
|
|
848
|
+
groups_info[key] = []
|
|
849
|
+
groups_info[key].extend(m.components)
|
|
850
|
+
|
|
851
|
+
if not groups_info:
|
|
852
|
+
print("No PFAS groups matched.")
|
|
853
|
+
return
|
|
854
|
+
|
|
855
|
+
print(f"\nMatched {len(groups_info)} PFAS group(s):")
|
|
856
|
+
print()
|
|
857
|
+
|
|
858
|
+
for (group_id, group_name), components in sorted(groups_info.items(), key=lambda x: x[0][0] or 0):
|
|
859
|
+
print(f"{_ANSI_BOLD}Group {group_id}: {group_name}{_ANSI_RESET}")
|
|
860
|
+
|
|
861
|
+
# Group components by SMARTS type
|
|
862
|
+
by_smarts: Dict[Optional[str], List[ComponentView]] = {}
|
|
863
|
+
for comp in components:
|
|
864
|
+
smarts = comp.smarts_label
|
|
865
|
+
smarts_key = str(smarts) if isinstance(smarts, list) else smarts
|
|
866
|
+
if smarts_key not in by_smarts:
|
|
867
|
+
by_smarts[smarts_key] = []
|
|
868
|
+
by_smarts[smarts_key].append(comp)
|
|
869
|
+
|
|
870
|
+
# Display components by SMARTS type
|
|
871
|
+
for smarts_label, comps in sorted(by_smarts.items(), key=lambda x: x[0] or ""):
|
|
872
|
+
label = smarts_label or "(no label)"
|
|
873
|
+
print(f" SMARTS: {label} ({len(comps)} component(s))")
|
|
874
|
+
for comp in sorted(comps, key=lambda c: c.size, reverse=True):
|
|
875
|
+
br_str = f"{comp.branching:.2f}" if comp.branching is not None else "\u2014"
|
|
876
|
+
ecc_str = f"{comp.mean_eccentricity:.2f}" if comp.mean_eccentricity is not None else "\u2014"
|
|
877
|
+
print(f" size={_ANSI_BOLD}{comp.size}{_ANSI_RESET} branching={br_str} eccentricity={ecc_str}")
|
|
878
|
+
|
|
879
|
+
print()
|
|
880
|
+
|
|
881
|
+
def table(self) -> str:
|
|
882
|
+
"""Return a text table with one row per match.
|
|
883
|
+
|
|
884
|
+
Columns:
|
|
885
|
+
- match_index: 1-based match index
|
|
886
|
+
- type: 'group' or 'definition'
|
|
887
|
+
- name: PFAS group name or definition name
|
|
888
|
+
- components: number of components (for groups)
|
|
889
|
+
"""
|
|
890
|
+
|
|
891
|
+
lines: List[str] = []
|
|
892
|
+
lines.append("match_index\ttype\tname\tcomponents")
|
|
893
|
+
|
|
894
|
+
for idx, m in enumerate(self.matches, start=1):
|
|
895
|
+
match_type = "group" if m.is_group else "definition"
|
|
896
|
+
if m.is_group:
|
|
897
|
+
name = m.group_name or str(m.get("match_id", ""))
|
|
898
|
+
components = len(m.components)
|
|
899
|
+
else:
|
|
900
|
+
name = m.get("definition_name") or m.get("name") or str(m.get("match_id", ""))
|
|
901
|
+
components = 0
|
|
902
|
+
|
|
903
|
+
lines.append(f"{idx}\t{match_type}\t{name}\t{components}")
|
|
904
|
+
|
|
905
|
+
return "\n".join(lines)
|
|
906
|
+
|
|
907
|
+
def classify(self) -> Tuple[str, int]:
|
|
908
|
+
"""Classify the molecule's PFAS content into a category label.
|
|
909
|
+
|
|
910
|
+
Returns
|
|
911
|
+
-------
|
|
912
|
+
(category, total_component_size) : Tuple[str, int]
|
|
913
|
+
*category* — a short label describing the main PFAS groups present:
|
|
914
|
+
|
|
915
|
+
* If one or more **OECD** groups are matched, returns their names
|
|
916
|
+
joined by ``", "``.
|
|
917
|
+
* Otherwise, returns a ``"per-"`` or ``"poly-"`` prefixed string
|
|
918
|
+
listing the matched **generic** / **telomeric** group names
|
|
919
|
+
(excluding groups 51 & 52), separated by ``", "``.
|
|
920
|
+
``"per-"`` is used only when **all** matched component SMARTS
|
|
921
|
+
carry ``saturation='per'``; ``"poly-"`` otherwise.
|
|
922
|
+
Name subsumption is applied: e.g. ``"amine"`` is suppressed
|
|
923
|
+
when ``"sulfonamide"`` or ``"amide"`` is present.
|
|
924
|
+
|
|
925
|
+
*total_component_size* — sum of :attr:`ComponentView.size` (C-atom
|
|
926
|
+
count) across **all** matched group components.
|
|
927
|
+
"""
|
|
928
|
+
group_info = _get_group_info()
|
|
929
|
+
meta = _get_component_meta()
|
|
930
|
+
|
|
931
|
+
oecd_names: List[str] = []
|
|
932
|
+
non_oecd: List[Tuple[int, str]] = [] # (group_id, name)
|
|
933
|
+
total_size: int = 0
|
|
934
|
+
seen_oecd: set = set()
|
|
935
|
+
seen_non_oecd_ids: set = set()
|
|
936
|
+
|
|
937
|
+
for m in self.iter_group_matches():
|
|
938
|
+
gid = m.group_id
|
|
939
|
+
if gid is None:
|
|
940
|
+
continue
|
|
941
|
+
info = group_info.get(gid, {})
|
|
942
|
+
cat = info.get("category", "other")
|
|
943
|
+
name = m.group_name or info.get("name", f"group_{gid}")
|
|
944
|
+
|
|
945
|
+
# Accumulate total component size across ALL groups
|
|
946
|
+
for comp in m.components:
|
|
947
|
+
total_size += comp.size
|
|
948
|
+
|
|
949
|
+
if cat == "OECD":
|
|
950
|
+
if name not in seen_oecd:
|
|
951
|
+
seen_oecd.add(name)
|
|
952
|
+
oecd_names.append(name)
|
|
953
|
+
elif cat in ("generic", "telomer"):
|
|
954
|
+
if gid not in _CLASSIFY_EXCLUDED_IDS and gid not in seen_non_oecd_ids:
|
|
955
|
+
seen_non_oecd_ids.add(gid)
|
|
956
|
+
non_oecd.append((gid, name))
|
|
957
|
+
|
|
958
|
+
# --- OECD priority -----------------------------------------------
|
|
959
|
+
if oecd_names:
|
|
960
|
+
return ", ".join(oecd_names), total_size
|
|
961
|
+
|
|
962
|
+
# --- Non-OECD (generic + telomeric) --------------------------------
|
|
963
|
+
if not non_oecd:
|
|
964
|
+
return "unclassified", total_size
|
|
965
|
+
|
|
966
|
+
# Apply name subsumption
|
|
967
|
+
matched_name_set = {name for _, name in non_oecd}
|
|
968
|
+
suppressed: set = set()
|
|
969
|
+
for name in matched_name_set:
|
|
970
|
+
for s in _SUBSUMES.get(name, []):
|
|
971
|
+
if s in matched_name_set:
|
|
972
|
+
suppressed.add(s)
|
|
973
|
+
|
|
974
|
+
filtered_names: List[str] = [
|
|
975
|
+
name for _, name in non_oecd if name not in suppressed
|
|
976
|
+
]
|
|
977
|
+
|
|
978
|
+
# Determine per/poly prefix from component saturation metadata
|
|
979
|
+
all_per = True
|
|
980
|
+
any_sat = False
|
|
981
|
+
for m in self.iter_group_matches():
|
|
982
|
+
gid = m.group_id
|
|
983
|
+
if gid is None or gid in _CLASSIFY_EXCLUDED_IDS:
|
|
984
|
+
continue
|
|
985
|
+
if group_info.get(gid, {}).get("category", "other") not in ("generic", "telomer"):
|
|
986
|
+
continue
|
|
987
|
+
for comp in m.components:
|
|
988
|
+
sl = comp.smarts_label
|
|
989
|
+
sl_str = str(sl) if isinstance(sl, list) else (sl or "")
|
|
990
|
+
sat = meta.get(sl_str, {}).get("saturation")
|
|
991
|
+
if sat is not None:
|
|
992
|
+
any_sat = True
|
|
993
|
+
if sat != "per":
|
|
994
|
+
all_per = False
|
|
995
|
+
|
|
996
|
+
prefix = "per" if (any_sat and all_per) else "poly"
|
|
997
|
+
label = (
|
|
998
|
+
f"{prefix}-{', '.join(filtered_names)}" if filtered_names
|
|
999
|
+
else "unclassified"
|
|
1000
|
+
)
|
|
1001
|
+
return label, total_size
|
|
1002
|
+
|
|
1003
|
+
def show(
|
|
1004
|
+
self,
|
|
1005
|
+
display: bool = True,
|
|
1006
|
+
subwidth: int = 350,
|
|
1007
|
+
subheight: int = 350,
|
|
1008
|
+
ncols: int = 4,
|
|
1009
|
+
) -> Image.Image:
|
|
1010
|
+
"""Show all component combinations for this molecule in a grid plot.
|
|
1011
|
+
|
|
1012
|
+
Components that share the same highlighted atoms are merged into a
|
|
1013
|
+
single panel. The legend lists every PFAS group (and its
|
|
1014
|
+
halogen / form / saturation metadata) that maps to those atoms as a
|
|
1015
|
+
bullet-point list, avoiding repeated panels for the same molecular
|
|
1016
|
+
fragment.
|
|
1017
|
+
|
|
1018
|
+
Atoms are highlighted with the colour of the first matching entry:
|
|
1019
|
+
- **Halogen**: cyan (F), green (Cl), amber (Br), violet (I)
|
|
1020
|
+
- **Form**: alkyl (full base colour) vs cyclic (warm-shifted tint)
|
|
1021
|
+
- **Saturation**: per- (full brightness) vs poly- (dimmed)
|
|
1022
|
+
|
|
1023
|
+
Parameters
|
|
1024
|
+
----------
|
|
1025
|
+
display : bool, default True
|
|
1026
|
+
Whether to display the image immediately.
|
|
1027
|
+
subwidth : int, default 350
|
|
1028
|
+
Width of each sub-image in pixels.
|
|
1029
|
+
subheight : int, default 350
|
|
1030
|
+
Minimum height of each sub-image in pixels. Panels with many
|
|
1031
|
+
matching groups are automatically made taller.
|
|
1032
|
+
ncols : int, default 4
|
|
1033
|
+
Number of columns in the grid.
|
|
1034
|
+
|
|
1035
|
+
Returns
|
|
1036
|
+
-------
|
|
1037
|
+
PIL.Image.Image
|
|
1038
|
+
Grid image containing all component visualizations.
|
|
1039
|
+
"""
|
|
1040
|
+
meta = _get_component_meta()
|
|
1041
|
+
|
|
1042
|
+
# Use the molecule with hydrogens that was used during component detection
|
|
1043
|
+
mol = self.mol_with_h
|
|
1044
|
+
if mol is None:
|
|
1045
|
+
# Fallback to reconstructing from SMILES if not available
|
|
1046
|
+
mol = Chem.MolFromSmiles(self.smiles)
|
|
1047
|
+
if mol is None:
|
|
1048
|
+
raise ValueError(f"Cannot parse SMILES: {self.smiles}")
|
|
1049
|
+
mol = Chem.AddHs(mol)
|
|
1050
|
+
|
|
1051
|
+
# Collect all (atoms_key -> entries) grouping
|
|
1052
|
+
from collections import OrderedDict
|
|
1053
|
+
comp_groups: Dict = OrderedDict()
|
|
1054
|
+
|
|
1055
|
+
for match in self.matches:
|
|
1056
|
+
if not match.is_group:
|
|
1057
|
+
continue
|
|
1058
|
+
base_label = match.group_name or match.get("match_id", "")
|
|
1059
|
+
for comp in match.components:
|
|
1060
|
+
atoms = comp.atoms
|
|
1061
|
+
if not atoms:
|
|
1062
|
+
continue
|
|
1063
|
+
key = frozenset(atoms)
|
|
1064
|
+
sl = comp.smarts_label
|
|
1065
|
+
sl_str = str(sl) if isinstance(sl, list) else (sl or "")
|
|
1066
|
+
info = meta.get(sl_str, {})
|
|
1067
|
+
halogen = info.get("halogen")
|
|
1068
|
+
form = info.get("form")
|
|
1069
|
+
saturation = info.get("saturation")
|
|
1070
|
+
colour = _component_color(halogen, form, saturation)
|
|
1071
|
+
if key not in comp_groups:
|
|
1072
|
+
# Build component-wide metrics once per unique atom set
|
|
1073
|
+
import math as _math
|
|
1074
|
+
size_str = str(comp.size)
|
|
1075
|
+
br_v = comp.branching
|
|
1076
|
+
ecc_v = comp.mean_eccentricity
|
|
1077
|
+
diam_v = comp.diameter
|
|
1078
|
+
rad_v = comp.radius
|
|
1079
|
+
egr_v = comp.effective_graph_resistance
|
|
1080
|
+
frc_v = comp.component_fraction
|
|
1081
|
+
br_str = f"{br_v:.2f}" if br_v is not None else "\u2014"
|
|
1082
|
+
ecc_str = f"{ecc_v:.2f}" if ecc_v is not None else "\u2014"
|
|
1083
|
+
diam_str = (f"{float(diam_v):.0f}" if diam_v is not None and not (isinstance(diam_v, float) and (_math.isnan(diam_v) or _math.isinf(diam_v))) else "\u2014")
|
|
1084
|
+
rad_str = (f"{float(rad_v):.0f}" if rad_v is not None and not (isinstance(rad_v, float) and (_math.isnan(rad_v) or _math.isinf(rad_v))) else "\u2014")
|
|
1085
|
+
egr_str = (f"{float(egr_v):.2f}" if egr_v is not None and not (isinstance(egr_v, float) and (_math.isnan(egr_v) or _math.isinf(egr_v))) else "\u2014")
|
|
1086
|
+
frc_str = f"{frc_v*100:.0f}%" if frc_v is not None else "\u2014"
|
|
1087
|
+
comp_groups[key] = {
|
|
1088
|
+
'atoms': sorted(atoms),
|
|
1089
|
+
'colour': colour,
|
|
1090
|
+
'halogen': halogen,
|
|
1091
|
+
'entries': [],
|
|
1092
|
+
'comp_metrics': (size_str, br_str, ecc_str, diam_str, rad_str, egr_str, frc_str),
|
|
1093
|
+
}
|
|
1094
|
+
# FG-specific entry: group name, SMARTS type, dist-to-centre, dist-to-periphery
|
|
1095
|
+
dct_v = comp.min_dist_to_centre
|
|
1096
|
+
dpr_v = comp.max_dist_to_periphery
|
|
1097
|
+
dct_str = str(dct_v) if dct_v is not None else "\u2014"
|
|
1098
|
+
dpr_str = str(dpr_v) if dpr_v is not None else "\u2014"
|
|
1099
|
+
entry = (base_label, sl_str, dct_str, dpr_str)
|
|
1100
|
+
if entry not in comp_groups[key]['entries']:
|
|
1101
|
+
comp_groups[key]['entries'].append(entry)
|
|
1102
|
+
|
|
1103
|
+
imgs: List[Image.Image] = []
|
|
1104
|
+
|
|
1105
|
+
for data in comp_groups.values():
|
|
1106
|
+
atoms = data['atoms']
|
|
1107
|
+
colour = data['colour']
|
|
1108
|
+
entries = data['entries']
|
|
1109
|
+
comp_metrics = data.get('comp_metrics')
|
|
1110
|
+
halogen_lbl = data.get('halogen') or ""
|
|
1111
|
+
|
|
1112
|
+
atom_colours: Dict[int, Color] = {a: colour for a in atoms}
|
|
1113
|
+
d2d = Draw.MolDraw2DCairo(subwidth, subheight)
|
|
1114
|
+
dopts = d2d.drawOptions()
|
|
1115
|
+
dopts.useBWAtomPalette()
|
|
1116
|
+
dopts.fixedBondLength = 20
|
|
1117
|
+
dopts.addAtomIndices = True
|
|
1118
|
+
dopts.addBondIndices = False
|
|
1119
|
+
dopts.maxFontSize = 14
|
|
1120
|
+
dopts.minFontSize = 12
|
|
1121
|
+
d2d.DrawMolecule(
|
|
1122
|
+
mol,
|
|
1123
|
+
highlightAtoms=atoms,
|
|
1124
|
+
highlightAtomColors=atom_colours,
|
|
1125
|
+
)
|
|
1126
|
+
d2d.FinishDrawing()
|
|
1127
|
+
mol_img = Image.open(BytesIO(d2d.GetDrawingText()))
|
|
1128
|
+
imgs.append(_mol_image_with_table(mol_img, entries, comp_metrics=comp_metrics, halogen_label=halogen_lbl))
|
|
1129
|
+
|
|
1130
|
+
if not imgs:
|
|
1131
|
+
raise ValueError("No PFAS group components found to display.")
|
|
1132
|
+
|
|
1133
|
+
grid, _, _ = _grid_images(imgs, buffer=4, ncols=ncols)
|
|
1134
|
+
if display:
|
|
1135
|
+
grid.show()
|
|
1136
|
+
return grid
|
|
1137
|
+
|
|
1138
|
+
# Alias so callers can use either mol_result.show() or mol_result.plot()
|
|
1139
|
+
plot = show
|
|
1140
|
+
|
|
1141
|
+
def svg(
|
|
1142
|
+
self,
|
|
1143
|
+
filename: str,
|
|
1144
|
+
subwidth: int = 350,
|
|
1145
|
+
subheight: int = 350,
|
|
1146
|
+
ncols: int = 4,
|
|
1147
|
+
) -> str:
|
|
1148
|
+
"""Export all component combinations to an SVG file (vector graphics).
|
|
1149
|
+
|
|
1150
|
+
Components that share the same highlighted atoms are merged into a
|
|
1151
|
+
single panel with a bullet-point legend listing all matching groups.
|
|
1152
|
+
|
|
1153
|
+
Parameters
|
|
1154
|
+
----------
|
|
1155
|
+
filename : str
|
|
1156
|
+
Path to the output SVG file.
|
|
1157
|
+
subwidth : int, default 350
|
|
1158
|
+
Width of each sub-image in pixels.
|
|
1159
|
+
subheight : int, default 350
|
|
1160
|
+
Minimum height of each sub-image in pixels.
|
|
1161
|
+
ncols : int, default 4
|
|
1162
|
+
Number of columns in the grid.
|
|
1163
|
+
|
|
1164
|
+
Returns
|
|
1165
|
+
-------
|
|
1166
|
+
str
|
|
1167
|
+
Path to the created SVG file.
|
|
1168
|
+
"""
|
|
1169
|
+
import svgutils.transform as sg
|
|
1170
|
+
|
|
1171
|
+
# Use the molecule with hydrogens that was used during component detection
|
|
1172
|
+
mol = self.mol_with_h
|
|
1173
|
+
if mol is None:
|
|
1174
|
+
# Fallback to reconstructing from SMILES if not available
|
|
1175
|
+
mol = Chem.MolFromSmiles(self.smiles)
|
|
1176
|
+
if mol is None:
|
|
1177
|
+
raise ValueError(f"Cannot parse SMILES: {self.smiles}")
|
|
1178
|
+
mol = Chem.AddHs(mol)
|
|
1179
|
+
|
|
1180
|
+
# Group by unique atom set
|
|
1181
|
+
from collections import OrderedDict
|
|
1182
|
+
comp_groups: Dict = OrderedDict()
|
|
1183
|
+
|
|
1184
|
+
for match in self.matches:
|
|
1185
|
+
if not match.is_group:
|
|
1186
|
+
continue
|
|
1187
|
+
base_label = match.group_name or match.get("match_id", "")
|
|
1188
|
+
for comp in match.components:
|
|
1189
|
+
atoms = comp.atoms
|
|
1190
|
+
if not atoms:
|
|
1191
|
+
continue
|
|
1192
|
+
key = frozenset(atoms)
|
|
1193
|
+
sl = comp.smarts_label
|
|
1194
|
+
sl_str = str(sl) if isinstance(sl, list) else (sl or "")
|
|
1195
|
+
if key not in comp_groups:
|
|
1196
|
+
import math as _math
|
|
1197
|
+
br_v = comp.branching
|
|
1198
|
+
ecc_v = comp.mean_eccentricity
|
|
1199
|
+
diam_v = comp.diameter
|
|
1200
|
+
rad_v = comp.radius
|
|
1201
|
+
egr_v = comp.effective_graph_resistance
|
|
1202
|
+
frc_v = comp.component_fraction
|
|
1203
|
+
br_str = f"{br_v:.2f}" if br_v is not None else "\u2014"
|
|
1204
|
+
ecc_str = f"{ecc_v:.2f}" if ecc_v is not None else "\u2014"
|
|
1205
|
+
diam_str = (f"{float(diam_v):.0f}" if diam_v is not None and not (isinstance(diam_v, float) and (_math.isnan(diam_v) or _math.isinf(diam_v))) else "\u2014")
|
|
1206
|
+
rad_str = (f"{float(rad_v):.0f}" if rad_v is not None and not (isinstance(rad_v, float) and (_math.isnan(rad_v) or _math.isinf(rad_v))) else "\u2014")
|
|
1207
|
+
egr_str = (f"{float(egr_v):.2f}" if egr_v is not None and not (isinstance(egr_v, float) and (_math.isnan(egr_v) or _math.isinf(egr_v))) else "\u2014")
|
|
1208
|
+
frc_str = f"{frc_v*100:.0f}%" if frc_v is not None else "\u2014"
|
|
1209
|
+
comp_groups[key] = {
|
|
1210
|
+
'atoms': sorted(atoms),
|
|
1211
|
+
'entries': [],
|
|
1212
|
+
'comp_metrics': (str(comp.size), br_str, ecc_str, diam_str, rad_str, egr_str, frc_str),
|
|
1213
|
+
}
|
|
1214
|
+
# FG-specific entry
|
|
1215
|
+
dct_v = comp.min_dist_to_centre
|
|
1216
|
+
dpr_v = comp.max_dist_to_periphery
|
|
1217
|
+
dct_str = str(dct_v) if dct_v is not None else "\u2014"
|
|
1218
|
+
dpr_str = str(dpr_v) if dpr_v is not None else "\u2014"
|
|
1219
|
+
entry = (base_label, sl_str, dct_str, dpr_str)
|
|
1220
|
+
if entry not in comp_groups[key]['entries']:
|
|
1221
|
+
comp_groups[key]['entries'].append(entry)
|
|
1222
|
+
|
|
1223
|
+
imgs: List[str] = []
|
|
1224
|
+
|
|
1225
|
+
for data in comp_groups.values():
|
|
1226
|
+
atoms = data['atoms']
|
|
1227
|
+
entries = data['entries']
|
|
1228
|
+
comp_metrics = data.get('comp_metrics')
|
|
1229
|
+
n = len(entries)
|
|
1230
|
+
|
|
1231
|
+
lines: List[str] = []
|
|
1232
|
+
for grp, sls, dct, dpr in entries:
|
|
1233
|
+
line = f"\u2022 {grp}"
|
|
1234
|
+
if sls:
|
|
1235
|
+
line += f" | {sls}"
|
|
1236
|
+
line += f" dct={dct} dpr={dpr}"
|
|
1237
|
+
lines.append(line)
|
|
1238
|
+
if comp_metrics:
|
|
1239
|
+
sz, br, ecc, diam, rad, egr, frc = comp_metrics
|
|
1240
|
+
lines.append(f" size={sz} br={br} ecc={ecc} diam={diam} rad={rad} egr={egr} chain={frc}")
|
|
1241
|
+
legend = "\n".join(lines)
|
|
1242
|
+
|
|
1243
|
+
effective_height = max(subheight, 220 + n * 38)
|
|
1244
|
+
|
|
1245
|
+
d2d = Draw.MolDraw2DSVG(subwidth, effective_height)
|
|
1246
|
+
dopts = d2d.drawOptions()
|
|
1247
|
+
dopts.useBWAtomPalette()
|
|
1248
|
+
dopts.fixedBondLength = 20
|
|
1249
|
+
dopts.addAtomIndices = True
|
|
1250
|
+
dopts.addBondIndices = False
|
|
1251
|
+
dopts.maxFontSize = 14
|
|
1252
|
+
dopts.minFontSize = 12
|
|
1253
|
+
d2d.DrawMolecule(mol, legend=legend, highlightAtoms=atoms)
|
|
1254
|
+
d2d.FinishDrawing()
|
|
1255
|
+
imgs.append(d2d.GetDrawingText())
|
|
1256
|
+
|
|
1257
|
+
if not imgs:
|
|
1258
|
+
raise ValueError("No PFAS group components found to display.")
|
|
1259
|
+
|
|
1260
|
+
# Convert SVG strings to svgutils figures
|
|
1261
|
+
svg_figs = [sg.fromstring(img) for img in imgs]
|
|
1262
|
+
|
|
1263
|
+
# Merge into grid
|
|
1264
|
+
from .draw_mols import merge_svg
|
|
1265
|
+
grid, _, _ = merge_svg(svg_figs, buffer=4, ncols=ncols)
|
|
1266
|
+
|
|
1267
|
+
grid.save(filename)
|
|
1268
|
+
return filename
|
|
1269
|
+
|
|
1270
|
+
# ------------------------------------------------------------------
|
|
1271
|
+
# Factory constructors
|
|
1272
|
+
# ------------------------------------------------------------------
|
|
1273
|
+
|
|
1274
|
+
@classmethod
|
|
1275
|
+
def from_smiles(cls, smiles: str, **kwargs) -> "PFASEmbedding":
|
|
1276
|
+
"""Parse a SMILES string and return a single :class:`PFASEmbedding`.
|
|
1277
|
+
|
|
1278
|
+
Parameters
|
|
1279
|
+
----------
|
|
1280
|
+
smiles : str
|
|
1281
|
+
SMILES string for one molecule.
|
|
1282
|
+
**kwargs
|
|
1283
|
+
Forwarded to :func:`~PFASGroups.parser.parse_smiles`
|
|
1284
|
+
(e.g. ``halogens``, ``saturation``, ``progress``).
|
|
1285
|
+
"""
|
|
1286
|
+
from .parser import parse_smiles
|
|
1287
|
+
return parse_smiles(smiles, **kwargs)[0]
|
|
1288
|
+
|
|
1289
|
+
@classmethod
|
|
1290
|
+
def from_mol(cls, mol, **kwargs) -> "PFASEmbedding":
|
|
1291
|
+
"""Parse an RDKit molecule and return a single :class:`PFASEmbedding`.
|
|
1292
|
+
|
|
1293
|
+
Parameters
|
|
1294
|
+
----------
|
|
1295
|
+
mol : rdkit.Chem.Mol
|
|
1296
|
+
RDKit molecule object.
|
|
1297
|
+
**kwargs
|
|
1298
|
+
Forwarded to :func:`~PFASGroups.parser.parse_mols`.
|
|
1299
|
+
"""
|
|
1300
|
+
from .parser import parse_mols
|
|
1301
|
+
return parse_mols([mol], **kwargs)[0]
|
|
1302
|
+
|
|
1303
|
+
@classmethod
|
|
1304
|
+
def from_inchi(cls, inchi: str, **kwargs) -> "PFASEmbedding":
|
|
1305
|
+
"""Parse an InChI string and return a single :class:`PFASEmbedding`.
|
|
1306
|
+
|
|
1307
|
+
Parameters
|
|
1308
|
+
----------
|
|
1309
|
+
inchi : str
|
|
1310
|
+
InChI string for one molecule.
|
|
1311
|
+
**kwargs
|
|
1312
|
+
Forwarded to :func:`~PFASGroups.parser.parse_mols`.
|
|
1313
|
+
"""
|
|
1314
|
+
from rdkit.Chem.inchi import MolFromInchi
|
|
1315
|
+
mol = MolFromInchi(inchi)
|
|
1316
|
+
if mol is None:
|
|
1317
|
+
raise ValueError(f"Cannot parse InChI: {inchi!r}")
|
|
1318
|
+
return cls.from_mol(mol, **kwargs)
|
|
1319
|
+
|
|
1320
|
+
def to_fingerprint(
|
|
1321
|
+
self,
|
|
1322
|
+
group_selection: str = 'all',
|
|
1323
|
+
component_metrics: Optional[List[str]] = None,
|
|
1324
|
+
selected_group_ids: Optional[List[int]] = None,
|
|
1325
|
+
halogens: Union[str, List[str]] = 'F',
|
|
1326
|
+
saturation: Optional[str] = 'per',
|
|
1327
|
+
molecule_metrics: Optional[List[str]] = None,
|
|
1328
|
+
pfas_groups: Optional[List[Dict]] = None,
|
|
1329
|
+
preset: Optional[str] = None,
|
|
1330
|
+
count_mode: Optional[str] = None,
|
|
1331
|
+
graph_metrics: Optional[List[str]] = None,
|
|
1332
|
+
progress: bool = False,
|
|
1333
|
+
**kwargs,
|
|
1334
|
+
) -> np.ndarray:
|
|
1335
|
+
"""Deprecated. Use :meth:`to_array` instead."""
|
|
1336
|
+
import warnings
|
|
1337
|
+
warnings.warn(
|
|
1338
|
+
"to_fingerprint() is deprecated; use to_array() instead.",
|
|
1339
|
+
DeprecationWarning,
|
|
1340
|
+
stacklevel=2,
|
|
1341
|
+
)
|
|
1342
|
+
if count_mode is not None or graph_metrics is not None:
|
|
1343
|
+
if component_metrics is None:
|
|
1344
|
+
component_metrics = [count_mode or 'binary'] + list(graph_metrics or [])
|
|
1345
|
+
return self.to_array(
|
|
1346
|
+
component_metrics=component_metrics,
|
|
1347
|
+
molecule_metrics=molecule_metrics,
|
|
1348
|
+
group_selection=group_selection,
|
|
1349
|
+
selected_group_ids=selected_group_ids,
|
|
1350
|
+
preset=preset,
|
|
1351
|
+
pfas_groups=pfas_groups,
|
|
1352
|
+
)
|
|
1353
|
+
|
|
1354
|
+
def to_array(
|
|
1355
|
+
self,
|
|
1356
|
+
component_metrics=_UNSET,
|
|
1357
|
+
molecule_metrics=_UNSET,
|
|
1358
|
+
group_selection=_UNSET,
|
|
1359
|
+
selected_group_ids=_UNSET,
|
|
1360
|
+
aggregation=_UNSET,
|
|
1361
|
+
preset=_UNSET,
|
|
1362
|
+
pfas_groups=_UNSET,
|
|
1363
|
+
halogens=_UNSET,
|
|
1364
|
+
) -> np.ndarray:
|
|
1365
|
+
"""Generate a 1-D embedding vector for this molecule.
|
|
1366
|
+
|
|
1367
|
+
When called with no arguments, returns the last cached embedding (or
|
|
1368
|
+
binary by default on the first call). Pass explicit arguments to
|
|
1369
|
+
override and update the cache.
|
|
1370
|
+
|
|
1371
|
+
Parameters
|
|
1372
|
+
----------
|
|
1373
|
+
component_metrics : list of str, default ['binary']
|
|
1374
|
+
Per-component metrics. Count modes: 'binary', 'count',
|
|
1375
|
+
'max_component', 'total_component'. Graph metrics:
|
|
1376
|
+
'effective_graph_resistance', 'effective_graph_resistance_BDE',
|
|
1377
|
+
'branching', 'mean_eccentricity', etc.
|
|
1378
|
+
molecule_metrics : list of str, optional
|
|
1379
|
+
Molecule-wide scalars appended after all component columns:
|
|
1380
|
+
'n_components', 'total_size', 'mean_branching', etc.
|
|
1381
|
+
group_selection : str, default 'all'
|
|
1382
|
+
'all', 'oecd', 'generic', 'telomers', or 'generic+telomers'.
|
|
1383
|
+
selected_group_ids : list of int, optional
|
|
1384
|
+
Explicit group IDs (overrides group_selection).
|
|
1385
|
+
aggregation : str, default 'mean'
|
|
1386
|
+
How to aggregate multiple matched components per group:
|
|
1387
|
+
'mean' or 'median'.
|
|
1388
|
+
preset : str, optional
|
|
1389
|
+
Named configuration from ``EMBEDDING_PRESETS``.
|
|
1390
|
+
pfas_groups : list, optional
|
|
1391
|
+
Custom group list (loaded from defaults when None).
|
|
1392
|
+
halogens : str or list of str, optional
|
|
1393
|
+
When provided, produce one block of ``n_groups`` columns *per
|
|
1394
|
+
halogen*, filtering each block's components to that halogen only.
|
|
1395
|
+
E.g. ``halogens=['F', 'Cl']`` yields a vector of length
|
|
1396
|
+
``2 × n_groups × len(component_metrics)``.
|
|
1397
|
+
``None`` (default) preserves the original behaviour: all
|
|
1398
|
+
components are used regardless of halogen.
|
|
1399
|
+
|
|
1400
|
+
Returns
|
|
1401
|
+
-------
|
|
1402
|
+
np.ndarray
|
|
1403
|
+
1-D float array of length
|
|
1404
|
+
``n_halogens × n_groups × len(component_metrics)
|
|
1405
|
+
+ len(molecule_metrics)``.
|
|
1406
|
+
"""
|
|
1407
|
+
# Return last cached result when called with no arguments
|
|
1408
|
+
_no_args = (
|
|
1409
|
+
component_metrics is _UNSET and molecule_metrics is _UNSET and
|
|
1410
|
+
group_selection is _UNSET and selected_group_ids is _UNSET and
|
|
1411
|
+
aggregation is _UNSET and preset is _UNSET and pfas_groups is _UNSET
|
|
1412
|
+
and halogens is _UNSET
|
|
1413
|
+
)
|
|
1414
|
+
if _no_args and getattr(self, '_last_array', None) is not None:
|
|
1415
|
+
return self._last_array
|
|
1416
|
+
|
|
1417
|
+
# Resolve sentinels to actual defaults
|
|
1418
|
+
if component_metrics is _UNSET: component_metrics = None
|
|
1419
|
+
if molecule_metrics is _UNSET: molecule_metrics = None
|
|
1420
|
+
if group_selection is _UNSET: group_selection = 'all'
|
|
1421
|
+
if selected_group_ids is _UNSET: selected_group_ids = None
|
|
1422
|
+
if aggregation is _UNSET: aggregation = 'mean'
|
|
1423
|
+
if preset is _UNSET: preset = None
|
|
1424
|
+
if pfas_groups is _UNSET: pfas_groups = None
|
|
1425
|
+
if halogens is _UNSET: halogens = None
|
|
1426
|
+
|
|
1427
|
+
from .embeddings import FINGERPRINT_PRESETS, _COUNT_MODES
|
|
1428
|
+
from .getter import get_compiled_HalogenGroups
|
|
1429
|
+
|
|
1430
|
+
resolved_cm: List[str] = list(component_metrics) if component_metrics else ['binary']
|
|
1431
|
+
resolved_mm: List[str] = list(molecule_metrics) if molecule_metrics else []
|
|
1432
|
+
|
|
1433
|
+
if preset is not None:
|
|
1434
|
+
if preset not in FINGERPRINT_PRESETS:
|
|
1435
|
+
raise ValueError(f"Unknown preset: {preset!r}. Available: {sorted(FINGERPRINT_PRESETS)}")
|
|
1436
|
+
_p = FINGERPRINT_PRESETS[preset]
|
|
1437
|
+
if _p.get('component_metrics') is not None:
|
|
1438
|
+
resolved_cm = list(_p['component_metrics'])
|
|
1439
|
+
if _p.get('molecule_metrics') is not None:
|
|
1440
|
+
resolved_mm = list(_p['molecule_metrics'])
|
|
1441
|
+
|
|
1442
|
+
if pfas_groups is None:
|
|
1443
|
+
pfas_groups = get_compiled_HalogenGroups()
|
|
1444
|
+
|
|
1445
|
+
sel_groups = _select_groups(group_selection, selected_group_ids, pfas_groups)
|
|
1446
|
+
match_by_id = {m['id']: m for m in self.get('matches', [])}
|
|
1447
|
+
|
|
1448
|
+
# Normalise halogens to a list (or None for legacy single-block mode)
|
|
1449
|
+
resolved_hal: Optional[List[str]] = (
|
|
1450
|
+
[halogens] if isinstance(halogens, str)
|
|
1451
|
+
else list(halogens) if halogens is not None
|
|
1452
|
+
else None
|
|
1453
|
+
)
|
|
1454
|
+
_meta = _get_component_meta() if resolved_hal is not None else {}
|
|
1455
|
+
hal_loop = resolved_hal if resolved_hal is not None else [None]
|
|
1456
|
+
|
|
1457
|
+
row: List[float] = []
|
|
1458
|
+
for hal in hal_loop:
|
|
1459
|
+
if hal is not None:
|
|
1460
|
+
hal_groups = [g for g in sel_groups
|
|
1461
|
+
if g.excludeHalogens is None or hal not in g.excludeHalogens]
|
|
1462
|
+
else:
|
|
1463
|
+
hal_groups = sel_groups
|
|
1464
|
+
for m in resolved_cm:
|
|
1465
|
+
for g in hal_groups:
|
|
1466
|
+
match = match_by_id.get(g.id)
|
|
1467
|
+
if match is None:
|
|
1468
|
+
row.append(0.0)
|
|
1469
|
+
else:
|
|
1470
|
+
comps = match.get('components', [])
|
|
1471
|
+
if hal is not None:
|
|
1472
|
+
comps = [c for c in comps
|
|
1473
|
+
if _meta.get(c.get('SMARTS', ''), {}).get('halogen') == hal]
|
|
1474
|
+
if m in _COUNT_MODES:
|
|
1475
|
+
row.append(_encode_count_comps(comps, m))
|
|
1476
|
+
else:
|
|
1477
|
+
row.append(_agg(comps, m, aggregation))
|
|
1478
|
+
|
|
1479
|
+
all_comps = [c for m in match_by_id.values() for c in m.get('components', [])]
|
|
1480
|
+
for m in resolved_mm:
|
|
1481
|
+
row.append(_mol_metric(all_comps, m))
|
|
1482
|
+
|
|
1483
|
+
result = EmbeddingArray(
|
|
1484
|
+
np.array(row, dtype=float),
|
|
1485
|
+
smiles=self.get('smiles', ''),
|
|
1486
|
+
inchi=self.get('inchi', ''),
|
|
1487
|
+
inchikey=self.get('inchikey', ''),
|
|
1488
|
+
source=self,
|
|
1489
|
+
)
|
|
1490
|
+
if _no_args:
|
|
1491
|
+
self._last_array = result
|
|
1492
|
+
return result
|
|
1493
|
+
|
|
1494
|
+
def column_names(
|
|
1495
|
+
self,
|
|
1496
|
+
component_metrics: Optional[List[str]] = None,
|
|
1497
|
+
molecule_metrics: Optional[List[str]] = None,
|
|
1498
|
+
group_selection: str = 'all',
|
|
1499
|
+
selected_group_ids: Optional[List[int]] = None,
|
|
1500
|
+
preset: Optional[str] = None,
|
|
1501
|
+
pfas_groups=None,
|
|
1502
|
+
halogens=None,
|
|
1503
|
+
) -> List[str]:
|
|
1504
|
+
"""Return the list of column labels for :meth:`to_array` without computing values.
|
|
1505
|
+
|
|
1506
|
+
Parameters match those of :meth:`to_array` (``aggregation`` is not
|
|
1507
|
+
relevant for column names).
|
|
1508
|
+
"""
|
|
1509
|
+
from .embeddings import FINGERPRINT_PRESETS
|
|
1510
|
+
from .getter import get_compiled_HalogenGroups
|
|
1511
|
+
|
|
1512
|
+
resolved_cm: List[str] = list(component_metrics) if component_metrics else ['binary']
|
|
1513
|
+
resolved_mm: List[str] = list(molecule_metrics) if molecule_metrics else []
|
|
1514
|
+
|
|
1515
|
+
if preset is not None:
|
|
1516
|
+
_p = FINGERPRINT_PRESETS.get(preset, {})
|
|
1517
|
+
if _p.get('component_metrics') is not None:
|
|
1518
|
+
resolved_cm = list(_p['component_metrics'])
|
|
1519
|
+
if _p.get('molecule_metrics') is not None:
|
|
1520
|
+
resolved_mm = list(_p['molecule_metrics'])
|
|
1521
|
+
|
|
1522
|
+
if pfas_groups is None:
|
|
1523
|
+
pfas_groups = get_compiled_HalogenGroups()
|
|
1524
|
+
|
|
1525
|
+
sel_groups = _select_groups(group_selection, selected_group_ids, pfas_groups)
|
|
1526
|
+
|
|
1527
|
+
resolved_hal: Optional[List[str]] = (
|
|
1528
|
+
[halogens] if isinstance(halogens, str)
|
|
1529
|
+
else list(halogens) if halogens is not None
|
|
1530
|
+
else None
|
|
1531
|
+
)
|
|
1532
|
+
hal_loop = resolved_hal if resolved_hal is not None else [None]
|
|
1533
|
+
|
|
1534
|
+
names = []
|
|
1535
|
+
for hal in hal_loop:
|
|
1536
|
+
if hal is not None:
|
|
1537
|
+
hal_groups = [g for g in sel_groups
|
|
1538
|
+
if g.excludeHalogens is None or hal not in g.excludeHalogens]
|
|
1539
|
+
else:
|
|
1540
|
+
hal_groups = sel_groups
|
|
1541
|
+
for m in resolved_cm:
|
|
1542
|
+
for g in hal_groups:
|
|
1543
|
+
names.append(f"{g.name} [{m}]{f' ({hal})' if hal is not None else ''}")
|
|
1544
|
+
names += [f"mol:{m}" for m in resolved_mm]
|
|
1545
|
+
return names
|
|
1546
|
+
|
|
1547
|
+
def to_sql(
|
|
1548
|
+
self,
|
|
1549
|
+
filename: Optional[str] = None,
|
|
1550
|
+
dbname: Optional[str] = None,
|
|
1551
|
+
user: Optional[str] = None,
|
|
1552
|
+
password: Optional[str] = None,
|
|
1553
|
+
host: Optional[str] = None,
|
|
1554
|
+
port: Optional[int] = None,
|
|
1555
|
+
components_table: str = "components",
|
|
1556
|
+
groups_table: str = "pfas_groups_in_compound",
|
|
1557
|
+
if_exists: str = "append",
|
|
1558
|
+
) -> None:
|
|
1559
|
+
"""Export this molecule result to a SQL database.
|
|
1560
|
+
|
|
1561
|
+
Can write to either SQLite (via filename) or PostgreSQL/MySQL (via connection parameters).
|
|
1562
|
+
|
|
1563
|
+
Parameters
|
|
1564
|
+
----------
|
|
1565
|
+
filename : str, optional
|
|
1566
|
+
Path to SQLite database file. If provided, uses SQLite.
|
|
1567
|
+
dbname : str, optional
|
|
1568
|
+
Database name (for PostgreSQL/MySQL).
|
|
1569
|
+
user : str, optional
|
|
1570
|
+
Database username. Defaults to os.environ['DB_USER'] if not provided.
|
|
1571
|
+
password : str, optional
|
|
1572
|
+
Database password. Defaults to os.environ['DB_PASSWORD'] if not provided.
|
|
1573
|
+
host : str, optional
|
|
1574
|
+
Database host. Defaults to os.environ.get('DB_HOST', 'localhost').
|
|
1575
|
+
port : int, optional
|
|
1576
|
+
Database port. Defaults to os.environ.get('DB_PORT', 5432 for PostgreSQL).
|
|
1577
|
+
components_table : str, default "components"
|
|
1578
|
+
Name of the table to store component-level data.
|
|
1579
|
+
groups_table : str, default "pfas_groups_in_compound"
|
|
1580
|
+
Name of the table to store PFAS group matches.
|
|
1581
|
+
if_exists : str, default "append"
|
|
1582
|
+
How to behave if tables exist: 'fail', 'replace', or 'append'.
|
|
1583
|
+
"""
|
|
1584
|
+
try:
|
|
1585
|
+
import pandas as pd
|
|
1586
|
+
import sqlalchemy
|
|
1587
|
+
except ImportError as exc:
|
|
1588
|
+
raise ImportError("pandas and sqlalchemy are required for to_sql. Install with: pip install pandas sqlalchemy") from exc
|
|
1589
|
+
# Determine connection
|
|
1590
|
+
if filename:
|
|
1591
|
+
engine = sqlalchemy.create_engine(f"sqlite:///{filename}")
|
|
1592
|
+
elif dbname:
|
|
1593
|
+
# Get credentials from environment if not provided
|
|
1594
|
+
if user is None:
|
|
1595
|
+
user = os.environ.get('DB_USER')
|
|
1596
|
+
if password is None:
|
|
1597
|
+
password = os.environ.get('DB_PASSWORD')
|
|
1598
|
+
if host is None:
|
|
1599
|
+
host = os.environ.get('DB_HOST', 'localhost')
|
|
1600
|
+
if port is None:
|
|
1601
|
+
port = int(os.environ.get('DB_PORT', 5432))
|
|
1602
|
+
|
|
1603
|
+
if not user or not password:
|
|
1604
|
+
raise ValueError("Database credentials required. Provide user/password or set DB_USER/DB_PASSWORD environment variables.")
|
|
1605
|
+
|
|
1606
|
+
# Assuming PostgreSQL; adjust for MySQL if needed
|
|
1607
|
+
connection_string = f"postgresql://{user}:{password}@{host}:{port}/{dbname}"
|
|
1608
|
+
engine = sqlalchemy.create_engine(connection_string)
|
|
1609
|
+
else:
|
|
1610
|
+
raise ValueError("Either filename (for SQLite) or dbname (for PostgreSQL) must be provided.")
|
|
1611
|
+
|
|
1612
|
+
# Prepare components data
|
|
1613
|
+
components_data = []
|
|
1614
|
+
for match in self.matches:
|
|
1615
|
+
if not match.is_group:
|
|
1616
|
+
continue
|
|
1617
|
+
for comp in match.components:
|
|
1618
|
+
smarts = comp.smarts_label
|
|
1619
|
+
components_data.append({
|
|
1620
|
+
'smiles': self.smiles,
|
|
1621
|
+
'group_id': match.group_id,
|
|
1622
|
+
'group_name': match.group_name,
|
|
1623
|
+
'smarts_label': str(smarts) if isinstance(smarts, list) else smarts,
|
|
1624
|
+
'component_atoms': ','.join(map(str, comp.atoms)),
|
|
1625
|
+
})
|
|
1626
|
+
|
|
1627
|
+
# Prepare groups data
|
|
1628
|
+
groups_data = []
|
|
1629
|
+
group_counts: Dict[Tuple[Optional[int], str], int] = {}
|
|
1630
|
+
for match in self.matches:
|
|
1631
|
+
if not match.is_group:
|
|
1632
|
+
continue
|
|
1633
|
+
key = (match.group_id, match.group_name or '')
|
|
1634
|
+
group_counts[key] = group_counts.get(key, 0) + 1
|
|
1635
|
+
|
|
1636
|
+
for (group_id, group_name), count in group_counts.items():
|
|
1637
|
+
groups_data.append({
|
|
1638
|
+
'smiles': self.smiles,
|
|
1639
|
+
'group_id': group_id,
|
|
1640
|
+
'group_name': group_name,
|
|
1641
|
+
'match_count': count,
|
|
1642
|
+
})
|
|
1643
|
+
|
|
1644
|
+
# Write to database
|
|
1645
|
+
if components_data:
|
|
1646
|
+
df_components = pd.DataFrame(components_data)
|
|
1647
|
+
df_components.to_sql(components_table, engine, if_exists=if_exists, index=False)
|
|
1648
|
+
|
|
1649
|
+
if groups_data:
|
|
1650
|
+
df_groups = pd.DataFrame(groups_data)
|
|
1651
|
+
df_groups.to_sql(groups_table, engine, if_exists=if_exists, index=False)
|
|
1652
|
+
|
|
1653
|
+
|
|
1654
|
+
class PFASEmbeddingSet(list):
|
|
1655
|
+
"""List-like container for multiple :class:`PFASEmbedding` results.
|
|
1656
|
+
|
|
1657
|
+
Subclasses :class:`list` so existing code that iterates over results
|
|
1658
|
+
continues to work. Call :meth:`to_array` to produce a
|
|
1659
|
+
``(n_molecules, n_columns)`` matrix from all stored results.
|
|
1660
|
+
"""
|
|
1661
|
+
|
|
1662
|
+
def __init__(self, iterable: Iterable[Dict[str, Any]] = ()): # type: ignore[override]
|
|
1663
|
+
super().__init__(PFASEmbedding(m) if not isinstance(m, PFASEmbedding) else m for m in iterable)
|
|
1664
|
+
|
|
1665
|
+
@property
|
|
1666
|
+
def matches(self) -> List[MatchView]:
|
|
1667
|
+
"""Flattened list of all MatchView objects across all molecules.
|
|
1668
|
+
|
|
1669
|
+
Some older code expects a ``matches`` attribute on a ResultsModel
|
|
1670
|
+
instance. Provide a read-only aggregated view by concatenating the
|
|
1671
|
+
per-molecule match lists.
|
|
1672
|
+
"""
|
|
1673
|
+
out: List[MatchView] = []
|
|
1674
|
+
for mol_res in self: # type: ignore[assignment]
|
|
1675
|
+
out.extend(mol_res.matches)
|
|
1676
|
+
return out
|
|
1677
|
+
|
|
1678
|
+
@classmethod
|
|
1679
|
+
def from_raw(cls, results: Iterable[Dict[str, Any]]) -> "PFASEmbeddingSet":
|
|
1680
|
+
"""Wrap an existing list of result dicts without changing them."""
|
|
1681
|
+
|
|
1682
|
+
return cls(results)
|
|
1683
|
+
|
|
1684
|
+
@classmethod
|
|
1685
|
+
def from_smiles(cls, smiles: Union[str, List[str]], **kwargs) -> "PFASEmbeddingSet":
|
|
1686
|
+
"""Parse SMILES string(s) and return a :class:`PFASEmbeddingSet`.
|
|
1687
|
+
|
|
1688
|
+
Parameters
|
|
1689
|
+
----------
|
|
1690
|
+
smiles : str or list of str
|
|
1691
|
+
One or more SMILES strings.
|
|
1692
|
+
**kwargs
|
|
1693
|
+
Forwarded to :func:`~PFASGroups.parser.parse_smiles`
|
|
1694
|
+
(e.g. ``halogens``, ``saturation``, ``progress``).
|
|
1695
|
+
"""
|
|
1696
|
+
from .parser import parse_smiles
|
|
1697
|
+
return parse_smiles(smiles, **kwargs)
|
|
1698
|
+
|
|
1699
|
+
@classmethod
|
|
1700
|
+
def from_mols(cls, mols, **kwargs) -> "PFASEmbeddingSet":
|
|
1701
|
+
"""Parse RDKit molecules and return a :class:`PFASEmbeddingSet`.
|
|
1702
|
+
|
|
1703
|
+
Parameters
|
|
1704
|
+
----------
|
|
1705
|
+
mols : list of rdkit.Chem.Mol
|
|
1706
|
+
List of RDKit molecule objects.
|
|
1707
|
+
**kwargs
|
|
1708
|
+
Forwarded to :func:`~PFASGroups.parser.parse_mols`.
|
|
1709
|
+
"""
|
|
1710
|
+
from .parser import parse_mols
|
|
1711
|
+
return parse_mols(list(mols), **kwargs)
|
|
1712
|
+
|
|
1713
|
+
@classmethod
|
|
1714
|
+
def from_inchis(cls, inchis: List[str], **kwargs) -> "PFASEmbeddingSet":
|
|
1715
|
+
"""Parse InChI strings and return a :class:`PFASEmbeddingSet`.
|
|
1716
|
+
|
|
1717
|
+
Parameters
|
|
1718
|
+
----------
|
|
1719
|
+
inchis : list of str
|
|
1720
|
+
List of InChI strings.
|
|
1721
|
+
**kwargs
|
|
1722
|
+
Forwarded to :func:`~PFASGroups.parser.parse_mols`.
|
|
1723
|
+
"""
|
|
1724
|
+
from rdkit.Chem.inchi import MolFromInchi
|
|
1725
|
+
mols = []
|
|
1726
|
+
for inchi in inchis:
|
|
1727
|
+
mol = MolFromInchi(inchi)
|
|
1728
|
+
if mol is None:
|
|
1729
|
+
raise ValueError(f"Cannot parse InChI: {inchi!r}")
|
|
1730
|
+
mols.append(mol)
|
|
1731
|
+
return cls.from_mols(mols, **kwargs)
|
|
1732
|
+
|
|
1733
|
+
def reorder(self, indices:Union[list,None] = None, key: Callable[["PFASEmbedding"], Any] = None, reverse: bool = False) -> "PFASEmbeddingSet":
|
|
1734
|
+
"""Return a new PFASEmbeddingSet with results reordered by a key function.
|
|
1735
|
+
|
|
1736
|
+
Parameters
|
|
1737
|
+
----------
|
|
1738
|
+
indices : list of int, optional
|
|
1739
|
+
Explicit list of indices defining the new order. If provided, this takes precedence over the key function.
|
|
1740
|
+
key : callable
|
|
1741
|
+
Function that takes a PFASEmbedding and returns a value to sort by.
|
|
1742
|
+
reverse : bool, default False
|
|
1743
|
+
Whether to sort in descending order.
|
|
1744
|
+
"""
|
|
1745
|
+
if indices is not None:
|
|
1746
|
+
if len(indices) != len(self):
|
|
1747
|
+
raise ValueError("Length of indices must match the number of results.")
|
|
1748
|
+
self[:] = [self[i] for i in indices]
|
|
1749
|
+
else:
|
|
1750
|
+
self[:] = sorted(self, key=key, reverse=reverse)
|
|
1751
|
+
return self
|
|
1752
|
+
|
|
1753
|
+
def iter_group_matches(
|
|
1754
|
+
self,
|
|
1755
|
+
group_id: Optional[int] = None,
|
|
1756
|
+
group_name: Optional[str] = None,
|
|
1757
|
+
) -> Iterator[Tuple["PFASEmbedding", MatchView]]:
|
|
1758
|
+
"""Iterate over all PFAS group matches across all molecules."""
|
|
1759
|
+
|
|
1760
|
+
for mol_res in self: # type: ignore[assignment]
|
|
1761
|
+
for match in mol_res.iter_group_matches(group_id=group_id, group_name=group_name):
|
|
1762
|
+
yield mol_res, match
|
|
1763
|
+
|
|
1764
|
+
# --- Plotting helpers ---------------------------------------------------
|
|
1765
|
+
|
|
1766
|
+
def _draw_single_molecule(
|
|
1767
|
+
self,
|
|
1768
|
+
mol: Chem.Mol,
|
|
1769
|
+
highlight_atoms: List[int],
|
|
1770
|
+
legend: str = "",
|
|
1771
|
+
subwidth: int = 300,
|
|
1772
|
+
subheight: int = 300,
|
|
1773
|
+
atom_colours: Optional[Dict[int, Color]] = None,
|
|
1774
|
+
maxFontSize: int = 14,
|
|
1775
|
+
minFontSize: int = 11,
|
|
1776
|
+
) -> Image.Image:
|
|
1777
|
+
"""Draw a single molecule with highlighted atoms as a PIL image.
|
|
1778
|
+
|
|
1779
|
+
Parameters
|
|
1780
|
+
----------
|
|
1781
|
+
atom_colours : dict, optional
|
|
1782
|
+
Mapping of atom index → (R, G, B) float colour. When provided,
|
|
1783
|
+
individual atoms are highlighted with their assigned colour.
|
|
1784
|
+
When ``None`` the default highlight colour is used for all atoms.
|
|
1785
|
+
maxFontSize : int, default 14
|
|
1786
|
+
Maximum font size for legend text.
|
|
1787
|
+
minFontSize : int, default 11
|
|
1788
|
+
Minimum font size for legend text.
|
|
1789
|
+
"""
|
|
1790
|
+
|
|
1791
|
+
d2d = Draw.MolDraw2DCairo(subwidth, subheight)
|
|
1792
|
+
dopts = d2d.drawOptions()
|
|
1793
|
+
dopts.useBWAtomPalette()
|
|
1794
|
+
dopts.fixedBondLength = 20
|
|
1795
|
+
dopts.addAtomIndices = True
|
|
1796
|
+
dopts.addBondIndices = False
|
|
1797
|
+
dopts.maxFontSize = maxFontSize
|
|
1798
|
+
dopts.minFontSize = minFontSize
|
|
1799
|
+
if atom_colours:
|
|
1800
|
+
d2d.DrawMolecule(
|
|
1801
|
+
mol,
|
|
1802
|
+
legend=legend,
|
|
1803
|
+
highlightAtoms=highlight_atoms,
|
|
1804
|
+
highlightAtomColors=atom_colours,
|
|
1805
|
+
)
|
|
1806
|
+
else:
|
|
1807
|
+
d2d.DrawMolecule(mol, legend=legend, highlightAtoms=highlight_atoms)
|
|
1808
|
+
d2d.FinishDrawing()
|
|
1809
|
+
png = d2d.GetDrawingText()
|
|
1810
|
+
return Image.open(BytesIO(png))
|
|
1811
|
+
|
|
1812
|
+
def plot_components_for_group(
|
|
1813
|
+
self,
|
|
1814
|
+
group_id: Optional[int] = None,
|
|
1815
|
+
group_name: Optional[str] = None,
|
|
1816
|
+
max_molecules: Optional[int] = None,
|
|
1817
|
+
subwidth: int = 300,
|
|
1818
|
+
subheight: int = 300,
|
|
1819
|
+
ncols: int = 3,
|
|
1820
|
+
) -> Tuple[Image.Image, int, int]:
|
|
1821
|
+
"""Plot all components for a specific PFAS group across molecules.
|
|
1822
|
+
|
|
1823
|
+
Either ``group_id`` or ``group_name`` (or both) can be provided to
|
|
1824
|
+
select the target group. Each panel corresponds to one molecule,
|
|
1825
|
+
with all its components for that group highlighted together.
|
|
1826
|
+
"""
|
|
1827
|
+
|
|
1828
|
+
imgs: List[Image.Image] = []
|
|
1829
|
+
count = 0
|
|
1830
|
+
|
|
1831
|
+
for mol_res in self: # type: ignore[assignment]
|
|
1832
|
+
if max_molecules is not None and count >= max_molecules:
|
|
1833
|
+
break
|
|
1834
|
+
|
|
1835
|
+
atoms = mol_res.collect_component_atoms(group_id=group_id, group_name=group_name)
|
|
1836
|
+
if not atoms:
|
|
1837
|
+
continue
|
|
1838
|
+
|
|
1839
|
+
# Use mol_with_h to preserve atom ordering from parse_groups_in_mol
|
|
1840
|
+
mol = mol_res.mol_with_h
|
|
1841
|
+
if mol is None:
|
|
1842
|
+
smiles = mol_res.smiles
|
|
1843
|
+
mol = Chem.MolFromSmiles(smiles)
|
|
1844
|
+
if mol is None:
|
|
1845
|
+
continue
|
|
1846
|
+
mol = Chem.AddHs(mol)
|
|
1847
|
+
|
|
1848
|
+
# Try to infer a label from the first matching group
|
|
1849
|
+
label = ""
|
|
1850
|
+
for m in mol_res.iter_group_matches(group_id=group_id, group_name=group_name):
|
|
1851
|
+
if m.group_name is not None:
|
|
1852
|
+
label = m.group_name
|
|
1853
|
+
break
|
|
1854
|
+
|
|
1855
|
+
img = self._draw_single_molecule(mol, atoms, legend=label, subwidth=subwidth, subheight=subheight)
|
|
1856
|
+
imgs.append(img)
|
|
1857
|
+
count += 1
|
|
1858
|
+
|
|
1859
|
+
if not imgs:
|
|
1860
|
+
raise ValueError("No matching components found for the requested group.")
|
|
1861
|
+
|
|
1862
|
+
return _grid_images(imgs, buffer=4, ncols=ncols)
|
|
1863
|
+
|
|
1864
|
+
def show(
|
|
1865
|
+
self,
|
|
1866
|
+
display: bool = True,
|
|
1867
|
+
subwidth: int = 350,
|
|
1868
|
+
subheight: int = 350,
|
|
1869
|
+
ncols: int = 4,
|
|
1870
|
+
) -> Image.Image:
|
|
1871
|
+
"""Show all component combinations in a grid plot.
|
|
1872
|
+
|
|
1873
|
+
Components that share the same highlighted atoms within a molecule are
|
|
1874
|
+
merged into a single panel. The table below each structure lists the
|
|
1875
|
+
matched PFAS group, the component SMARTS type, and three graph metrics:
|
|
1876
|
+
size (C-atom count), branching (1.0 = linear) and mean eccentricity.
|
|
1877
|
+
|
|
1878
|
+
Atoms are highlighted with the colour derived from the component SMARTS
|
|
1879
|
+
metadata (halogen / form / saturation) of the first entry in each panel.
|
|
1880
|
+
"""
|
|
1881
|
+
meta = _get_component_meta()
|
|
1882
|
+
|
|
1883
|
+
imgs: List[Image.Image] = []
|
|
1884
|
+
|
|
1885
|
+
for mol_index, mol_res in enumerate(self): # type: ignore[assignment]
|
|
1886
|
+
# Use the molecule with hydrogens from parsing
|
|
1887
|
+
mol = mol_res.mol_with_h
|
|
1888
|
+
if mol is None:
|
|
1889
|
+
smiles = mol_res.smiles
|
|
1890
|
+
mol = Chem.MolFromSmiles(smiles)
|
|
1891
|
+
if mol is None:
|
|
1892
|
+
continue
|
|
1893
|
+
mol = Chem.AddHs(mol)
|
|
1894
|
+
|
|
1895
|
+
# Group by unique atom set within this molecule
|
|
1896
|
+
from collections import OrderedDict
|
|
1897
|
+
comp_groups: Dict = OrderedDict()
|
|
1898
|
+
|
|
1899
|
+
for match in mol_res.matches:
|
|
1900
|
+
if not match.is_group:
|
|
1901
|
+
continue
|
|
1902
|
+
base_label = match.group_name or match.get("match_id", "")
|
|
1903
|
+
for comp in match.components:
|
|
1904
|
+
atoms = comp.atoms
|
|
1905
|
+
if not atoms:
|
|
1906
|
+
continue
|
|
1907
|
+
key = frozenset(atoms)
|
|
1908
|
+
sl = comp.smarts_label
|
|
1909
|
+
sl_str = str(sl) if isinstance(sl, list) else (sl or "")
|
|
1910
|
+
info = meta.get(sl_str, {})
|
|
1911
|
+
halogen = info.get("halogen")
|
|
1912
|
+
form = info.get("form")
|
|
1913
|
+
saturation = info.get("saturation")
|
|
1914
|
+
colour = _component_color(halogen, form, saturation)
|
|
1915
|
+
if key not in comp_groups:
|
|
1916
|
+
import math as _math
|
|
1917
|
+
br_v = comp.branching
|
|
1918
|
+
ecc_v = comp.mean_eccentricity
|
|
1919
|
+
diam_v = comp.diameter
|
|
1920
|
+
rad_v = comp.radius
|
|
1921
|
+
egr_v = comp.effective_graph_resistance
|
|
1922
|
+
frc_v = comp.component_fraction
|
|
1923
|
+
br_str = f"{br_v:.2f}" if br_v is not None else "\u2014"
|
|
1924
|
+
ecc_str = f"{ecc_v:.2f}" if ecc_v is not None else "\u2014"
|
|
1925
|
+
diam_str = (f"{float(diam_v):.0f}" if diam_v is not None and not (isinstance(diam_v, float) and (_math.isnan(diam_v) or _math.isinf(diam_v))) else "\u2014")
|
|
1926
|
+
rad_str = (f"{float(rad_v):.0f}" if rad_v is not None and not (isinstance(rad_v, float) and (_math.isnan(rad_v) or _math.isinf(rad_v))) else "\u2014")
|
|
1927
|
+
egr_str = (f"{float(egr_v):.2f}" if egr_v is not None and not (isinstance(egr_v, float) and (_math.isnan(egr_v) or _math.isinf(egr_v))) else "\u2014")
|
|
1928
|
+
frc_str = f"{frc_v*100:.0f}%" if frc_v is not None else "\u2014"
|
|
1929
|
+
comp_groups[key] = {
|
|
1930
|
+
'atoms': sorted(atoms),
|
|
1931
|
+
'colour': colour,
|
|
1932
|
+
'halogen': halogen,
|
|
1933
|
+
'entries': [],
|
|
1934
|
+
'comp_metrics': (str(comp.size), br_str, ecc_str, diam_str, rad_str, egr_str, frc_str),
|
|
1935
|
+
}
|
|
1936
|
+
dct_v = comp.min_dist_to_centre
|
|
1937
|
+
dpr_v = comp.max_dist_to_periphery
|
|
1938
|
+
dct_str = str(dct_v) if dct_v is not None else "\u2014"
|
|
1939
|
+
dpr_str = str(dpr_v) if dpr_v is not None else "\u2014"
|
|
1940
|
+
entry = (base_label, sl_str, dct_str, dpr_str)
|
|
1941
|
+
if entry not in comp_groups[key]['entries']:
|
|
1942
|
+
comp_groups[key]['entries'].append(entry)
|
|
1943
|
+
|
|
1944
|
+
for data in comp_groups.values():
|
|
1945
|
+
atoms = data['atoms']
|
|
1946
|
+
colour = data['colour']
|
|
1947
|
+
entries = data['entries']
|
|
1948
|
+
comp_metrics = data.get('comp_metrics')
|
|
1949
|
+
halogen_lbl = data.get('halogen') or ""
|
|
1950
|
+
|
|
1951
|
+
atom_colours: Dict[int, Color] = {a: colour for a in atoms}
|
|
1952
|
+
d2d = Draw.MolDraw2DCairo(subwidth, subheight)
|
|
1953
|
+
dopts = d2d.drawOptions()
|
|
1954
|
+
dopts.useBWAtomPalette()
|
|
1955
|
+
dopts.fixedBondLength = 20
|
|
1956
|
+
dopts.addAtomIndices = True
|
|
1957
|
+
dopts.addBondIndices = False
|
|
1958
|
+
dopts.maxFontSize = 14
|
|
1959
|
+
dopts.minFontSize = 12
|
|
1960
|
+
d2d.DrawMolecule(
|
|
1961
|
+
mol,
|
|
1962
|
+
highlightAtoms=atoms,
|
|
1963
|
+
highlightAtomColors=atom_colours,
|
|
1964
|
+
)
|
|
1965
|
+
d2d.FinishDrawing()
|
|
1966
|
+
mol_img = Image.open(BytesIO(d2d.GetDrawingText()))
|
|
1967
|
+
imgs.append(_mol_image_with_table(
|
|
1968
|
+
mol_img, entries, comp_metrics=comp_metrics,
|
|
1969
|
+
mol_label=f"mol#{mol_index + 1}", halogen_label=halogen_lbl
|
|
1970
|
+
))
|
|
1971
|
+
|
|
1972
|
+
if not imgs:
|
|
1973
|
+
raise ValueError("No PFAS group components found to display.")
|
|
1974
|
+
|
|
1975
|
+
grid, _, _ = _grid_images(imgs, buffer=4, ncols=ncols)
|
|
1976
|
+
if display:
|
|
1977
|
+
grid.show()
|
|
1978
|
+
return grid
|
|
1979
|
+
|
|
1980
|
+
# Alias so callers can use either results.show() or results.plot()
|
|
1981
|
+
plot = show
|
|
1982
|
+
|
|
1983
|
+
def to_sql(
|
|
1984
|
+
self,
|
|
1985
|
+
filename: Optional[str] = None,
|
|
1986
|
+
dbname: Optional[str] = None,
|
|
1987
|
+
user: Optional[str] = None,
|
|
1988
|
+
password: Optional[str] = None,
|
|
1989
|
+
host: Optional[str] = None,
|
|
1990
|
+
port: Optional[int] = None,
|
|
1991
|
+
components_table: str = "components",
|
|
1992
|
+
groups_table: str = "pfas_groups_in_compound",
|
|
1993
|
+
if_exists: str = "append",
|
|
1994
|
+
) -> None:
|
|
1995
|
+
"""Export this molecule result to a SQL database.
|
|
1996
|
+
|
|
1997
|
+
Can write to either SQLite (via filename) or PostgreSQL/MySQL (via connection parameters).
|
|
1998
|
+
|
|
1999
|
+
Parameters
|
|
2000
|
+
----------
|
|
2001
|
+
filename : str, optional
|
|
2002
|
+
Path to SQLite database file. If provided, uses SQLite.
|
|
2003
|
+
dbname : str, optional
|
|
2004
|
+
Database name (for PostgreSQL/MySQL).
|
|
2005
|
+
user : str, optional
|
|
2006
|
+
Database username. Defaults to os.environ['DB_USER'] if not provided.
|
|
2007
|
+
password : str, optional
|
|
2008
|
+
Database password. Defaults to os.environ['DB_PASSWORD'] if not provided.
|
|
2009
|
+
host : str, optional
|
|
2010
|
+
Database host. Defaults to os.environ.get('DB_HOST', 'localhost').
|
|
2011
|
+
port : int, optional
|
|
2012
|
+
Database port. Defaults to os.environ.get('DB_PORT', 5432 for PostgreSQL).
|
|
2013
|
+
components_table : str, default "components"
|
|
2014
|
+
Name of the table to store component-level data.
|
|
2015
|
+
groups_table : str, default "pfas_groups_in_compound"
|
|
2016
|
+
Name of the table to store PFAS group matches.
|
|
2017
|
+
if_exists : str, default "append"
|
|
2018
|
+
How to behave if tables exist: 'fail', 'replace', or 'append'.
|
|
2019
|
+
"""
|
|
2020
|
+
try:
|
|
2021
|
+
import pandas as pd
|
|
2022
|
+
import sqlalchemy
|
|
2023
|
+
except ImportError as exc:
|
|
2024
|
+
raise ImportError("pandas and sqlalchemy are required for to_sql. Install with: pip install pandas sqlalchemy") from exc
|
|
2025
|
+
# Determine connection
|
|
2026
|
+
if filename:
|
|
2027
|
+
engine = sqlalchemy.create_engine(f"sqlite:///{filename}")
|
|
2028
|
+
elif dbname:
|
|
2029
|
+
# Get credentials from environment if not provided
|
|
2030
|
+
if user is None:
|
|
2031
|
+
user = os.environ.get('DB_USER')
|
|
2032
|
+
if password is None:
|
|
2033
|
+
password = os.environ.get('DB_PASSWORD')
|
|
2034
|
+
if host is None:
|
|
2035
|
+
host = os.environ.get('DB_HOST', 'localhost')
|
|
2036
|
+
if port is None:
|
|
2037
|
+
port = int(os.environ.get('DB_PORT', 5432))
|
|
2038
|
+
|
|
2039
|
+
if not user or not password:
|
|
2040
|
+
raise ValueError("Database credentials required. Provide user/password or set DB_USER/DB_PASSWORD environment variables.")
|
|
2041
|
+
|
|
2042
|
+
# Assuming PostgreSQL; adjust for MySQL if needed
|
|
2043
|
+
connection_string = f"postgresql://{user}:{password}@{host}:{port}/{dbname}"
|
|
2044
|
+
engine = sqlalchemy.create_engine(connection_string)
|
|
2045
|
+
else:
|
|
2046
|
+
raise ValueError("Either filename (for SQLite) or dbname (for PostgreSQL) must be provided.")
|
|
2047
|
+
|
|
2048
|
+
# Prepare components data across all molecules in this ResultsModel
|
|
2049
|
+
components_data = []
|
|
2050
|
+
groups_data = []
|
|
2051
|
+
|
|
2052
|
+
for mol_res in self: # type: ignore[assignment]
|
|
2053
|
+
# local counts per molecule
|
|
2054
|
+
local_group_counts: Dict[Tuple[Optional[int], str], int] = {}
|
|
2055
|
+
for match in mol_res.matches:
|
|
2056
|
+
if not match.is_group:
|
|
2057
|
+
continue
|
|
2058
|
+
for comp in match.components:
|
|
2059
|
+
smarts = comp.smarts_label
|
|
2060
|
+
components_data.append({
|
|
2061
|
+
'smiles': mol_res.smiles,
|
|
2062
|
+
'group_id': match.group_id,
|
|
2063
|
+
'group_name': match.group_name,
|
|
2064
|
+
'smarts_label': str(smarts) if isinstance(smarts, list) else smarts,
|
|
2065
|
+
'component_atoms': ','.join(map(str, comp.atoms)),
|
|
2066
|
+
})
|
|
2067
|
+
|
|
2068
|
+
key = (match.group_id, match.group_name or '')
|
|
2069
|
+
local_group_counts[key] = local_group_counts.get(key, 0) + 1
|
|
2070
|
+
|
|
2071
|
+
for (group_id, group_name), count in local_group_counts.items():
|
|
2072
|
+
groups_data.append({
|
|
2073
|
+
'smiles': mol_res.smiles,
|
|
2074
|
+
'group_id': group_id,
|
|
2075
|
+
'group_name': group_name,
|
|
2076
|
+
'match_count': count,
|
|
2077
|
+
})
|
|
2078
|
+
|
|
2079
|
+
# Write to database
|
|
2080
|
+
if components_data:
|
|
2081
|
+
df_components = pd.DataFrame(components_data)
|
|
2082
|
+
df_components.to_sql(components_table, engine, if_exists=if_exists, index=False)
|
|
2083
|
+
|
|
2084
|
+
if groups_data:
|
|
2085
|
+
df_groups = pd.DataFrame(groups_data)
|
|
2086
|
+
df_groups.to_sql(groups_table, engine, if_exists=if_exists, index=False)
|
|
2087
|
+
|
|
2088
|
+
def svg(
|
|
2089
|
+
self,
|
|
2090
|
+
filename: str,
|
|
2091
|
+
subwidth: int = 350,
|
|
2092
|
+
subheight: int = 350,
|
|
2093
|
+
ncols: int = 4,
|
|
2094
|
+
) -> str:
|
|
2095
|
+
"""Export all component combinations to an SVG file (vector graphics).
|
|
2096
|
+
|
|
2097
|
+
Components that share the same highlighted atoms within a molecule are
|
|
2098
|
+
merged into a single panel with a bullet-point legend.
|
|
2099
|
+
|
|
2100
|
+
Parameters
|
|
2101
|
+
----------
|
|
2102
|
+
filename : str
|
|
2103
|
+
Path to the output SVG file.
|
|
2104
|
+
subwidth : int, default 350
|
|
2105
|
+
Width of each sub-image in pixels.
|
|
2106
|
+
subheight : int, default 350
|
|
2107
|
+
Minimum height of each sub-image in pixels.
|
|
2108
|
+
ncols : int, default 4
|
|
2109
|
+
Number of columns in the grid.
|
|
2110
|
+
|
|
2111
|
+
Returns
|
|
2112
|
+
-------
|
|
2113
|
+
str
|
|
2114
|
+
Path to the created SVG file.
|
|
2115
|
+
"""
|
|
2116
|
+
import svgutils.transform as sg
|
|
2117
|
+
|
|
2118
|
+
imgs: List[str] = []
|
|
2119
|
+
|
|
2120
|
+
for mol_index, mol_res in enumerate(self): # type: ignore[assignment]
|
|
2121
|
+
# Use the molecule with hydrogens from parsing
|
|
2122
|
+
mol = mol_res.mol_with_h
|
|
2123
|
+
if mol is None:
|
|
2124
|
+
smiles = mol_res.smiles
|
|
2125
|
+
mol = Chem.MolFromSmiles(smiles)
|
|
2126
|
+
if mol is None:
|
|
2127
|
+
continue
|
|
2128
|
+
mol = Chem.AddHs(mol)
|
|
2129
|
+
|
|
2130
|
+
# Group by unique atom set within this molecule
|
|
2131
|
+
from collections import OrderedDict
|
|
2132
|
+
comp_groups: Dict = OrderedDict()
|
|
2133
|
+
|
|
2134
|
+
for match in mol_res.matches:
|
|
2135
|
+
if not match.is_group:
|
|
2136
|
+
continue
|
|
2137
|
+
base_label = match.group_name or match.get("match_id", "")
|
|
2138
|
+
for comp in match.components:
|
|
2139
|
+
atoms = comp.atoms
|
|
2140
|
+
if not atoms:
|
|
2141
|
+
continue
|
|
2142
|
+
key = frozenset(atoms)
|
|
2143
|
+
sl = comp.smarts_label
|
|
2144
|
+
sl_str = str(sl) if isinstance(sl, list) else (sl or "")
|
|
2145
|
+
if key not in comp_groups:
|
|
2146
|
+
import math as _math
|
|
2147
|
+
br_v = comp.branching
|
|
2148
|
+
ecc_v = comp.mean_eccentricity
|
|
2149
|
+
diam_v = comp.diameter
|
|
2150
|
+
rad_v = comp.radius
|
|
2151
|
+
egr_v = comp.effective_graph_resistance
|
|
2152
|
+
frc_v = comp.component_fraction
|
|
2153
|
+
br_str = f"{br_v:.2f}" if br_v is not None else "\u2014"
|
|
2154
|
+
ecc_str = f"{ecc_v:.2f}" if ecc_v is not None else "\u2014"
|
|
2155
|
+
diam_str = (f"{float(diam_v):.0f}" if diam_v is not None and not (isinstance(diam_v, float) and (_math.isnan(diam_v) or _math.isinf(diam_v))) else "\u2014")
|
|
2156
|
+
rad_str = (f"{float(rad_v):.0f}" if rad_v is not None and not (isinstance(rad_v, float) and (_math.isnan(rad_v) or _math.isinf(rad_v))) else "\u2014")
|
|
2157
|
+
egr_str = (f"{float(egr_v):.2f}" if egr_v is not None and not (isinstance(egr_v, float) and (_math.isnan(egr_v) or _math.isinf(egr_v))) else "\u2014")
|
|
2158
|
+
frc_str = f"{frc_v*100:.0f}%" if frc_v is not None else "\u2014"
|
|
2159
|
+
comp_groups[key] = {
|
|
2160
|
+
'atoms': sorted(atoms),
|
|
2161
|
+
'entries': [],
|
|
2162
|
+
'comp_metrics': (str(comp.size), br_str, ecc_str, diam_str, rad_str, egr_str, frc_str),
|
|
2163
|
+
}
|
|
2164
|
+
dct_v = comp.min_dist_to_centre
|
|
2165
|
+
dpr_v = comp.max_dist_to_periphery
|
|
2166
|
+
dct_str = str(dct_v) if dct_v is not None else "\u2014"
|
|
2167
|
+
dpr_str = str(dpr_v) if dpr_v is not None else "\u2014"
|
|
2168
|
+
entry = (base_label, sl_str, dct_str, dpr_str)
|
|
2169
|
+
if entry not in comp_groups[key]['entries']:
|
|
2170
|
+
comp_groups[key]['entries'].append(entry)
|
|
2171
|
+
|
|
2172
|
+
for data in comp_groups.values():
|
|
2173
|
+
atoms = data['atoms']
|
|
2174
|
+
entries = data['entries']
|
|
2175
|
+
comp_metrics = data.get('comp_metrics')
|
|
2176
|
+
n = len(entries)
|
|
2177
|
+
|
|
2178
|
+
lines: List[str] = [f"mol#{mol_index + 1}"]
|
|
2179
|
+
for grp, sls, dct, dpr in entries:
|
|
2180
|
+
line = f"\u2022 {grp}"
|
|
2181
|
+
if sls:
|
|
2182
|
+
line += f" | {sls}"
|
|
2183
|
+
line += f" dct={dct} dpr={dpr}"
|
|
2184
|
+
lines.append(line)
|
|
2185
|
+
if comp_metrics:
|
|
2186
|
+
sz, br, ecc, diam, rad, egr, frc = comp_metrics
|
|
2187
|
+
lines.append(f" size={sz} br={br} ecc={ecc} diam={diam} rad={rad} egr={egr} chain={frc}")
|
|
2188
|
+
legend = "\n".join(lines)
|
|
2189
|
+
|
|
2190
|
+
effective_height = max(subheight, 220 + (n + 1) * 38)
|
|
2191
|
+
|
|
2192
|
+
d2d = Draw.MolDraw2DSVG(subwidth, effective_height)
|
|
2193
|
+
dopts = d2d.drawOptions()
|
|
2194
|
+
dopts.useBWAtomPalette()
|
|
2195
|
+
dopts.fixedBondLength = 20
|
|
2196
|
+
dopts.addAtomIndices = True
|
|
2197
|
+
dopts.addBondIndices = False
|
|
2198
|
+
dopts.maxFontSize = 14
|
|
2199
|
+
dopts.minFontSize = 11
|
|
2200
|
+
d2d.DrawMolecule(mol, legend=legend, highlightAtoms=atoms)
|
|
2201
|
+
d2d.FinishDrawing()
|
|
2202
|
+
imgs.append(d2d.GetDrawingText())
|
|
2203
|
+
|
|
2204
|
+
if not imgs:
|
|
2205
|
+
raise ValueError("No PFAS group components found to display.")
|
|
2206
|
+
|
|
2207
|
+
# Convert SVG strings to svgutils figures
|
|
2208
|
+
svg_figs = [sg.fromstring(img) for img in imgs]
|
|
2209
|
+
|
|
2210
|
+
# Merge into grid
|
|
2211
|
+
from .draw_mols import merge_svg
|
|
2212
|
+
grid, _, _ = merge_svg(svg_figs, buffer=4, ncols=ncols)
|
|
2213
|
+
|
|
2214
|
+
grid.save(filename)
|
|
2215
|
+
return filename
|
|
2216
|
+
|
|
2217
|
+
def summarise(self) -> str:
|
|
2218
|
+
"""Return a coloured text summary of the results.
|
|
2219
|
+
|
|
2220
|
+
The summary includes:
|
|
2221
|
+
- number of molecules
|
|
2222
|
+
- counts of PFAS group and definition matches
|
|
2223
|
+
- total number of components across all group matches
|
|
2224
|
+
- the most frequent PFAS groups (colour-coded by halogen)
|
|
2225
|
+
"""
|
|
2226
|
+
meta = _get_component_meta()
|
|
2227
|
+
|
|
2228
|
+
total_molecules = len(self)
|
|
2229
|
+
total_group_matches = 0
|
|
2230
|
+
total_definition_matches = 0
|
|
2231
|
+
total_components = 0
|
|
2232
|
+
group_counts: Dict[str, int] = {}
|
|
2233
|
+
# map group name → first halogen seen
|
|
2234
|
+
group_halogen: Dict[str, Optional[str]] = {}
|
|
2235
|
+
|
|
2236
|
+
for mol_res in self: # type: ignore[assignment]
|
|
2237
|
+
for m in mol_res.matches:
|
|
2238
|
+
if m.is_group:
|
|
2239
|
+
total_group_matches += 1
|
|
2240
|
+
total_components += len(m.components)
|
|
2241
|
+
name = m.group_name or str(m.get("match_id", ""))
|
|
2242
|
+
group_counts[name] = group_counts.get(name, 0) + 1
|
|
2243
|
+
if name not in group_halogen:
|
|
2244
|
+
for comp in m.components:
|
|
2245
|
+
sl = comp.smarts_label
|
|
2246
|
+
sl_str = str(sl) if isinstance(sl, list) else sl
|
|
2247
|
+
info = meta.get(sl_str or "", {})
|
|
2248
|
+
group_halogen[name] = info.get("halogen")
|
|
2249
|
+
break
|
|
2250
|
+
elif m.is_definition:
|
|
2251
|
+
total_definition_matches += 1
|
|
2252
|
+
|
|
2253
|
+
unique_groups = len(group_counts)
|
|
2254
|
+
|
|
2255
|
+
lines: List[str] = []
|
|
2256
|
+
lines.append(_ansi("PFASEmbeddingSet summary", _ANSI_BOLD))
|
|
2257
|
+
lines.append(f"- Molecules: {total_molecules}")
|
|
2258
|
+
lines.append(
|
|
2259
|
+
f"- PFAS group matches: {total_group_matches} (unique groups: {unique_groups})"
|
|
2260
|
+
)
|
|
2261
|
+
lines.append(f"- PFAS definition matches: {total_definition_matches}")
|
|
2262
|
+
lines.append(
|
|
2263
|
+
f"- Total components across all PFAS group matches: {total_components}"
|
|
2264
|
+
)
|
|
2265
|
+
|
|
2266
|
+
if group_counts:
|
|
2267
|
+
lines.append("- Top PFAS groups (by number of matches):")
|
|
2268
|
+
for name, count in sorted(
|
|
2269
|
+
group_counts.items(), key=lambda kv: kv[1], reverse=True
|
|
2270
|
+
)[:10]:
|
|
2271
|
+
halogen = group_halogen.get(name)
|
|
2272
|
+
hal_code = _ANSI_HALOGEN.get(halogen or "", "")
|
|
2273
|
+
lines.append(f" * {hal_code}{_ANSI_BOLD}{name}{_ANSI_RESET}: {count} match(es)")
|
|
2274
|
+
|
|
2275
|
+
return "\n".join(lines)
|
|
2276
|
+
|
|
2277
|
+
def __str__(self) -> str:
|
|
2278
|
+
return self.summarise()
|
|
2279
|
+
|
|
2280
|
+
def table(self) -> str:
|
|
2281
|
+
"""Return a more detailed text table with one row per molecule.
|
|
2282
|
+
|
|
2283
|
+
The TSV table has the following columns: ``index`` (1-based),
|
|
2284
|
+
``smiles``, ``group_matches`` (count), ``definition_matches`` (count),
|
|
2285
|
+
and ``groups`` (per-molecule PFAS groups with counts, e.g.
|
|
2286
|
+
``"Perfluoroalkyl (2); Polyfluoroalkyl (1)"``).
|
|
2287
|
+
"""
|
|
2288
|
+
|
|
2289
|
+
lines: List[str] = []
|
|
2290
|
+
lines.append(
|
|
2291
|
+
"index\tsmiles\tgroup_matches\tdefinition_matches\tgroups"
|
|
2292
|
+
)
|
|
2293
|
+
|
|
2294
|
+
for idx, mol_res in enumerate(self, start=1): # type: ignore[assignment]
|
|
2295
|
+
smiles = mol_res.smiles
|
|
2296
|
+
group_matches = [m for m in mol_res.matches if m.is_group]
|
|
2297
|
+
definition_matches = [m for m in mol_res.matches if m.is_definition]
|
|
2298
|
+
|
|
2299
|
+
per_mol_group_counts: Dict[str, int] = {}
|
|
2300
|
+
for m in group_matches:
|
|
2301
|
+
name = m.group_name or str(m.get("match_id", ""))
|
|
2302
|
+
per_mol_group_counts[name] = per_mol_group_counts.get(name, 0) + 1
|
|
2303
|
+
|
|
2304
|
+
if per_mol_group_counts:
|
|
2305
|
+
groups_summary = "; ".join(
|
|
2306
|
+
f"{name} ({count})"
|
|
2307
|
+
for name, count in sorted(
|
|
2308
|
+
per_mol_group_counts.items(), key=lambda kv: kv[1], reverse=True
|
|
2309
|
+
)
|
|
2310
|
+
)
|
|
2311
|
+
else:
|
|
2312
|
+
groups_summary = "None"
|
|
2313
|
+
|
|
2314
|
+
lines.append(
|
|
2315
|
+
f"{idx}\t{smiles}\t{len(group_matches)}\t{len(definition_matches)}\t{groups_summary}"
|
|
2316
|
+
)
|
|
2317
|
+
|
|
2318
|
+
return "\n".join(lines)
|
|
2319
|
+
|
|
2320
|
+
def classify(self) -> pd.DataFrame:
|
|
2321
|
+
"""Return a classification DataFrame with one row per molecule.
|
|
2322
|
+
|
|
2323
|
+
Each molecule is classified by :meth:`MoleculeResult.classify`.
|
|
2324
|
+
|
|
2325
|
+
Returns
|
|
2326
|
+
-------
|
|
2327
|
+
pandas.DataFrame
|
|
2328
|
+
Columns:
|
|
2329
|
+
|
|
2330
|
+
* ``smiles`` — molecule SMILES.
|
|
2331
|
+
* ``category`` — classification label: OECD group name(s) if
|
|
2332
|
+
matched, otherwise ``"per-"``/``"poly-"`` + generic/telomeric
|
|
2333
|
+
group names (comma-separated).
|
|
2334
|
+
* ``total_component_size`` — sum of C-atom counts across all
|
|
2335
|
+
matched group components.
|
|
2336
|
+
"""
|
|
2337
|
+
rows = []
|
|
2338
|
+
for mol_res in self: # type: ignore[assignment]
|
|
2339
|
+
category, total_size = mol_res.classify()
|
|
2340
|
+
rows.append({
|
|
2341
|
+
"smiles": mol_res.smiles,
|
|
2342
|
+
"category": category,
|
|
2343
|
+
"total_component_size": total_size,
|
|
2344
|
+
})
|
|
2345
|
+
return pd.DataFrame(rows, columns=["smiles", "category", "total_component_size"])
|
|
2346
|
+
|
|
2347
|
+
def summary(self) -> None:
|
|
2348
|
+
"""Print a detailed coloured summary of matched groups across all molecules.
|
|
2349
|
+
|
|
2350
|
+
For each group, shows the component SMARTS type and, per component,
|
|
2351
|
+
the graph metrics: size (C-atom count), branching and mean eccentricity.
|
|
2352
|
+
Component size statistics (min, max, mean) are also shown.
|
|
2353
|
+
"""
|
|
2354
|
+
|
|
2355
|
+
print("=" * 80)
|
|
2356
|
+
print(f"{_ANSI_BOLD}RESULTS SUMMARY:{_ANSI_RESET} {len(self)} molecule(s)")
|
|
2357
|
+
print("=" * 80)
|
|
2358
|
+
|
|
2359
|
+
if not self:
|
|
2360
|
+
print("No molecules in results.")
|
|
2361
|
+
return
|
|
2362
|
+
|
|
2363
|
+
# Collect all group matches across all molecules
|
|
2364
|
+
all_groups_info: Dict[Tuple[Optional[int], str], List[ComponentView]] = {}
|
|
2365
|
+
|
|
2366
|
+
for mol_res in self: # type: ignore[assignment]
|
|
2367
|
+
for m in mol_res.matches:
|
|
2368
|
+
if not m.is_group:
|
|
2369
|
+
continue
|
|
2370
|
+
key = (m.group_id, m.group_name or "Unknown")
|
|
2371
|
+
if key not in all_groups_info:
|
|
2372
|
+
all_groups_info[key] = []
|
|
2373
|
+
all_groups_info[key].extend(m.components)
|
|
2374
|
+
|
|
2375
|
+
if not all_groups_info:
|
|
2376
|
+
print("\nNo PFAS groups matched across all molecules.")
|
|
2377
|
+
return
|
|
2378
|
+
|
|
2379
|
+
print(f"\nMatched {len(all_groups_info)} unique PFAS group(s) across all molecules:")
|
|
2380
|
+
print()
|
|
2381
|
+
|
|
2382
|
+
for (group_id, group_name), components in sorted(all_groups_info.items(), key=lambda x: x[0][0] or 0):
|
|
2383
|
+
print(f"{_ANSI_BOLD}Group {group_id}: {group_name}{_ANSI_RESET}")
|
|
2384
|
+
|
|
2385
|
+
# Group components by SMARTS type
|
|
2386
|
+
by_smarts: Dict[Optional[str], List[ComponentView]] = {}
|
|
2387
|
+
for comp in components:
|
|
2388
|
+
smarts = comp.smarts_label
|
|
2389
|
+
smarts_key = str(smarts) if isinstance(smarts, list) else smarts
|
|
2390
|
+
if smarts_key not in by_smarts:
|
|
2391
|
+
by_smarts[smarts_key] = []
|
|
2392
|
+
by_smarts[smarts_key].append(comp)
|
|
2393
|
+
|
|
2394
|
+
# Display components by SMARTS type
|
|
2395
|
+
for smarts_label, comps in sorted(by_smarts.items(), key=lambda x: x[0] or ""):
|
|
2396
|
+
label = smarts_label or "(no label)"
|
|
2397
|
+
sizes = [comp.size for comp in comps]
|
|
2398
|
+
min_sz = min(sizes) if sizes else 0
|
|
2399
|
+
max_sz = max(sizes) if sizes else 0
|
|
2400
|
+
mean_sz = sum(sizes) / len(sizes) if sizes else 0.0
|
|
2401
|
+
print(f" SMARTS: {label} ({len(comps)} component(s)) "
|
|
2402
|
+
f"size {min_sz}\u2013{max_sz} (mean {mean_sz:.1f})")
|
|
2403
|
+
for comp in sorted(comps, key=lambda c: c.size, reverse=True):
|
|
2404
|
+
br_str = f"{comp.branching:.2f}" if comp.branching is not None else "\u2014"
|
|
2405
|
+
ecc_str = f"{comp.mean_eccentricity:.2f}" if comp.mean_eccentricity is not None else "\u2014"
|
|
2406
|
+
print(f" size={_ANSI_BOLD}{comp.size}{_ANSI_RESET} branching={br_str} eccentricity={ecc_str}")
|
|
2407
|
+
|
|
2408
|
+
print()
|
|
2409
|
+
|
|
2410
|
+
print("=" * 80)
|
|
2411
|
+
|
|
2412
|
+
def plot_all_components_with_group_colours(
|
|
2413
|
+
self,
|
|
2414
|
+
max_molecules: Optional[int] = None,
|
|
2415
|
+
subwidth: int = 300,
|
|
2416
|
+
subheight: int = 300,
|
|
2417
|
+
ncols: int = 3,
|
|
2418
|
+
) -> Tuple[Image.Image, int, int]:
|
|
2419
|
+
"""Plot all matched components, coloured by PFAS group.
|
|
2420
|
+
|
|
2421
|
+
Each panel corresponds to one molecule; atoms are highlighted with
|
|
2422
|
+
colours assigned per PFAS group. The legend lists the groups found
|
|
2423
|
+
in that molecule.
|
|
2424
|
+
"""
|
|
2425
|
+
|
|
2426
|
+
imgs: List[Image.Image] = []
|
|
2427
|
+
count = 0
|
|
2428
|
+
|
|
2429
|
+
for mol_res in self: # type: ignore[assignment]
|
|
2430
|
+
if max_molecules is not None and count >= max_molecules:
|
|
2431
|
+
break
|
|
2432
|
+
|
|
2433
|
+
# Use the molecule with hydrogens from parsing
|
|
2434
|
+
mol = mol_res.mol_with_h
|
|
2435
|
+
if mol is None:
|
|
2436
|
+
# Fallback to SMILES
|
|
2437
|
+
smiles = mol_res.smiles
|
|
2438
|
+
mol = Chem.MolFromSmiles(smiles)
|
|
2439
|
+
if mol is None:
|
|
2440
|
+
continue
|
|
2441
|
+
mol = Chem.AddHs(mol)
|
|
2442
|
+
|
|
2443
|
+
# Build colour map per group
|
|
2444
|
+
atom_colours: Dict[int, Color] = {}
|
|
2445
|
+
group_labels: List[str] = []
|
|
2446
|
+
group_index: Dict[str, int] = {}
|
|
2447
|
+
|
|
2448
|
+
for match in mol_res.matches:
|
|
2449
|
+
if not match.is_group:
|
|
2450
|
+
continue
|
|
2451
|
+
gid = match.get("match_id") or f"G{match.group_id}" # type: ignore[operator]
|
|
2452
|
+
if gid not in group_index:
|
|
2453
|
+
group_index[gid] = len(group_index)
|
|
2454
|
+
colour = _GROUP_COLORS[group_index[gid] % len(_GROUP_COLORS)]
|
|
2455
|
+
|
|
2456
|
+
label_name = match.group_name or gid
|
|
2457
|
+
if label_name not in group_labels:
|
|
2458
|
+
group_labels.append(label_name)
|
|
2459
|
+
|
|
2460
|
+
for comp in match.components:
|
|
2461
|
+
for atom_idx in comp.atoms:
|
|
2462
|
+
atom_colours.setdefault(atom_idx, colour)
|
|
2463
|
+
|
|
2464
|
+
if not atom_colours:
|
|
2465
|
+
continue
|
|
2466
|
+
|
|
2467
|
+
# Legend lists all group names present in this molecule
|
|
2468
|
+
legend = ", ".join(group_labels)
|
|
2469
|
+
|
|
2470
|
+
d2d = Draw.MolDraw2DCairo(subwidth, subheight)
|
|
2471
|
+
dopts = d2d.drawOptions()
|
|
2472
|
+
dopts.useBWAtomPalette()
|
|
2473
|
+
dopts.fixedBondLength = 20
|
|
2474
|
+
dopts.addAtomIndices = True
|
|
2475
|
+
dopts.addBondIndices = False
|
|
2476
|
+
dopts.maxFontSize = 16
|
|
2477
|
+
dopts.minFontSize = 13
|
|
2478
|
+
|
|
2479
|
+
highlight_atoms = list(atom_colours.keys())
|
|
2480
|
+
d2d.DrawMolecule(mol, legend=legend, highlightAtoms=highlight_atoms, highlightAtomColors=atom_colours)
|
|
2481
|
+
d2d.FinishDrawing()
|
|
2482
|
+
png = d2d.GetDrawingText()
|
|
2483
|
+
imgs.append(Image.open(BytesIO(png)))
|
|
2484
|
+
count += 1
|
|
2485
|
+
|
|
2486
|
+
if not imgs:
|
|
2487
|
+
raise ValueError("No PFAS group components found in results.")
|
|
2488
|
+
|
|
2489
|
+
return _grid_images(imgs, buffer=4, ncols=ncols)
|
|
2490
|
+
|
|
2491
|
+
def to_sql_all(
|
|
2492
|
+
self,
|
|
2493
|
+
conn: Optional[Union[str, 'sqlalchemy.engine.Engine']] = None,
|
|
2494
|
+
filename: Optional[str] = None,
|
|
2495
|
+
components_table: str = "components",
|
|
2496
|
+
groups_table: str = "pfas_groups_in_compound",
|
|
2497
|
+
if_exists: str = "append",
|
|
2498
|
+
) -> None:
|
|
2499
|
+
"""Export all molecule results to a SQL database.
|
|
2500
|
+
|
|
2501
|
+
This method efficiently batches all molecules into the database in a single operation.
|
|
2502
|
+
|
|
2503
|
+
Parameters
|
|
2504
|
+
----------
|
|
2505
|
+
conn : str or sqlalchemy.engine.Engine, optional
|
|
2506
|
+
Database connection. Can be:
|
|
2507
|
+
- SQLAlchemy Engine object
|
|
2508
|
+
- Connection string (e.g., 'postgresql://user:pass@host:port/db')
|
|
2509
|
+
- SQLite path with 'sqlite:///' prefix
|
|
2510
|
+
filename : str, optional
|
|
2511
|
+
Path to SQLite database file (legacy parameter, use conn instead).
|
|
2512
|
+
components_table : str, default "components"
|
|
2513
|
+
Name of the table to store component-level data.
|
|
2514
|
+
groups_table : str, default "pfas_groups_in_compound"
|
|
2515
|
+
Name of the table to store PFAS group matches.
|
|
2516
|
+
if_exists : str, default "append"
|
|
2517
|
+
How to behave if tables exist: 'fail', 'replace', or 'append'.
|
|
2518
|
+
|
|
2519
|
+
Examples
|
|
2520
|
+
--------
|
|
2521
|
+
>>> # Using connection string
|
|
2522
|
+
>>> results.to_sql(conn='postgresql://user:pass@localhost/pfas_db')
|
|
2523
|
+
>>>
|
|
2524
|
+
>>> # Using SQLAlchemy engine
|
|
2525
|
+
>>> from sqlalchemy import create_engine
|
|
2526
|
+
>>> engine = create_engine('sqlite:///pfas.db')
|
|
2527
|
+
>>> results.to_sql_all(conn=engine)
|
|
2528
|
+
>>>
|
|
2529
|
+
>>> # Using filename (legacy)
|
|
2530
|
+
>>> results.to_sql(filename='pfas.db')
|
|
2531
|
+
"""
|
|
2532
|
+
try:
|
|
2533
|
+
import pandas as pd
|
|
2534
|
+
import sqlalchemy
|
|
2535
|
+
except ImportError as exc:
|
|
2536
|
+
raise ImportError("pandas and sqlalchemy are required for to_sql. Install with: pip install pandas sqlalchemy") from exc
|
|
2537
|
+
# Determine connection
|
|
2538
|
+
if conn is None and filename is None:
|
|
2539
|
+
raise ValueError("Either 'conn' or 'filename' must be provided.")
|
|
2540
|
+
|
|
2541
|
+
if conn is not None:
|
|
2542
|
+
# Handle conn parameter
|
|
2543
|
+
if isinstance(conn, str):
|
|
2544
|
+
# If it's a string, create engine from connection string
|
|
2545
|
+
engine = sqlalchemy.create_engine(conn)
|
|
2546
|
+
else:
|
|
2547
|
+
# Assume it's already a SQLAlchemy engine
|
|
2548
|
+
engine = conn
|
|
2549
|
+
else:
|
|
2550
|
+
# Legacy filename parameter
|
|
2551
|
+
engine = sqlalchemy.create_engine(f"sqlite:///{filename}")
|
|
2552
|
+
|
|
2553
|
+
# Prepare components data for all molecules
|
|
2554
|
+
components_data = []
|
|
2555
|
+
for mol_res in self: # type: ignore[assignment]
|
|
2556
|
+
for match in mol_res.matches:
|
|
2557
|
+
if not match.is_group:
|
|
2558
|
+
continue
|
|
2559
|
+
for comp in match.components:
|
|
2560
|
+
smarts = comp.smarts_label
|
|
2561
|
+
components_data.append({
|
|
2562
|
+
'smiles': mol_res.smiles,
|
|
2563
|
+
'group_id': match.group_id,
|
|
2564
|
+
'group_name': match.group_name,
|
|
2565
|
+
'smarts_label': str(smarts) if isinstance(smarts, list) else smarts,
|
|
2566
|
+
'component_atoms': ','.join(map(str, comp.atoms)),
|
|
2567
|
+
})
|
|
2568
|
+
|
|
2569
|
+
# Prepare groups data for all molecules
|
|
2570
|
+
groups_data = []
|
|
2571
|
+
for mol_res in self: # type: ignore[assignment]
|
|
2572
|
+
group_counts: Dict[Tuple[Optional[int], str], int] = {}
|
|
2573
|
+
for match in mol_res.matches:
|
|
2574
|
+
if not match.is_group:
|
|
2575
|
+
continue
|
|
2576
|
+
key = (match.group_id, match.group_name or '')
|
|
2577
|
+
group_counts[key] = group_counts.get(key, 0) + 1
|
|
2578
|
+
|
|
2579
|
+
for (group_id, group_name), count in group_counts.items():
|
|
2580
|
+
groups_data.append({
|
|
2581
|
+
'smiles': mol_res.smiles,
|
|
2582
|
+
'group_id': group_id,
|
|
2583
|
+
'group_name': group_name,
|
|
2584
|
+
'match_count': count,
|
|
2585
|
+
})
|
|
2586
|
+
|
|
2587
|
+
# Write to database
|
|
2588
|
+
if components_data:
|
|
2589
|
+
df_components = pd.DataFrame(components_data)
|
|
2590
|
+
df_components.to_sql(components_table, engine, if_exists=if_exists, index=False)
|
|
2591
|
+
|
|
2592
|
+
if groups_data:
|
|
2593
|
+
df_groups = pd.DataFrame(groups_data)
|
|
2594
|
+
df_groups.to_sql(groups_table, engine, if_exists=if_exists, index=False)
|
|
2595
|
+
def to_fingerprint(
|
|
2596
|
+
self,
|
|
2597
|
+
group_selection: str = 'all',
|
|
2598
|
+
component_metrics: Optional[List[str]] = None,
|
|
2599
|
+
selected_group_ids: Optional[List[int]] = None,
|
|
2600
|
+
halogens: Union[str, List[str]] = 'F',
|
|
2601
|
+
saturation: Optional[str] = 'per',
|
|
2602
|
+
molecule_metrics: Optional[List[str]] = None,
|
|
2603
|
+
pfas_groups: Optional[List[Dict]] = None,
|
|
2604
|
+
preset: Optional[str] = None,
|
|
2605
|
+
count_mode: Optional[str] = None,
|
|
2606
|
+
graph_metrics: Optional[List[str]] = None,
|
|
2607
|
+
progress: bool = False,
|
|
2608
|
+
**kwargs,
|
|
2609
|
+
) -> np.ndarray:
|
|
2610
|
+
"""Deprecated. Use :meth:`to_array` instead."""
|
|
2611
|
+
import warnings
|
|
2612
|
+
warnings.warn(
|
|
2613
|
+
"to_fingerprint() is deprecated; use to_array() instead.",
|
|
2614
|
+
DeprecationWarning,
|
|
2615
|
+
stacklevel=2,
|
|
2616
|
+
)
|
|
2617
|
+
if count_mode is not None or graph_metrics is not None:
|
|
2618
|
+
if component_metrics is None:
|
|
2619
|
+
component_metrics = [count_mode or 'binary'] + list(graph_metrics or [])
|
|
2620
|
+
return self.to_array(
|
|
2621
|
+
component_metrics=component_metrics,
|
|
2622
|
+
molecule_metrics=molecule_metrics,
|
|
2623
|
+
group_selection=group_selection,
|
|
2624
|
+
selected_group_ids=selected_group_ids,
|
|
2625
|
+
preset=preset,
|
|
2626
|
+
pfas_groups=pfas_groups,
|
|
2627
|
+
)
|
|
2628
|
+
|
|
2629
|
+
# ------------------------------------------------------------------
|
|
2630
|
+
# Convenience aliases for backward compatibility / PFASFingerprint API
|
|
2631
|
+
# ------------------------------------------------------------------
|
|
2632
|
+
|
|
2633
|
+
@property
|
|
2634
|
+
def n_molecules(self) -> int:
|
|
2635
|
+
"""Number of molecules in this set."""
|
|
2636
|
+
return len(self)
|
|
2637
|
+
|
|
2638
|
+
@property
|
|
2639
|
+
def has_cache(self) -> bool:
|
|
2640
|
+
"""Always True — PFASEmbeddingSet stores pre-parsed results."""
|
|
2641
|
+
return True
|
|
2642
|
+
|
|
2643
|
+
@property
|
|
2644
|
+
def match_cache(self) -> "PFASEmbeddingSet":
|
|
2645
|
+
"""Alias for the set itself (backward compat with PFASFingerprint API)."""
|
|
2646
|
+
return self
|
|
2647
|
+
|
|
2648
|
+
def get_embedding(self, **kwargs) -> "EmbeddingArray":
|
|
2649
|
+
"""Alias for :meth:`to_array` (backward compat with PFASFingerprint API)."""
|
|
2650
|
+
return self.to_array(**kwargs)
|
|
2651
|
+
|
|
2652
|
+
def to_array(
|
|
2653
|
+
self,
|
|
2654
|
+
component_metrics=_UNSET,
|
|
2655
|
+
molecule_metrics=_UNSET,
|
|
2656
|
+
group_selection=_UNSET,
|
|
2657
|
+
selected_group_ids=_UNSET,
|
|
2658
|
+
aggregation=_UNSET,
|
|
2659
|
+
preset=_UNSET,
|
|
2660
|
+
pfas_groups=_UNSET,
|
|
2661
|
+
halogens=_UNSET,
|
|
2662
|
+
progress: bool = True,
|
|
2663
|
+
) -> "EmbeddingArray":
|
|
2664
|
+
"""Stack per-molecule embedding rows into a ``(n_mols, n_cols)`` matrix.
|
|
2665
|
+
|
|
2666
|
+
When called with no arguments, returns the last cached embedding (or
|
|
2667
|
+
binary by default on the first call). Pass explicit arguments to
|
|
2668
|
+
override and update the cache.
|
|
2669
|
+
|
|
2670
|
+
Parameters match those of :meth:`PFASEmbedding.to_array`, plus:
|
|
2671
|
+
|
|
2672
|
+
progress : bool, default True
|
|
2673
|
+
Show a tqdm progress bar while computing embeddings.
|
|
2674
|
+
"""
|
|
2675
|
+
_no_args = (
|
|
2676
|
+
component_metrics is _UNSET and molecule_metrics is _UNSET and
|
|
2677
|
+
group_selection is _UNSET and selected_group_ids is _UNSET and
|
|
2678
|
+
aggregation is _UNSET and preset is _UNSET and pfas_groups is _UNSET
|
|
2679
|
+
and halogens is _UNSET
|
|
2680
|
+
)
|
|
2681
|
+
if _no_args and getattr(self, '_last_array', None) is not None:
|
|
2682
|
+
return self._last_array
|
|
2683
|
+
|
|
2684
|
+
# Resolve sentinels to defaults
|
|
2685
|
+
if component_metrics is _UNSET: component_metrics = None
|
|
2686
|
+
if molecule_metrics is _UNSET: molecule_metrics = None
|
|
2687
|
+
if group_selection is _UNSET: group_selection = 'all'
|
|
2688
|
+
if selected_group_ids is _UNSET: selected_group_ids = None
|
|
2689
|
+
if aggregation is _UNSET: aggregation = 'mean'
|
|
2690
|
+
if preset is _UNSET: preset = None
|
|
2691
|
+
if pfas_groups is _UNSET: pfas_groups = None
|
|
2692
|
+
if halogens is _UNSET: halogens = None
|
|
2693
|
+
|
|
2694
|
+
from .getter import get_compiled_HalogenGroups
|
|
2695
|
+
if pfas_groups is None:
|
|
2696
|
+
pfas_groups = get_compiled_HalogenGroups()
|
|
2697
|
+
|
|
2698
|
+
_iter = self
|
|
2699
|
+
if progress and len(self) > 1:
|
|
2700
|
+
try:
|
|
2701
|
+
from tqdm.auto import tqdm as _tqdm
|
|
2702
|
+
except ImportError:
|
|
2703
|
+
from tqdm import tqdm as _tqdm
|
|
2704
|
+
_iter = _tqdm(self, desc='Computing embeddings', total=len(self))
|
|
2705
|
+
|
|
2706
|
+
rows = [
|
|
2707
|
+
mol.to_array(
|
|
2708
|
+
component_metrics=component_metrics,
|
|
2709
|
+
molecule_metrics=molecule_metrics,
|
|
2710
|
+
group_selection=group_selection,
|
|
2711
|
+
selected_group_ids=selected_group_ids,
|
|
2712
|
+
aggregation=aggregation,
|
|
2713
|
+
preset=preset,
|
|
2714
|
+
pfas_groups=pfas_groups,
|
|
2715
|
+
halogens=halogens,
|
|
2716
|
+
)
|
|
2717
|
+
for mol in _iter
|
|
2718
|
+
]
|
|
2719
|
+
if not rows:
|
|
2720
|
+
mat = np.zeros((0, 0), dtype=float)
|
|
2721
|
+
else:
|
|
2722
|
+
mat = np.vstack(rows)
|
|
2723
|
+
|
|
2724
|
+
all_smiles = [m.get('smiles', '') for m in self]
|
|
2725
|
+
all_inchi = [m.get('inchi', '') for m in self]
|
|
2726
|
+
all_inchikey = [m.get('inchikey', '') for m in self]
|
|
2727
|
+
result = EmbeddingArray(mat, smiles=all_smiles, inchi=all_inchi,
|
|
2728
|
+
inchikey=all_inchikey, source=self)
|
|
2729
|
+
if _no_args:
|
|
2730
|
+
self._last_array = result
|
|
2731
|
+
return result
|
|
2732
|
+
|
|
2733
|
+
# ------------------------------------------------------------------
|
|
2734
|
+
# Analysis methods
|
|
2735
|
+
# ------------------------------------------------------------------
|
|
2736
|
+
|
|
2737
|
+
def compare_kld(
|
|
2738
|
+
self,
|
|
2739
|
+
other: "PFASEmbeddingSet",
|
|
2740
|
+
method: str = 'minmax',
|
|
2741
|
+
) -> float:
|
|
2742
|
+
"""Compare two sets using KL divergence on group-occurrence frequencies.
|
|
2743
|
+
|
|
2744
|
+
Parameters
|
|
2745
|
+
----------
|
|
2746
|
+
other : PFASEmbeddingSet
|
|
2747
|
+
Second set to compare against.
|
|
2748
|
+
method : str, default ``'minmax'``
|
|
2749
|
+
``'forward'``, ``'reverse'``, ``'symmetric'``, or ``'minmax'``
|
|
2750
|
+
(normalised symmetric KLD).
|
|
2751
|
+
|
|
2752
|
+
Returns
|
|
2753
|
+
-------
|
|
2754
|
+
float
|
|
2755
|
+
KL divergence value (lower = more similar).
|
|
2756
|
+
"""
|
|
2757
|
+
from scipy.stats import entropy as _entropy
|
|
2758
|
+
|
|
2759
|
+
p_mat = self.to_array()
|
|
2760
|
+
q_mat = other.to_array()
|
|
2761
|
+
|
|
2762
|
+
if p_mat.shape[1] != q_mat.shape[1]:
|
|
2763
|
+
raise ValueError(
|
|
2764
|
+
f"Column count mismatch: {p_mat.shape[1]} vs {q_mat.shape[1]}. "
|
|
2765
|
+
"Call to_array() with explicit arguments on both sets to ensure "
|
|
2766
|
+
"the same embedding configuration."
|
|
2767
|
+
)
|
|
2768
|
+
|
|
2769
|
+
eps = 1e-10
|
|
2770
|
+
p = np.sum(np.atleast_2d(np.asarray(p_mat)) > 0, axis=0).astype(float) + eps
|
|
2771
|
+
q = np.sum(np.atleast_2d(np.asarray(q_mat)) > 0, axis=0).astype(float) + eps
|
|
2772
|
+
p /= p.sum()
|
|
2773
|
+
q /= q.sum()
|
|
2774
|
+
|
|
2775
|
+
if method == 'forward':
|
|
2776
|
+
return float(_entropy(p, q))
|
|
2777
|
+
if method == 'reverse':
|
|
2778
|
+
return float(_entropy(q, p))
|
|
2779
|
+
if method == 'symmetric':
|
|
2780
|
+
return float((_entropy(p, q) + _entropy(q, p)) / 2)
|
|
2781
|
+
if method == 'minmax':
|
|
2782
|
+
kl_fwd = _entropy(p, q)
|
|
2783
|
+
kl_rev = _entropy(q, p)
|
|
2784
|
+
kl_sym = (kl_fwd + kl_rev) / 2
|
|
2785
|
+
max_kl = np.log(len(p))
|
|
2786
|
+
return float(kl_sym / max_kl) if max_kl > 0 else 0.0
|
|
2787
|
+
raise ValueError(f"Unknown method: {method!r}. Choose from 'forward', 'reverse', 'symmetric', 'minmax'.")
|
|
2788
|
+
|
|
2789
|
+
# ------------------------------------------------------------------
|
|
2790
|
+
# Internal helper shared by all DR methods
|
|
2791
|
+
# ------------------------------------------------------------------
|
|
2792
|
+
def _color_labels(
|
|
2793
|
+
self,
|
|
2794
|
+
color_by,
|
|
2795
|
+
):
|
|
2796
|
+
"""Return (colors, unique_labels, colour_map) for scatter plots.
|
|
2797
|
+
|
|
2798
|
+
Parameters
|
|
2799
|
+
----------
|
|
2800
|
+
color_by : None | ``'top_group'`` | list of str
|
|
2801
|
+
* ``None`` – return ``(None, None, None)``; callers fall back to a
|
|
2802
|
+
single default colour.
|
|
2803
|
+
* ``'top_group'`` – derive one label per molecule from the
|
|
2804
|
+
highest ``match_count`` ``HalogenGroup`` match.
|
|
2805
|
+
* list of str – use directly as per-molecule labels (must be the
|
|
2806
|
+
same length as ``self``).
|
|
2807
|
+
"""
|
|
2808
|
+
import matplotlib.pyplot as plt
|
|
2809
|
+
import matplotlib.patches as mpatches
|
|
2810
|
+
|
|
2811
|
+
if color_by is None:
|
|
2812
|
+
return None, None, None
|
|
2813
|
+
|
|
2814
|
+
if color_by == 'top_group':
|
|
2815
|
+
def _top(emb):
|
|
2816
|
+
hits = [m for m in emb.get('matches', [])
|
|
2817
|
+
if m.get('type') == 'HalogenGroup']
|
|
2818
|
+
if not hits:
|
|
2819
|
+
return 'No match'
|
|
2820
|
+
return max(hits, key=lambda m: m['match_count'])['group_name']
|
|
2821
|
+
labels = [_top(e) for e in self]
|
|
2822
|
+
elif isinstance(color_by, (list, tuple)):
|
|
2823
|
+
labels = list(color_by)
|
|
2824
|
+
else:
|
|
2825
|
+
raise ValueError(
|
|
2826
|
+
"color_by must be None, 'top_group', or a list of per-molecule labels."
|
|
2827
|
+
)
|
|
2828
|
+
|
|
2829
|
+
unique_labels = sorted(set(labels))
|
|
2830
|
+
cmap = plt.colormaps.get_cmap('tab20')
|
|
2831
|
+
colour_map = {
|
|
2832
|
+
lbl: cmap(i / max(len(unique_labels) - 1, 1))
|
|
2833
|
+
for i, lbl in enumerate(unique_labels)
|
|
2834
|
+
}
|
|
2835
|
+
colors = [colour_map[lbl] for lbl in labels]
|
|
2836
|
+
handles = [
|
|
2837
|
+
mpatches.Patch(color=colour_map[lbl], label=lbl)
|
|
2838
|
+
for lbl in unique_labels
|
|
2839
|
+
]
|
|
2840
|
+
return colors, handles, labels
|
|
2841
|
+
|
|
2842
|
+
def perform_pca(
|
|
2843
|
+
self,
|
|
2844
|
+
n_components: int = 2,
|
|
2845
|
+
plot: bool = True,
|
|
2846
|
+
output_file: Optional[str] = None,
|
|
2847
|
+
color_by=None,
|
|
2848
|
+
) -> Dict:
|
|
2849
|
+
"""Perform PCA on the embedding matrix.
|
|
2850
|
+
|
|
2851
|
+
Parameters
|
|
2852
|
+
----------
|
|
2853
|
+
n_components : int, default 2
|
|
2854
|
+
plot : bool, default True
|
|
2855
|
+
output_file : str, optional
|
|
2856
|
+
color_by : None | ``'top_group'`` | list of str, default None
|
|
2857
|
+
Colour scatter-plot points. Pass ``'top_group'`` to colour by the
|
|
2858
|
+
PFAS group with the highest match count per molecule, or pass a
|
|
2859
|
+
list of per-molecule label strings. ``None`` uses a single colour.
|
|
2860
|
+
|
|
2861
|
+
Returns
|
|
2862
|
+
-------
|
|
2863
|
+
dict
|
|
2864
|
+
Keys: ``'transformed'``, ``'explained_variance'``, ``'components'``,
|
|
2865
|
+
``'pca_model'``, ``'scaler'``, ``'labels'`` (if *color_by* is set).
|
|
2866
|
+
"""
|
|
2867
|
+
try:
|
|
2868
|
+
from sklearn.decomposition import PCA
|
|
2869
|
+
from sklearn.preprocessing import StandardScaler
|
|
2870
|
+
import matplotlib
|
|
2871
|
+
import matplotlib.pyplot as plt
|
|
2872
|
+
except ImportError as exc:
|
|
2873
|
+
raise ImportError(
|
|
2874
|
+
"scikit-learn and matplotlib required: pip install scikit-learn matplotlib"
|
|
2875
|
+
) from exc
|
|
2876
|
+
|
|
2877
|
+
mat = np.atleast_2d(np.asarray(self.to_array()))
|
|
2878
|
+
scaler = StandardScaler()
|
|
2879
|
+
X = scaler.fit_transform(mat)
|
|
2880
|
+
pca = PCA(n_components=n_components)
|
|
2881
|
+
X_pca = pca.fit_transform(X)
|
|
2882
|
+
|
|
2883
|
+
colors, handles, labels = self._color_labels(color_by)
|
|
2884
|
+
|
|
2885
|
+
if plot and n_components >= 2:
|
|
2886
|
+
fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(14, 5))
|
|
2887
|
+
ax1.scatter(X_pca[:, 0], X_pca[:, 1],
|
|
2888
|
+
c=colors, alpha=0.6, s=50)
|
|
2889
|
+
ax1.set_xlabel(f'PC1 ({pca.explained_variance_ratio_[0]:.1%} variance)')
|
|
2890
|
+
ax1.set_ylabel(f'PC2 ({pca.explained_variance_ratio_[1]:.1%} variance)')
|
|
2891
|
+
ax1.set_title('PCA of PFAS Group Embeddings')
|
|
2892
|
+
ax1.grid(True, alpha=0.3)
|
|
2893
|
+
if handles:
|
|
2894
|
+
ax1.legend(handles=handles, fontsize=7, loc='best',
|
|
2895
|
+
title='Group', title_fontsize=8)
|
|
2896
|
+
ax2.bar(range(1, n_components + 1), pca.explained_variance_ratio_)
|
|
2897
|
+
ax2.set_xlabel('Principal Component')
|
|
2898
|
+
ax2.set_ylabel('Explained Variance Ratio')
|
|
2899
|
+
ax2.set_title('Scree Plot')
|
|
2900
|
+
ax2.grid(True, alpha=0.3)
|
|
2901
|
+
plt.tight_layout()
|
|
2902
|
+
if output_file:
|
|
2903
|
+
plt.savefig(output_file, dpi=300, bbox_inches='tight')
|
|
2904
|
+
elif matplotlib.get_backend() != 'agg':
|
|
2905
|
+
plt.show()
|
|
2906
|
+
plt.close()
|
|
2907
|
+
|
|
2908
|
+
result = {
|
|
2909
|
+
'transformed': X_pca,
|
|
2910
|
+
'explained_variance': pca.explained_variance_ratio_,
|
|
2911
|
+
'components': pca.components_,
|
|
2912
|
+
'pca_model': pca,
|
|
2913
|
+
'scaler': scaler,
|
|
2914
|
+
}
|
|
2915
|
+
if labels is not None:
|
|
2916
|
+
result['labels'] = labels
|
|
2917
|
+
return result
|
|
2918
|
+
|
|
2919
|
+
def perform_kernel_pca(
|
|
2920
|
+
self,
|
|
2921
|
+
n_components: int = 2,
|
|
2922
|
+
kernel: str = 'rbf',
|
|
2923
|
+
gamma: Optional[float] = None,
|
|
2924
|
+
plot: bool = True,
|
|
2925
|
+
output_file: Optional[str] = None,
|
|
2926
|
+
color_by=None,
|
|
2927
|
+
) -> Dict:
|
|
2928
|
+
"""Perform kernel PCA on the embedding matrix.
|
|
2929
|
+
|
|
2930
|
+
Parameters
|
|
2931
|
+
----------
|
|
2932
|
+
n_components : int, default 2
|
|
2933
|
+
kernel : str, default ``'rbf'``
|
|
2934
|
+
gamma : float, optional
|
|
2935
|
+
plot : bool, default True
|
|
2936
|
+
output_file : str, optional
|
|
2937
|
+
color_by : None | ``'top_group'`` | list of str, default None
|
|
2938
|
+
Colour scatter-plot points. Pass ``'top_group'`` to colour by the
|
|
2939
|
+
PFAS group with the highest match count per molecule, or pass a
|
|
2940
|
+
list of per-molecule label strings. ``None`` uses a single colour.
|
|
2941
|
+
|
|
2942
|
+
Returns
|
|
2943
|
+
-------
|
|
2944
|
+
dict
|
|
2945
|
+
Keys: ``'transformed'``, ``'kpca_model'``, ``'scaler'``, ``'kernel'``,
|
|
2946
|
+
``'gamma'``, ``'labels'`` (if *color_by* is set).
|
|
2947
|
+
"""
|
|
2948
|
+
try:
|
|
2949
|
+
from sklearn.decomposition import KernelPCA
|
|
2950
|
+
from sklearn.preprocessing import StandardScaler
|
|
2951
|
+
import matplotlib
|
|
2952
|
+
import matplotlib.pyplot as plt
|
|
2953
|
+
except ImportError as exc:
|
|
2954
|
+
raise ImportError(
|
|
2955
|
+
"scikit-learn and matplotlib required: pip install scikit-learn matplotlib"
|
|
2956
|
+
) from exc
|
|
2957
|
+
|
|
2958
|
+
mat = np.atleast_2d(np.asarray(self.to_array()))
|
|
2959
|
+
scaler = StandardScaler()
|
|
2960
|
+
X = scaler.fit_transform(mat)
|
|
2961
|
+
gamma = gamma if gamma is not None else 1.0 / X.shape[1]
|
|
2962
|
+
kpca = KernelPCA(n_components=n_components, kernel=kernel, gamma=gamma)
|
|
2963
|
+
X_kpca = kpca.fit_transform(X)
|
|
2964
|
+
|
|
2965
|
+
colors, handles, labels = self._color_labels(color_by)
|
|
2966
|
+
|
|
2967
|
+
if plot and n_components >= 2:
|
|
2968
|
+
fig, ax = plt.subplots(figsize=(8, 6))
|
|
2969
|
+
ax.scatter(X_kpca[:, 0], X_kpca[:, 1],
|
|
2970
|
+
c=colors, alpha=0.6, s=50)
|
|
2971
|
+
ax.set_xlabel('Kernel PC1')
|
|
2972
|
+
ax.set_ylabel('Kernel PC2')
|
|
2973
|
+
ax.set_title(f'Kernel PCA ({kernel}) of PFAS Group Embeddings')
|
|
2974
|
+
ax.grid(True, alpha=0.3)
|
|
2975
|
+
if handles:
|
|
2976
|
+
ax.legend(handles=handles, fontsize=7, loc='best',
|
|
2977
|
+
title='Group', title_fontsize=8)
|
|
2978
|
+
plt.tight_layout()
|
|
2979
|
+
if output_file:
|
|
2980
|
+
plt.savefig(output_file, dpi=300, bbox_inches='tight')
|
|
2981
|
+
elif matplotlib.get_backend() != 'agg':
|
|
2982
|
+
plt.show()
|
|
2983
|
+
plt.close()
|
|
2984
|
+
|
|
2985
|
+
result = {
|
|
2986
|
+
'transformed': X_kpca,
|
|
2987
|
+
'kpca_model': kpca,
|
|
2988
|
+
'scaler': scaler,
|
|
2989
|
+
'kernel': kernel,
|
|
2990
|
+
'gamma': gamma,
|
|
2991
|
+
}
|
|
2992
|
+
if labels is not None:
|
|
2993
|
+
result['labels'] = labels
|
|
2994
|
+
return result
|
|
2995
|
+
|
|
2996
|
+
def perform_tsne(
|
|
2997
|
+
self,
|
|
2998
|
+
n_components: int = 2,
|
|
2999
|
+
perplexity: float = 30.0,
|
|
3000
|
+
learning_rate: float = 200.0,
|
|
3001
|
+
max_iter: int = 1000,
|
|
3002
|
+
plot: bool = True,
|
|
3003
|
+
output_file: Optional[str] = None,
|
|
3004
|
+
color_by=None,
|
|
3005
|
+
) -> Dict:
|
|
3006
|
+
"""Perform t-SNE dimensionality reduction on the embedding matrix.
|
|
3007
|
+
|
|
3008
|
+
Parameters
|
|
3009
|
+
----------
|
|
3010
|
+
n_components : int, default 2
|
|
3011
|
+
perplexity : float, default 30.0
|
|
3012
|
+
learning_rate : float, default 200.0
|
|
3013
|
+
max_iter : int, default 1000
|
|
3014
|
+
plot : bool, default True
|
|
3015
|
+
output_file : str, optional
|
|
3016
|
+
color_by : None | ``'top_group'`` | list of str, default None
|
|
3017
|
+
Colour scatter-plot points. Pass ``'top_group'`` to colour by the
|
|
3018
|
+
PFAS group with the highest match count per molecule, or pass a
|
|
3019
|
+
list of per-molecule label strings. ``None`` uses a single colour.
|
|
3020
|
+
|
|
3021
|
+
Returns
|
|
3022
|
+
-------
|
|
3023
|
+
dict
|
|
3024
|
+
Keys: ``'transformed'``, ``'tsne_model'``, ``'scaler'``, ``'perplexity'``,
|
|
3025
|
+
``'labels'`` (if *color_by* is set).
|
|
3026
|
+
"""
|
|
3027
|
+
try:
|
|
3028
|
+
from sklearn.manifold import TSNE
|
|
3029
|
+
from sklearn.preprocessing import StandardScaler
|
|
3030
|
+
import matplotlib
|
|
3031
|
+
import matplotlib.pyplot as plt
|
|
3032
|
+
except ImportError as exc:
|
|
3033
|
+
raise ImportError(
|
|
3034
|
+
"scikit-learn and matplotlib required: pip install scikit-learn matplotlib"
|
|
3035
|
+
) from exc
|
|
3036
|
+
|
|
3037
|
+
mat = np.atleast_2d(np.asarray(self.to_array()))
|
|
3038
|
+
scaler = StandardScaler()
|
|
3039
|
+
X = scaler.fit_transform(mat)
|
|
3040
|
+
tsne = TSNE(
|
|
3041
|
+
n_components=n_components,
|
|
3042
|
+
perplexity=perplexity,
|
|
3043
|
+
learning_rate=learning_rate,
|
|
3044
|
+
max_iter=max_iter,
|
|
3045
|
+
random_state=42,
|
|
3046
|
+
)
|
|
3047
|
+
X_tsne = tsne.fit_transform(X)
|
|
3048
|
+
|
|
3049
|
+
colors, handles, labels = self._color_labels(color_by)
|
|
3050
|
+
|
|
3051
|
+
if plot and n_components >= 2:
|
|
3052
|
+
fig, ax = plt.subplots(figsize=(8, 6))
|
|
3053
|
+
ax.scatter(X_tsne[:, 0], X_tsne[:, 1],
|
|
3054
|
+
c=colors, alpha=0.6, s=50)
|
|
3055
|
+
ax.set_xlabel('t-SNE 1')
|
|
3056
|
+
ax.set_ylabel('t-SNE 2')
|
|
3057
|
+
ax.set_title(f't-SNE (perplexity={perplexity}) of PFAS Group Embeddings')
|
|
3058
|
+
ax.grid(True, alpha=0.3)
|
|
3059
|
+
if handles:
|
|
3060
|
+
ax.legend(handles=handles, fontsize=7, loc='best',
|
|
3061
|
+
title='Group', title_fontsize=8)
|
|
3062
|
+
plt.tight_layout()
|
|
3063
|
+
if output_file:
|
|
3064
|
+
plt.savefig(output_file, dpi=300, bbox_inches='tight')
|
|
3065
|
+
elif matplotlib.get_backend() != 'agg':
|
|
3066
|
+
plt.show()
|
|
3067
|
+
plt.close()
|
|
3068
|
+
|
|
3069
|
+
result = {
|
|
3070
|
+
'transformed': X_tsne,
|
|
3071
|
+
'tsne_model': tsne,
|
|
3072
|
+
'scaler': scaler,
|
|
3073
|
+
'perplexity': perplexity,
|
|
3074
|
+
}
|
|
3075
|
+
if labels is not None:
|
|
3076
|
+
result['labels'] = labels
|
|
3077
|
+
return result
|
|
3078
|
+
|
|
3079
|
+
def perform_umap(
|
|
3080
|
+
self,
|
|
3081
|
+
n_components: int = 2,
|
|
3082
|
+
n_neighbors: int = 15,
|
|
3083
|
+
min_dist: float = 0.1,
|
|
3084
|
+
metric: str = 'euclidean',
|
|
3085
|
+
plot: bool = True,
|
|
3086
|
+
output_file: Optional[str] = None,
|
|
3087
|
+
color_by=None,
|
|
3088
|
+
) -> Dict:
|
|
3089
|
+
"""Perform UMAP dimensionality reduction on the embedding matrix.
|
|
3090
|
+
|
|
3091
|
+
Parameters
|
|
3092
|
+
----------
|
|
3093
|
+
n_components : int, default 2
|
|
3094
|
+
n_neighbors : int, default 15
|
|
3095
|
+
min_dist : float, default 0.1
|
|
3096
|
+
metric : str, default ``'euclidean'``
|
|
3097
|
+
plot : bool, default True
|
|
3098
|
+
output_file : str, optional
|
|
3099
|
+
color_by : None | ``'top_group'`` | list of str, default None
|
|
3100
|
+
Colour scatter-plot points. Pass ``'top_group'`` to colour by the
|
|
3101
|
+
PFAS group with the highest match count per molecule, or pass a
|
|
3102
|
+
list of per-molecule label strings. ``None`` uses a single colour.
|
|
3103
|
+
|
|
3104
|
+
Returns
|
|
3105
|
+
-------
|
|
3106
|
+
dict
|
|
3107
|
+
Keys: ``'transformed'``, ``'umap_model'``, ``'scaler'``, ``'n_neighbors'``,
|
|
3108
|
+
``'min_dist'``, ``'labels'`` (if *color_by* is set).
|
|
3109
|
+
"""
|
|
3110
|
+
try:
|
|
3111
|
+
import umap
|
|
3112
|
+
from sklearn.preprocessing import StandardScaler
|
|
3113
|
+
import matplotlib
|
|
3114
|
+
import matplotlib.pyplot as plt
|
|
3115
|
+
except ImportError as exc:
|
|
3116
|
+
raise ImportError(
|
|
3117
|
+
"umap-learn and matplotlib required: pip install umap-learn matplotlib"
|
|
3118
|
+
) from exc
|
|
3119
|
+
|
|
3120
|
+
import warnings as _warnings
|
|
3121
|
+
import os as _os
|
|
3122
|
+
_os.environ.setdefault('KMP_DUPLICATE_LIB_OK', 'TRUE')
|
|
3123
|
+
|
|
3124
|
+
mat = np.atleast_2d(np.asarray(self.to_array()))
|
|
3125
|
+
scaler = StandardScaler()
|
|
3126
|
+
X = scaler.fit_transform(mat)
|
|
3127
|
+
|
|
3128
|
+
with _warnings.catch_warnings():
|
|
3129
|
+
_warnings.filterwarnings('ignore', message=r'.*n_jobs.*overridden.*', category=UserWarning)
|
|
3130
|
+
_warnings.filterwarnings('ignore', message=r'.*Intel OpenMP.*LLVM OpenMP.*', category=RuntimeWarning)
|
|
3131
|
+
reducer = umap.UMAP(
|
|
3132
|
+
n_components=n_components,
|
|
3133
|
+
n_neighbors=n_neighbors,
|
|
3134
|
+
min_dist=min_dist,
|
|
3135
|
+
metric=metric,
|
|
3136
|
+
random_state=42,
|
|
3137
|
+
)
|
|
3138
|
+
X_umap = reducer.fit_transform(X)
|
|
3139
|
+
|
|
3140
|
+
colors, handles, labels = self._color_labels(color_by)
|
|
3141
|
+
|
|
3142
|
+
if plot and n_components >= 2:
|
|
3143
|
+
fig, ax = plt.subplots(figsize=(8, 6))
|
|
3144
|
+
ax.scatter(X_umap[:, 0], X_umap[:, 1],
|
|
3145
|
+
c=colors, alpha=0.6, s=50)
|
|
3146
|
+
ax.set_xlabel('UMAP 1')
|
|
3147
|
+
ax.set_ylabel('UMAP 2')
|
|
3148
|
+
ax.set_title(f'UMAP (n_neighbors={n_neighbors}) of PFAS Group Embeddings')
|
|
3149
|
+
ax.grid(True, alpha=0.3)
|
|
3150
|
+
if handles:
|
|
3151
|
+
ax.legend(handles=handles, fontsize=7, loc='best',
|
|
3152
|
+
title='Group', title_fontsize=8)
|
|
3153
|
+
plt.tight_layout()
|
|
3154
|
+
if output_file:
|
|
3155
|
+
plt.savefig(output_file, dpi=300, bbox_inches='tight')
|
|
3156
|
+
elif matplotlib.get_backend() != 'agg':
|
|
3157
|
+
plt.show()
|
|
3158
|
+
plt.close()
|
|
3159
|
+
|
|
3160
|
+
result = {
|
|
3161
|
+
'transformed': X_umap,
|
|
3162
|
+
'umap_model': reducer,
|
|
3163
|
+
'scaler': scaler,
|
|
3164
|
+
'n_neighbors': n_neighbors,
|
|
3165
|
+
'min_dist': min_dist,
|
|
3166
|
+
}
|
|
3167
|
+
if labels is not None:
|
|
3168
|
+
result['labels'] = labels
|
|
3169
|
+
return result
|
|
3170
|
+
|
|
3171
|
+
def column_names(
|
|
3172
|
+
self,
|
|
3173
|
+
component_metrics: Optional[List[str]] = None,
|
|
3174
|
+
molecule_metrics: Optional[List[str]] = None,
|
|
3175
|
+
group_selection: str = 'all',
|
|
3176
|
+
selected_group_ids: Optional[List[int]] = None,
|
|
3177
|
+
preset: Optional[str] = None,
|
|
3178
|
+
pfas_groups=None,
|
|
3179
|
+
halogens=None,
|
|
3180
|
+
) -> List[str]:
|
|
3181
|
+
"""Return column labels (delegates to first element)."""
|
|
3182
|
+
if self:
|
|
3183
|
+
return self[0].column_names(
|
|
3184
|
+
component_metrics=component_metrics,
|
|
3185
|
+
molecule_metrics=molecule_metrics,
|
|
3186
|
+
group_selection=group_selection,
|
|
3187
|
+
selected_group_ids=selected_group_ids,
|
|
3188
|
+
preset=preset,
|
|
3189
|
+
pfas_groups=pfas_groups,
|
|
3190
|
+
halogens=halogens,
|
|
3191
|
+
)
|
|
3192
|
+
return []
|
|
3193
|
+
|
|
3194
|
+
@classmethod
|
|
3195
|
+
def from_sql(
|
|
3196
|
+
cls,
|
|
3197
|
+
conn: Optional[Union[str, 'sqlalchemy.engine.Engine']] = None,
|
|
3198
|
+
filename: Optional[str] = None,
|
|
3199
|
+
components_table: str = "components",
|
|
3200
|
+
groups_table: str = "pfas_groups_in_compound",
|
|
3201
|
+
limit: Optional[int] = None,
|
|
3202
|
+
) -> "PFASEmbeddingSet":
|
|
3203
|
+
"""Load results from SQL database.
|
|
3204
|
+
|
|
3205
|
+
Parameters
|
|
3206
|
+
----------
|
|
3207
|
+
conn : str or SQLAlchemy Engine, optional
|
|
3208
|
+
Database connection string or engine
|
|
3209
|
+
filename : str, optional
|
|
3210
|
+
SQLite database filename (alternative to conn)
|
|
3211
|
+
components_table : str, default "components"
|
|
3212
|
+
Name of the components table
|
|
3213
|
+
groups_table : str, default "pfas_groups_in_compound"
|
|
3214
|
+
Name of the groups table
|
|
3215
|
+
limit : int, optional
|
|
3216
|
+
Limit number of molecules to load
|
|
3217
|
+
|
|
3218
|
+
Returns
|
|
3219
|
+
-------
|
|
3220
|
+
ResultsModel
|
|
3221
|
+
Loaded results
|
|
3222
|
+
"""
|
|
3223
|
+
try:
|
|
3224
|
+
import sqlalchemy
|
|
3225
|
+
except ImportError as exc:
|
|
3226
|
+
raise ImportError("sqlalchemy is required for SQL operations. Install with: pip install sqlalchemy") from exc
|
|
3227
|
+
# Create engine
|
|
3228
|
+
if conn is not None:
|
|
3229
|
+
if isinstance(conn, str):
|
|
3230
|
+
engine = sqlalchemy.create_engine(conn)
|
|
3231
|
+
else:
|
|
3232
|
+
engine = conn
|
|
3233
|
+
elif filename is not None:
|
|
3234
|
+
engine = sqlalchemy.create_engine(f"sqlite:///{filename}")
|
|
3235
|
+
else:
|
|
3236
|
+
raise ValueError("Either conn or filename must be provided")
|
|
3237
|
+
|
|
3238
|
+
# Load groups data
|
|
3239
|
+
query = f"SELECT * FROM {groups_table}"
|
|
3240
|
+
if limit is not None:
|
|
3241
|
+
query += f" LIMIT {limit}"
|
|
3242
|
+
|
|
3243
|
+
df_groups = pd.read_sql(query, engine)
|
|
3244
|
+
|
|
3245
|
+
# Reconstruct results
|
|
3246
|
+
results = []
|
|
3247
|
+
for smiles in df_groups['smiles'].unique():
|
|
3248
|
+
mol_groups = df_groups[df_groups['smiles'] == smiles]
|
|
3249
|
+
|
|
3250
|
+
matches = []
|
|
3251
|
+
for _, row in mol_groups.iterrows():
|
|
3252
|
+
matches.append({
|
|
3253
|
+
'type': 'HalogenGroup',
|
|
3254
|
+
'match_id': row['group_id'],
|
|
3255
|
+
'group_id': row['group_id'],
|
|
3256
|
+
'group_name': row['group_name'],
|
|
3257
|
+
'match_count': row['match_count'],
|
|
3258
|
+
'components': [], # Components not stored in basic SQL format
|
|
3259
|
+
})
|
|
3260
|
+
|
|
3261
|
+
results.append({
|
|
3262
|
+
'smiles': smiles,
|
|
3263
|
+
'matches': matches,
|
|
3264
|
+
})
|
|
3265
|
+
|
|
3266
|
+
return cls(results)
|
|
3267
|
+
|
|
3268
|
+
|
|
3269
|
+
# Backward-compatible aliases
|
|
3270
|
+
MoleculeResult = PFASEmbedding
|
|
3271
|
+
ResultsModel = PFASEmbeddingSet
|
|
3272
|
+
|
|
3273
|
+
|
|
3274
|
+
def generate_fingerprint(
|
|
3275
|
+
smiles,
|
|
3276
|
+
*,
|
|
3277
|
+
selected_groups=None,
|
|
3278
|
+
representation='vector',
|
|
3279
|
+
component_metrics=None,
|
|
3280
|
+
halogens='F',
|
|
3281
|
+
saturation='per',
|
|
3282
|
+
count_mode=None,
|
|
3283
|
+
**kwargs,
|
|
3284
|
+
):
|
|
3285
|
+
"""Generate a fingerprint vector for one or more SMILES strings.
|
|
3286
|
+
|
|
3287
|
+
Parameters
|
|
3288
|
+
----------
|
|
3289
|
+
smiles : str or list of str
|
|
3290
|
+
Input SMILES.
|
|
3291
|
+
halogens : str or list of str, default 'F'
|
|
3292
|
+
Halogens to include in the fingerprint.
|
|
3293
|
+
saturation : str or None, default 'per'
|
|
3294
|
+
Saturation filter.
|
|
3295
|
+
count_mode : str or None
|
|
3296
|
+
Fingerprint mode, e.g. 'binary' or 'count'. When provided, overrides
|
|
3297
|
+
*component_metrics*.
|
|
3298
|
+
component_metrics : list of str or None
|
|
3299
|
+
Explicit list of per-component metrics.
|
|
3300
|
+
|
|
3301
|
+
Returns
|
|
3302
|
+
-------
|
|
3303
|
+
tuple (array, column_names)
|
|
3304
|
+
array : EmbeddingArray — 1-D for a single SMILES, 2-D for a list.
|
|
3305
|
+
column_names : list of str.
|
|
3306
|
+
"""
|
|
3307
|
+
from .parser import parse_smiles as _parse_smiles
|
|
3308
|
+
resolved_cm = ([count_mode] if count_mode is not None else component_metrics) or None
|
|
3309
|
+
result = _parse_smiles(smiles, halogens=halogens, saturation=saturation, **kwargs)
|
|
3310
|
+
arr = result.to_array(component_metrics=resolved_cm)
|
|
3311
|
+
cols = result.column_names(component_metrics=resolved_cm)
|
|
3312
|
+
if isinstance(smiles, str):
|
|
3313
|
+
arr = arr[0] # squeeze (1, n_cols) → (n_cols,) for single SMILES
|
|
3314
|
+
return arr, cols
|
|
3315
|
+
|