PFASGroups 3.2.2__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,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
+