OSCR

Drug screening for α-synuclein aggregation inhibitors via multimodal graph neural network.

Code ↔ Paper

6 matches between paragraphs of the paper and lines of its authors' code, computed by the harvester (lexical-v1). Click a colored paragraph or line to see its counterpart.

The 6 matches
  1. [1] § Results and discussion › Comparative analysis of diverse machine learning models ↔ pycaret.py, lines 243–302 · score 0.91 · gradient boosting, CatBoost, LightGBM, XGBoost, Extra Trees, linear regression
  2. [2] § Results and discussion › Interpretability analysis of drug molecules ↔ gnnexplainer.py, lines 115–232 · score 0.66 · pyridine rings, aromatic rings, Amide, methylene, amino, hydroxyl
  3. [3] § Results and discussion › Interpretability analysis of drug molecules ↔ gnnexplainer.py, lines 115–232 · score 0.63 · Aromatic rings, vinyl, Ether, amides, methylene, amino
  4. [4] § Results and discussion › Superiority of GraphSAGE over baseline Graph Neural Network (GNN) models ↔ statistical analysis.py, lines 135–232 · score 0.61 · odds ratio, log scaled, molecular descriptors odds, Forest
  5. [5] § Results and discussion › Distribution and differential analysis of molecular descriptors for high-/low-affinity ligands ↔ statistical analysis.py, lines 48–74 · score 0.59 · pChEMBL, aromatic rings, median, donors, affinity, bond
  6. [6] § Results and discussion › Interpretability analysis of drug molecules ↔ gnnexplainer.py, lines 1468–1512 · score 0.51 · stratified sampling, representative molecules, GNNExplainer, atom, predictive

Paper

Loaded from Europe PMC by your browser, not stored by OSCR: doi.org · Europe PMC

The paper is loaded when this pane is shown.

The authors' code

Python · 1,896 lines · 79 KB · no license · 3 matches

  1. import numpy as np
  2. import pandas as pd
  3. import torch
  4. import torch.nn as nn
  5. import torch.nn.functional as F
  6. from torch_geometric.data import Data, DataLoader
  7. from torch_geometric.explain import Explainer, GNNExplainer
  8. from torch_geometric.explain.config import ExplainerConfig, ModelConfig
  9. from torch_geometric.nn import GATConv, SAGEConv, global_max_pool
  10. import matplotlib.pyplot as plt
  11. import seaborn as sns
  12. import networkx as nx
  13. from rdkit import Chem
  14. from rdkit.Chem import Draw, AllChem, Descriptors, rdMolDescriptors
  15. from rdkit.Chem.Draw import rdDepictor
  16. import io
  17. import base64
  18. from PIL import Image
  19. from collections import defaultdict, Counter
  20. import warnings
  21. plt.rcParams['font.sans-serif'] = ['DejaVu Sans', 'Arial', 'sans-serif']
  22. plt.rcParams['axes.unicode_minus'] = False
  23. plt.rcParams['font.size'] = 14
  24. plt.rcParams['axes.titlesize'] = 18
  25. plt.rcParams['axes.labelsize'] = 16
  26. plt.rcParams['xtick.labelsize'] = 14
  27. plt.rcParams['ytick.labelsize'] = 14
  28. plt.rcParams['legend.fontsize'] = 14
  29. try:
  30. from rdkit.Chem.Draw import rdMolDraw2D
  31. except ImportError:
  32. try:
  33. from rdkit.Chem import rdMolDraw2D
  34. except ImportError:
  35. rdMolDraw2D = None
  36. print("Warning: rdMolDraw2D not available")
  37. warnings.filterwarnings('ignore')
  38. seed = 42
  39. torch.manual_seed(seed)
  40. np.random.seed(seed)
  41. def one_of_k_encoding_unk(x, valid_entries):
  42. if x not in valid_entries:
  43. x = 'Unknown'
  44. return [1 if entry == x else 0 for entry in valid_entries]
  45. class ModifiedGATLayer(nn.Module):
  46. def __init__(self, in_features, out_features):
  47. super(ModifiedGATLayer, self).__init__()
  48. self.query_transform = nn.Linear(in_features, out_features)
  49. self.key_transform = nn.Linear(in_features, out_features)
  50. self.value_transform = nn.Linear(in_features, out_features)
  51. self.conv3 = nn.Conv1d(in_channels=out_features, out_channels=out_features, kernel_size=3, padding=1)
  52. self.conv5 = nn.Conv1d(in_channels=out_features, out_channels=out_features, kernel_size=5, padding=2)
  53. self.linear_transform = nn.Linear(out_features * 2 + out_features, out_features)
  54. def forward(self, x):
  55. Q = self.query_transform(x)
  56. K = self.key_transform(x)
  57. V = self.value_transform(x)
  58. K = K.unsqueeze(2)
  59. K_conv3 = self.conv3(K)
  60. K_conv5 = self.conv5(K)
  61. K_concat = torch.cat((K_conv3, K_conv5, K), dim=1)
  62. K_new = self.linear_transform(K_concat.transpose(1, 2))
  63. attention_scores = torch.matmul(Q, K_new.transpose(1, 2)) / (K_new.size(-1) ** 0.5)
  64. attention_weights = F.softmax(attention_scores.squeeze(-1), dim=-1)
  65. output = torch.matmul(attention_weights, V) + V
  66. return output
  67. class GAT_GraphSAGE(nn.Module):
  68. def __init__(self, n_output=1, num_features_xd=35, output_dim=128, dropout=0.3):
  69. super(GAT_GraphSAGE, self).__init__()
  70. self.conv1 = ModifiedGATLayer(in_features=num_features_xd, out_features=num_features_xd)
  71. self.conv2 = SAGEConv(num_features_xd, num_features_xd)
  72. self.fc_g1 = nn.Linear(num_features_xd, 1500)
  73. self.fc_g2 = nn.Linear(1500, output_dim)
  74. self.relu = nn.ReLU()
  75. self.dropout = nn.Dropout(dropout)
  76. self.out = nn.Linear(output_dim, n_output)
  77. def forward(self, data):
  78. x, edge_index, batch = data.x, data.edge_index, data.batch
  79. x = self.conv1(x)
  80. x = self.relu(x)
  81. x = self.conv2(x, edge_index)
  82. x = self.relu(x)
  83. x = global_max_pool(x, batch)
  84. x = self.relu(self.fc_g1(x))
  85. x = self.dropout(x)
  86. x = self.fc_g2(x)
  87. out = self.out(x)
  88. return out
  89. class ExplainableGATGraphSAGE(nn.Module):
  90. def __init__(self, gat_graphsage_model):
  91. super(ExplainableGATGraphSAGE, self).__init__()
  92. self.gat_graphsage = gat_graphsage_model
  93. def forward(self, x, edge_index, batch=None, edge_attr=None):
  94. if batch is None:
  95. batch = torch.zeros(x.size(0), dtype=torch.long, device=x.device)
  96. data = Data(x=x, edge_index=edge_index, batch=batch)
  97. return self.gat_graphsage(data)
  98. class SubstructureIdentifier:
  99. def __init__(self):
  100. self.common_substructures = {
  101. 'hydroxyl': 'O',
  102. 'amino': 'N',
  103. 'carboxyl': 'C(=O)O',
  104. 'carbonyl': 'C=O',
  105. 'ester': 'C(=O)O[C,c]',
  106. 'amide': 'C(=O)N',
  107. 'ether': '[C,c]O[C,c]',
  108. 'nitro': 'N(=O)=O',
  109. 'sulfonyl': 'S(=O)(=O)',
  110. 'phosphate': 'P(=O)',
  111. 'benzene': 'c1ccccc1',
  112. 'pyridine': 'c1ccncc1',
  113. 'pyrimidine': 'c1cncnc1',
  114. 'imidazole': 'c1c[nH]cn1',
  115. 'thiophene': 'c1ccsc1',
  116. 'furan': 'c1ccoc1',
  117. 'indole': 'c1ccc2[nH]ccc2c1',
  118. 'quinoline': 'c1ccc2ncccc2c1',
  119. 'piperidine': 'C1CCNCC1',
  120. 'piperazine': 'C1CNCCN1',
  121. 'morpholine': 'C1COCCN1',
  122. 'pyrrolidine': 'C1CCNC1',
  123. 'tetrahydrofuran': 'C1CCOC1',
  124. 'methylene': 'CC',
  125. 'ethylene': 'CCC',
  126. 'propylene': 'CCCC',
  127. 'vinyl': 'C=C',
  128. 'acetylene': 'C#C',
  129. }
  130. def identify_substructures_in_molecule(self, mol):
  131. if mol is None:
  132. return {}
  133. found_substructures = {}
  134. for name, smarts in self.common_substructures.items():
  135. try:
  136. matches = mol.GetSubstructMatches(Chem.MolFromSmarts(smarts))
  137. if matches:
  138. found_substructures[name] = {
  139. 'pattern': smarts,
  140. 'matches': matches,
  141. 'count': len(matches)
  142. }
  143. except:
  144. continue
  145. return found_substructures
  146. def extract_important_substructures(self, mol, important_atom_indices, radius=2):
  147. if mol is None:
  148. return []
  149. important_substructures = []
  150. for atom_idx in important_atom_indices:
  151. if atom_idx >= mol.GetNumAtoms():
  152. continue
  153. env = Chem.FindAtomEnvironmentOfRadiusN(mol, radius, atom_idx)
  154. submol = Chem.PathToSubmol(mol, env)
  155. if submol:
  156. try:
  157. sub_smiles = Chem.MolToSmiles(submol)
  158. if sub_smiles and len(sub_smiles) > 1:
  159. important_substructures.append({
  160. 'center_atom': atom_idx,
  161. 'smiles': sub_smiles,
  162. 'num_atoms': submol.GetNumAtoms(),
  163. 'radius': radius
  164. })
  165. except:
  166. continue
  167. return important_substructures
  168. def get_functional_groups(self, mol):
  169. if mol is None:
  170. return []
  171. functional_groups = []
  172. try:
  173. from rdkit.Chem import Descriptors, Fragments
  174. groups_to_check = {
  175. 'Aromatic_Rings': Descriptors.NumAromaticRings(mol),
  176. 'Aliphatic_Rings': Descriptors.NumAliphaticRings(mol),
  177. 'Hydroxyl_Groups': Fragments.fr_Al_OH(mol),
  178. 'Carboxyl_Groups': Fragments.fr_COO(mol),
  179. 'Amino_Groups': Fragments.fr_NH2(mol),
  180. 'Ester_Groups': Fragments.fr_ester(mol),
  181. 'Ether_Groups': Fragments.fr_ether(mol),
  182. 'Amide_Groups': Fragments.fr_amide(mol),
  183. 'Nitro_Groups': Fragments.fr_nitro(mol),
  184. 'Benzene_Rings': Fragments.fr_benzene(mol),
  185. 'Pyridine_Rings': Fragments.fr_pyridine(mol),
  186. }
  187. for group_name, count in groups_to_check.items():
  188. if count > 0:
  189. functional_groups.append({
  190. 'name': group_name,
  191. 'count': count
  192. })
  193. except Exception as e:
  194. print(f"Error in functional group analysis: {e}")
  195. return functional_groups
  196. class SubstructureVisualizer:
  197. def __init__(self):
  198. plt.rcParams['font.sans-serif'] = ['DejaVu Sans', 'Arial', 'sans-serif']
  199. plt.rcParams['axes.unicode_minus'] = False
  200. def visualize_substructure_importance_summary(self, important_substructures, save_path=None):
  201. all_substructure_types = {}
  202. all_functional_groups = Counter()
  203. for struct in important_substructures:
  204. for sub_name, sub_info in struct['known_substructures'].items():
  205. if sub_name not in all_substructure_types:
  206. all_substructure_types[sub_name] = {
  207. 'count': 0,
  208. 'total_importance': 0,
  209. 'molecules': []
  210. }
  211. all_substructure_types[sub_name]['count'] += len(sub_info['matches'])
  212. all_substructure_types[sub_name]['total_importance'] += sum(sub_info['importance_scores'])
  213. all_substructure_types[sub_name]['molecules'].append(struct['smiles'][:30])
  214. for fg in struct['functional_groups']:
  215. all_functional_groups[fg['name']] += fg['count']
  216. if not all_substructure_types:
  217. print("No substructure data available for visualization")
  218. return
  219. sorted_substructures = sorted(all_substructure_types.items(),
  220. key=lambda x: x[1]['count'], reverse=True)[:15]
  221. sub_names = [item[0] for item in sorted_substructures]
  222. sub_counts = [item[1]['count'] for item in sorted_substructures]
  223. plt.figure(figsize=(14, 10))
  224. bars1 = plt.barh(sub_names, sub_counts, color='skyblue', alpha=0.8)
  225. plt.title('Substructure Frequency (Top 15)', fontsize=20, fontweight='bold', pad=25)
  226. plt.xlabel('Frequency Count', fontsize=18, fontweight='bold')
  227. plt.ylabel('Substructure Type', fontsize=18, fontweight='bold')
  228. plt.xticks(fontsize=16)
  229. plt.yticks(fontsize=16)
  230. plt.grid(axis='x', alpha=0.3)
  231. for bar, count in zip(bars1, sub_counts):
  232. plt.text(bar.get_width() + max(sub_counts) * 0.01, bar.get_y() + bar.get_height() / 2,
  233. str(count), ha='left', va='center', fontsize=14, fontweight='bold')
  234. plt.subplots_adjust(left=0.25, right=0.95, top=0.9, bottom=0.15)
  235. if save_path:
  236. plt.savefig(f"{save_path}_subplot1_frequency.png", dpi=300, bbox_inches='tight')
  237. print(f"Substructure subplot 1 saved: {save_path}_subplot1_frequency.png")
  238. plt.show()
  239. avg_importance = []
  240. for item in sorted_substructures:
  241. count = item[1]['count']
  242. total_imp = item[1]['total_importance']
  243. avg_imp = total_imp / count if count > 0 else 0
  244. avg_importance.append(avg_imp)
  245. plt.figure(figsize=(14, 10))
  246. bars2 = plt.barh(sub_names, avg_importance, color='lightcoral', alpha=0.8)
  247. plt.title('Average Substructure Importance (Top 15)', fontsize=20, fontweight='bold', pad=25)
  248. plt.xlabel('Average Importance Score', fontsize=18, fontweight='bold')
  249. plt.ylabel('Substructure Type', fontsize=18, fontweight='bold')
  250. plt.xticks(fontsize=16)
  251. plt.yticks(fontsize=16)
  252. plt.grid(axis='x', alpha=0.3)
  253. for bar, imp in zip(bars2, avg_importance):
  254. plt.text(bar.get_width() + max(avg_importance) * 0.01, bar.get_y() + bar.get_height() / 2,
  255. f'{imp:.3f}', ha='left', va='center', fontsize=14, fontweight='bold')
  256. plt.subplots_adjust(left=0.25, right=0.95, top=0.9, bottom=0.15)
  257. if save_path:
  258. plt.savefig(f"{save_path}_subplot2_importance.png", dpi=300, bbox_inches='tight')
  259. print(f"Substructure subplot 2 saved: {save_path}_subplot2_importance.png")
  260. plt.show()
  261. top_fg = all_functional_groups.most_common(12)
  262. fg_names = [item[0].replace('_', ' ') for item in top_fg]
  263. fg_counts = [item[1] for item in top_fg]
  264. plt.figure(figsize=(14, 12))
  265. colors = plt.cm.Set3(np.linspace(0, 1, len(fg_names)))
  266. def autopct_func_fg(pct, idx=0):
  267. if idx < len(fg_counts) - 4:
  268. return f'{pct:.1f}%'
  269. else:
  270. return ''
  271. counter = [0]
  272. def autopct_with_counter(pct):
  273. result = autopct_func_fg(pct, counter[0])
  274. counter[0] += 1
  275. return result
  276. wedges, texts, autotexts = plt.pie(fg_counts, autopct=autopct_with_counter,
  277. colors=colors, startangle=90,
  278. textprops={'fontsize': 14, 'fontweight': 'bold'})
  279. plt.title('Functional Group Distribution', fontsize=20, fontweight='bold', pad=25)
  280. for i, autotext in enumerate(autotexts):
  281. autotext.set_fontsize(14)
  282. autotext.set_color('black')
  283. autotext.set_fontweight('bold')
  284. plt.legend(wedges, fg_names, title="Functional Groups",
  285. title_fontsize=16, fontsize=14,
  286. loc="center left", bbox_to_anchor=(1.1, 0.5))
  287. plt.subplots_adjust(left=0.1, right=0.75, top=0.9, bottom=0.1)
  288. if save_path:
  289. plt.savefig(f"{save_path}_subplot3_functional_groups.png", dpi=300, bbox_inches='tight')
  290. print(f"Substructure subplot 3 saved: {save_path}_subplot3_functional_groups.png")
  291. plt.show()
  292. plt.figure(figsize=(14, 10))
  293. x_counts = [item[1]['count'] for item in sorted_substructures[:20]]
  294. y_importance = []
  295. colors_scatter = []
  296. labels = []
  297. for item in sorted_substructures[:20]:
  298. count = item[1]['count']
  299. total_imp = item[1]['total_importance']
  300. avg_imp = total_imp / count if count > 0 else 0
  301. y_importance.append(avg_imp)
  302. if avg_imp > 0.6:
  303. colors_scatter.append('red')
  304. elif avg_imp > 0.4:
  305. colors_scatter.append('orange')
  306. else:
  307. colors_scatter.append('blue')
  308. labels.append(item[0])
  309. scatter = plt.scatter(x_counts, y_importance, c=colors_scatter,
  310. s=150, alpha=0.7, edgecolors='black')
  311. for i, label in enumerate(labels):
  312. if y_importance[i] > 0.5 or x_counts[i] > 200:
  313. plt.annotate(label, (x_counts[i], y_importance[i]),
  314. xytext=(5, 5), textcoords='offset points',
  315. fontsize=12, alpha=0.8, fontweight='bold')
  316. plt.title('Substructure Importance vs Frequency', fontsize=20, fontweight='bold', pad=25)
  317. plt.xlabel('Frequency Count', fontsize=18, fontweight='bold')
  318. plt.ylabel('Average Importance', fontsize=18, fontweight='bold')
  319. plt.xticks(fontsize=16)
  320. plt.yticks(fontsize=16)
  321. plt.grid(True, alpha=0.3)
  322. from matplotlib.patches import Patch
  323. legend_elements = [Patch(facecolor='red', alpha=0.7, label='High Importance (>0.6)'),
  324. Patch(facecolor='orange', alpha=0.7, label='Medium Importance (0.4-0.6)'),
  325. Patch(facecolor='blue', alpha=0.7, label='Low Importance (<0.4)')]
  326. plt.legend(handles=legend_elements, loc='upper right', fontsize=14)
  327. plt.subplots_adjust(left=0.15, right=0.95, top=0.9, bottom=0.15)
  328. if save_path:
  329. plt.savefig(f"{save_path}_subplot4_scatter.png", dpi=300, bbox_inches='tight')
  330. print(f"Substructure subplot 4 saved: {save_path}_subplot4_scatter.png")
  331. plt.show()
  332. print(f"\nAll four substructure subplots have been generated and saved separately!")
  333. def visualize_substructure_in_molecules(self, full_dataset_substructures, num_examples=6, save_path=None):
  334. if not full_dataset_substructures:
  335. print("No important substructures available for visualization")
  336. return
  337. print(f"Searching for qualified molecules from {len(full_dataset_substructures)} molecules in full dataset...")
  338. selected_molecules = []
  339. for struct in full_dataset_substructures:
  340. if 'target_value' not in struct or struct['target_value'] <= 6:
  341. continue
  342. has_high_importance_substructure = False
  343. for sub_name, sub_info in struct['known_substructures'].items():
  344. if sub_info['importance_scores'] and max(sub_info['importance_scores']) > 0.5:
  345. has_high_importance_substructure = True
  346. break
  347. if has_high_importance_substructure:
  348. selected_molecules.append(struct)
  349. if len(selected_molecules) >= num_examples:
  350. break
  351. print(f"Found {len(selected_molecules)} molecules with y > 6 and high importance substructures (> 0.5)")
  352. if not selected_molecules:
  353. print("No molecules found with y > 6 and high importance substructures (> 0.5)")
  354. print("Falling back to molecules with any high importance substructures...")
  355. for struct in full_dataset_substructures:
  356. if len(selected_molecules) >= num_examples:
  357. break
  358. has_high_importance_substructure = False
  359. for sub_name, sub_info in struct['known_substructures'].items():
  360. if sub_info['importance_scores'] and max(sub_info['importance_scores']) > 0.5:
  361. has_high_importance_substructure = True
  362. break
  363. if has_high_importance_substructure:
  364. selected_molecules.append(struct)
  365. if not selected_molecules:
  366. print("No molecules with high importance substructures found")
  367. selected_molecules = full_dataset_substructures[:num_examples]
  368. fig, axes = plt.subplots(2, 3, figsize=(20, 14))
  369. axes = axes.flatten()
  370. for i, struct in enumerate(selected_molecules[:num_examples]):
  371. try:
  372. mol = Chem.MolFromSmiles(struct['smiles'])
  373. if mol is None:
  374. continue
  375. best_substructure = None
  376. best_importance = 0
  377. for sub_name, sub_info in struct['known_substructures'].items():
  378. if sub_info['importance_scores']:
  379. avg_importance = np.mean(sub_info['importance_scores'])
  380. if avg_importance > best_importance:
  381. best_importance = avg_importance
  382. best_substructure = (sub_name, sub_info)
  383. if best_substructure is None:
  384. img = Draw.MolToImage(mol, size=(400, 400))
  385. axes[i].imshow(img)
  386. y_value = struct.get('target_value', 'N/A')
  387. axes[i].set_title(f'Molecule {i + 1}\nTrue y: {y_value}' +
  388. (f' (y={y_value:.3f})' if isinstance(y_value, (int, float)) else ''),
  389. fontsize=16, fontweight='bold', pad=10)
  390. else:
  391. sub_name, sub_info = best_substructure
  392. highlight_atoms = []
  393. for match in sub_info['matches']:
  394. highlight_atoms.extend(match)
  395. highlight_atoms = list(set(highlight_atoms))
  396. img = Draw.MolToImage(mol, size=(400, 400),
  397. highlightAtoms=highlight_atoms,
  398. highlightColors={atom: (1, 0.8, 0.8) for atom in highlight_atoms})
  399. axes[i].imshow(img)
  400. y_value = struct.get('target_value', 'N/A')
  401. title_text = f'Molecule {i + 1}: {sub_name}\n'
  402. if isinstance(y_value, (int, float)):
  403. title_text += f'True y: {y_value:.3f}\n'
  404. else:
  405. title_text += f'True y: {y_value}\n'
  406. title_text += f'Importance: {best_importance:.3f}'
  407. axes[i].set_title(title_text, fontsize=16, fontweight='bold', pad=10)
  408. axes[i].axis('off')
  409. except Exception as e:
  410. print(f"Error drawing molecule {i + 1}: {e}")
  411. axes[i].text(0.5, 0.5, f'Molecule {i + 1}\nDrawing Failed',
  412. ha='center', va='center', transform=axes[i].transAxes,
  413. fontsize=16, fontweight='bold')
  414. axes[i].axis('off')
  415. for i in range(len(selected_molecules), len(axes)):
  416. axes[i].axis('off')
  417. plt.suptitle('Important Substructures Highlighted in Molecules (Full Dataset: y > 6, Importance > 0.5)',
  418. fontsize=20, fontweight='bold', y=0.95)
  419. plt.subplots_adjust(top=0.85, bottom=0.05, left=0.05, right=0.95, hspace=0.3, wspace=0.2)
  420. if save_path:
  421. plt.savefig(save_path, dpi=300, bbox_inches='tight')
  422. print(f"Highlighted molecules plot saved: {save_path}")
  423. plt.show()
  424. return selected_molecules
  425. def create_substructure_heatmap(self, important_substructures, save_path=None):
  426. if not important_substructures:
  427. print("No data available for heatmap")
  428. return
  429. all_substructures = set()
  430. for struct in important_substructures:
  431. all_substructures.update(struct['known_substructures'].keys())
  432. all_substructures = sorted(list(all_substructures))
  433. if len(all_substructures) == 0:
  434. print("No substructures found")
  435. return
  436. matrix_data = []
  437. molecule_labels = []
  438. for i, struct in enumerate(important_substructures[:40]):
  439. row = []
  440. molecule_labels.append(f"Mol{i + 1}")
  441. for sub_name in all_substructures:
  442. if sub_name in struct['known_substructures']:
  443. importance_scores = struct['known_substructures'][sub_name]['importance_scores']
  444. avg_importance = np.mean(importance_scores) if importance_scores else 0
  445. row.append(avg_importance)
  446. else:
  447. row.append(0)
  448. matrix_data.append(row)
  449. df_heatmap = pd.DataFrame(matrix_data,
  450. columns=all_substructures,
  451. index=molecule_labels)
  452. df_heatmap = df_heatmap.loc[:, (df_heatmap != 0).any(axis=0)]
  453. if df_heatmap.empty:
  454. print("Insufficient data for heatmap")
  455. return
  456. fig_width = max(18, len(df_heatmap.columns) * 1.2)
  457. fig_height = max(14, len(df_heatmap.index) * 0.45)
  458. plt.figure(figsize=(fig_width, fig_height))
  459. try:
  460. ax = sns.heatmap(df_heatmap,
  461. cmap='Blues',
  462. annot=False,
  463. fmt='.2f',
  464. cbar_kws={'label': 'Importance Score', 'shrink': 0.8},
  465. xticklabels=True,
  466. yticklabels=True,
  467. square=False)
  468. plt.title('Molecule-Substructure Importance Heatmap (Top 40 Molecules)',
  469. fontsize=22, fontweight='bold', pad=30)
  470. plt.xlabel('Substructure Type', fontsize=18, fontweight='bold', labelpad=15)
  471. plt.ylabel('Molecules (Top 40)', fontsize=18, fontweight='bold', labelpad=15)
  472. plt.xticks(rotation=45, ha='right', fontsize=14, fontweight='normal')
  473. plt.yticks(rotation=0, fontsize=14, fontweight='normal')
  474. cbar = ax.collections[0].colorbar
  475. cbar.ax.tick_params(labelsize=16)
  476. cbar.set_label('Importance Score', fontsize=18, fontweight='bold', labelpad=20)
  477. plt.subplots_adjust(left=0.15, bottom=0.25, right=0.85, top=0.85)
  478. except Exception as e:
  479. print(f"Error creating heatmap: {e}")
  480. if save_path:
  481. plt.savefig(save_path, dpi=300, bbox_inches='tight',
  482. facecolor='white', edgecolor='none')
  483. print(f"Substructure heatmap (Top 40 molecules) saved: {save_path}")
  484. plt.show()
  485. class MolecularExplainer:
  486. def __init__(self, model, device='cpu'):
  487. self.device = device
  488. self.model = model.to(device)
  489. self.model.eval()
  490. self.atom_symbols = ['C', 'N', 'O', 'S', 'F', 'P', 'Cl', 'Br', 'I', 'Unknown']
  491. self.substructure_identifier = SubstructureIdentifier()
  492. self.substructure_visualizer = SubstructureVisualizer()
  493. try:
  494. self.explainer = Explainer(
  495. model=self.model,
  496. algorithm=GNNExplainer(epochs=100, lr=0.01),
  497. explanation_type='model',
  498. node_mask_type='attributes',
  499. edge_mask_type='object',
  500. model_config=ModelConfig(
  501. mode='regression',
  502. task_level='graph',
  503. return_type='raw'
  504. )
  505. )
  506. except Exception as e:
  507. print(f"Warning: Could not initialize GNNExplainer: {e}")
  508. self.explainer = None
  509. def get_atom_symbol_from_features(self, atom_features):
  510. atom_idx = torch.argmax(atom_features[:10]).item()
  511. return self.atom_symbols[atom_idx]
  512. def simple_gradient_explanation(self, data):
  513. self.model.eval()
  514. data = data.to(self.device)
  515. if not hasattr(data, 'batch') or data.batch is None:
  516. data.batch = torch.zeros(data.x.size(0), dtype=torch.long, device=self.device)
  517. data.x.requires_grad_(True)
  518. prediction = self.model(data.x, data.edge_index, data.batch)
  519. prediction.backward()
  520. node_importance = torch.norm(data.x.grad, dim=1)
  521. return {
  522. 'node_mask': node_importance.cpu().detach(),
  523. 'edge_mask': None,
  524. 'prediction': prediction.cpu().detach(),
  525. 'data': data.cpu()
  526. }
  527. def explain_molecule(self, data):
  528. try:
  529. if self.explainer is not None:
  530. data = data.to(self.device)
  531. if not hasattr(data, 'batch') or data.batch is None:
  532. data.batch = torch.zeros(data.x.size(0), dtype=torch.long, device=self.device)
  533. explanation = self.explainer(
  534. x=data.x,
  535. edge_index=data.edge_index,
  536. batch=data.batch
  537. )
  538. return {
  539. 'node_mask': explanation.node_mask.cpu() if explanation.node_mask is not None else None,
  540. 'edge_mask': explanation.edge_mask.cpu() if explanation.edge_mask is not None else None,
  541. 'prediction': explanation.prediction.cpu() if hasattr(explanation, 'prediction') else None,
  542. 'data': data.cpu()
  543. }
  544. else:
  545. return self.simple_gradient_explanation(data)
  546. except Exception as e:
  547. print(f"Error in explanation: {e}")
  548. try:
  549. return self.simple_gradient_explanation(data)
  550. except Exception as e2:
  551. print(f"Error in gradient explanation: {e2}")
  552. return None
  553. def process_node_importance(self, node_mask, num_nodes):
  554. if node_mask is None:
  555. return np.full(num_nodes, 0.5)
  556. node_mask_np = node_mask.numpy()
  557. if len(node_mask_np.shape) > 1:
  558. if node_mask_np.shape[0] == num_nodes:
  559. node_colors = np.linalg.norm(node_mask_np, axis=1)
  560. elif node_mask_np.shape[1] == num_nodes:
  561. node_colors = np.linalg.norm(node_mask_np, axis=0)
  562. else:
  563. node_colors = np.full(num_nodes, np.mean(node_mask_np))
  564. else:
  565. node_colors = node_mask_np
  566. if len(node_colors) > num_nodes:
  567. node_colors = node_colors[:num_nodes]
  568. elif len(node_colors) < num_nodes:
  569. avg_importance = node_colors.mean() if len(node_colors) > 0 else 0.5
  570. padded_colors = np.full(num_nodes, avg_importance)
  571. padded_colors[:len(node_colors)] = node_colors
  572. node_colors = padded_colors
  573. if len(node_colors) > 0 and node_colors.max() > node_colors.min():
  574. node_colors = (node_colors - node_colors.min()) / (node_colors.max() - node_colors.min())
  575. else:
  576. node_colors = np.full(num_nodes, 0.5)
  577. return node_colors
  578. def visualize_molecule_explanation(self, explanation, smiles, target_value=None, title="Molecule Explanation",
  579. save_path=None):
  580. if explanation is None:
  581. print("No explanation to visualize")
  582. return
  583. data = explanation['data']
  584. node_mask = explanation['node_mask']
  585. edge_mask = explanation['edge_mask']
  586. prediction = explanation['prediction']
  587. fig, axes = plt.subplots(1, 2, figsize=(18, 8))
  588. try:
  589. mol = Chem.MolFromSmiles(smiles)
  590. if mol is not None:
  591. img = Draw.MolToImage(mol, size=(400, 400))
  592. axes[0].imshow(img)
  593. axes[0].set_title(f'Original Molecule\nSMILES: {smiles[:50]}...',
  594. fontsize=16, fontweight='bold', pad=15)
  595. axes[0].axis('off')
  596. except Exception as e:
  597. print(f"Warning: Could not draw molecule with RDKit: {e}")
  598. axes[0].text(0.5, 0.5, 'Could not draw molecule', ha='center', va='center',
  599. fontsize=16, fontweight='bold')
  600. axes[0].set_title('Original Molecule', fontsize=16, fontweight='bold')
  601. axes[0].axis('off')
  602. num_nodes = data.x.size(0)
  603. node_colors = self.process_node_importance(node_mask, num_nodes)
  604. G = nx.Graph()
  605. for i in range(num_nodes):
  606. atom_symbol = self.get_atom_symbol_from_features(data.x[i])
  607. G.add_node(i, symbol=atom_symbol)
  608. edge_index = data.edge_index.cpu().numpy()
  609. for i in range(edge_index.shape[1]):
  610. G.add_edge(edge_index[0, i], edge_index[1, i])
  611. try:
  612. pos = nx.spring_layout(G, seed=42)
  613. except:
  614. pos = nx.circular_layout(G)
  615. node_x = [pos[i][0] for i in range(num_nodes)]
  616. node_y = [pos[i][1] for i in range(num_nodes)]
  617. try:
  618. scatter = axes[1].scatter(
  619. node_x, node_y,
  620. c=node_colors,
  621. cmap='RdYlBu_r',
  622. s=600,
  623. alpha=0.8,
  624. edgecolors='black',
  625. linewidth=1.5
  626. )
  627. cbar = plt.colorbar(scatter, ax=axes[1], label='Node Importance', shrink=0.8)
  628. cbar.ax.tick_params(labelsize=14)
  629. cbar.set_label('Node Importance', fontsize=16, fontweight='bold')
  630. except Exception as e:
  631. print(f"Error in scatter plot: {e}")
  632. axes[1].scatter(node_x, node_y, c='lightblue', s=600, alpha=0.8, edgecolors='black')
  633. try:
  634. for edge in G.edges():
  635. x_coords = [pos[edge[0]][0], pos[edge[1]][0]]
  636. y_coords = [pos[edge[0]][1], pos[edge[1]][1]]
  637. axes[1].plot(x_coords, y_coords, 'k-', alpha=0.5, linewidth=2)
  638. except Exception as e:
  639. print(f"Warning: Could not draw edges: {e}")
  640. try:
  641. for i in range(num_nodes):
  642. atom_symbol = self.get_atom_symbol_from_features(data.x[i])
  643. axes[1].text(pos[i][0], pos[i][1], atom_symbol,
  644. ha='center', va='center', fontsize=14, fontweight='bold')
  645. except Exception as e:
  646. print(f"Warning: Could not add labels: {e}")
  647. pred_val = prediction.item() if prediction is not None else "N/A"
  648. if target_value is not None:
  649. title_text = f'Node Importance\nTrue y: {target_value:.4f}\nPrediction: {pred_val:.4f}'
  650. else:
  651. title_text = f'Node Importance\nPrediction: {pred_val:.4f}'
  652. axes[1].set_title(title_text, fontsize=16, fontweight='bold', pad=15)
  653. axes[1].axis('off')
  654. axes[1].set_aspect('equal')
  655. plt.subplots_adjust(left=0.05, right=0.95, top=0.9, bottom=0.1, wspace=0.3)
  656. if save_path:
  657. try:
  658. plt.savefig(save_path, dpi=300, bbox_inches='tight')
  659. print(f"Image saved: {save_path}")
  660. except Exception as e:
  661. print(f"Failed to save image: {e}")
  662. plt.show()
  663. def visualize_selected_molecule(self, explanation, smiles, target_value=None, title="Selected Molecule",
  664. save_path=None):
  665. if explanation is None:
  666. print("No explanation to visualize")
  667. return
  668. data = explanation['data']
  669. node_mask = explanation['node_mask']
  670. edge_mask = explanation['edge_mask']
  671. prediction = explanation['prediction']
  672. fig, axes = plt.subplots(1, 2, figsize=(18, 8))
  673. try:
  674. mol = Chem.MolFromSmiles(smiles)
  675. if mol is not None:
  676. img = Draw.MolToImage(mol, size=(400, 400))
  677. axes[0].imshow(img)
  678. axes[0].set_title(f'Original Molecule\nSMILES: {smiles[:50]}...',
  679. fontsize=16, fontweight='bold', pad=15)
  680. axes[0].axis('off')
  681. except Exception as e:
  682. print(f"Warning: Could not draw molecule with RDKit: {e}")
  683. axes[0].text(0.5, 0.5, 'Could not draw molecule', ha='center', va='center',
  684. fontsize=16, fontweight='bold')
  685. axes[0].set_title('Original Molecule', fontsize=16, fontweight='bold')
  686. axes[0].axis('off')
  687. num_nodes = data.x.size(0)
  688. node_colors = self.process_node_importance(node_mask, num_nodes)
  689. G = nx.Graph()
  690. for i in range(num_nodes):
  691. atom_symbol = self.get_atom_symbol_from_features(data.x[i])
  692. G.add_node(i, symbol=atom_symbol)
  693. edge_index = data.edge_index.cpu().numpy()
  694. for i in range(edge_index.shape[1]):
  695. G.add_edge(edge_index[0, i], edge_index[1, i])
  696. try:
  697. pos = nx.spring_layout(G, seed=42)
  698. except:
  699. pos = nx.circular_layout(G)
  700. node_x = [pos[i][0] for i in range(num_nodes)]
  701. node_y = [pos[i][1] for i in range(num_nodes)]
  702. try:
  703. scatter = axes[1].scatter(
  704. node_x, node_y,
  705. c=node_colors,
  706. cmap='RdYlBu_r',
  707. s=600,
  708. alpha=0.8,
  709. edgecolors='black',
  710. linewidth=1.5
  711. )
  712. cbar = plt.colorbar(scatter, ax=axes[1], label='Importance Score', shrink=0.8)
  713. cbar.ax.tick_params(labelsize=14)
  714. cbar.set_label('Importance Score', fontsize=16, fontweight='bold')
  715. except Exception as e:
  716. print(f"Error in scatter plot: {e}")
  717. axes[1].scatter(node_x, node_y, c='lightblue', s=600, alpha=0.8, edgecolors='black')
  718. try:
  719. for edge in G.edges():
  720. x_coords = [pos[edge[0]][0], pos[edge[1]][0]]
  721. y_coords = [pos[edge[0]][1], pos[edge[1]][1]]
  722. axes[1].plot(x_coords, y_coords, 'k-', alpha=0.5, linewidth=2)
  723. except Exception as e:
  724. print(f"Warning: Could not draw edges: {e}")
  725. try:
  726. for i in range(num_nodes):
  727. atom_symbol = self.get_atom_symbol_from_features(data.x[i])
  728. axes[1].text(pos[i][0], pos[i][1], atom_symbol,
  729. ha='center', va='center', fontsize=14, fontweight='bold')
  730. except Exception as e:
  731. print(f"Warning: Could not add labels: {e}")
  732. if target_value is not None:
  733. title_text = f'True y: {target_value:.4f}'
  734. else:
  735. title_text = 'True y: N/A'
  736. axes[1].set_title(title_text, fontsize=16, fontweight='bold', pad=15)
  737. axes[1].axis('off')
  738. axes[1].set_aspect('equal')
  739. plt.subplots_adjust(left=0.05, right=0.95, top=0.9, bottom=0.1, wspace=0.3)
  740. if save_path:
  741. try:
  742. plt.savefig(save_path, dpi=300, bbox_inches='tight')
  743. print(f"Image saved: {save_path}")
  744. except Exception as e:
  745. print(f"Failed to save image: {e}")
  746. plt.show()
  747. def analyze_feature_importance(self, explanations_data):
  748. all_node_importances = []
  749. all_atom_types = []
  750. all_predictions = []
  751. for exp_data in explanations_data:
  752. explanation = exp_data['explanation']
  753. if explanation is None or explanation['node_mask'] is None:
  754. continue
  755. node_mask = explanation['node_mask']
  756. data = explanation['data']
  757. prediction = explanation['prediction']
  758. num_nodes = data.x.size(0)
  759. node_importances = self.process_node_importance(node_mask, num_nodes)
  760. for i in range(num_nodes):
  761. atom_features = data.x[i]
  762. atom_symbol = self.get_atom_symbol_from_features(atom_features)
  763. all_node_importances.append(node_importances[i])
  764. all_atom_types.append(atom_symbol)
  765. all_predictions.append(prediction.item() if prediction is not None else 0)
  766. importance_df = pd.DataFrame({
  767. 'atom_type': all_atom_types,
  768. 'importance': all_node_importances,
  769. 'prediction': all_predictions
  770. })
  771. return importance_df
  772. def find_important_substructures(self, explanations_data, importance_threshold=0.5):
  773. important_substructures = []
  774. for exp_data in explanations_data:
  775. explanation = exp_data['explanation']
  776. smiles = exp_data['smiles']
  777. target_value = exp_data.get('original_target', 0)
  778. if explanation is None or explanation['node_mask'] is None:
  779. continue
  780. data = explanation['data']
  781. node_mask = explanation['node_mask']
  782. prediction = explanation['prediction']
  783. num_nodes = data.x.size(0)
  784. node_importances = self.process_node_importance(node_mask, num_nodes)
  785. important_atoms = []
  786. important_atom_indices = []
  787. for i in range(num_nodes):
  788. if node_importances[i] > importance_threshold:
  789. atom_symbol = self.get_atom_symbol_from_features(data.x[i])
  790. important_atoms.append((i, atom_symbol, node_importances[i]))
  791. important_atom_indices.append(i)
  792. if not important_atoms:
  793. continue
  794. try:
  795. mol = Chem.MolFromSmiles(smiles)
  796. if mol is None:
  797. continue
  798. all_substructures = self.substructure_identifier.identify_substructures_in_molecule(mol)
  799. functional_groups = self.substructure_identifier.get_functional_groups(mol)
  800. local_substructures = self.substructure_identifier.extract_important_substructures(
  801. mol, important_atom_indices, radius=2
  802. )
  803. relevant_substructures = {}
  804. for sub_name, sub_info in all_substructures.items():
  805. for match in sub_info['matches']:
  806. if any(atom_idx in match for atom_idx in important_atom_indices):
  807. if sub_name not in relevant_substructures:
  808. relevant_substructures[sub_name] = {
  809. 'pattern': sub_info['pattern'],
  810. 'matches': [],
  811. 'importance_scores': []
  812. }
  813. match_importance = []
  814. for atom_idx in match:
  815. if atom_idx in important_atom_indices:
  816. idx_in_list = important_atom_indices.index(atom_idx)
  817. match_importance.append(important_atoms[idx_in_list][2])
  818. if match_importance:
  819. avg_importance = np.mean(match_importance)
  820. relevant_substructures[sub_name]['matches'].append(match)
  821. relevant_substructures[sub_name]['importance_scores'].append(avg_importance)
  822. edge_index = data.edge_index.cpu().numpy()
  823. important_edges = []
  824. for i in range(edge_index.shape[1]):
  825. if (edge_index[0, i] in important_atom_indices and
  826. edge_index[1, i] in important_atom_indices):
  827. important_edges.append((edge_index[0, i], edge_index[1, i]))
  828. important_substructures.append({
  829. 'smiles': smiles,
  830. 'important_atoms': important_atoms,
  831. 'important_edges': important_edges,
  832. 'prediction': prediction.item() if prediction is not None else 0,
  833. 'target_value': target_value,
  834. 'num_important_atoms': len(important_atoms),
  835. 'known_substructures': relevant_substructures,
  836. 'functional_groups': functional_groups,
  837. 'local_substructures': local_substructures,
  838. 'substructure_summary': {
  839. 'num_known_substructures': len(relevant_substructures),
  840. 'num_functional_groups': len(functional_groups),
  841. 'num_local_fragments': len(local_substructures)
  842. }
  843. })
  844. except Exception as e:
  845. print(f"Error analyzing substructures for {smiles}: {e}")
  846. important_substructures.append({
  847. 'smiles': smiles,
  848. 'important_atoms': important_atoms,
  849. 'important_edges': important_edges,
  850. 'prediction': prediction.item() if prediction is not None else 0,
  851. 'target_value': target_value,
  852. 'num_important_atoms': len(important_atoms),
  853. 'known_substructures': {},
  854. 'functional_groups': [],
  855. 'local_substructures': [],
  856. 'substructure_summary': {
  857. 'num_known_substructures': 0,
  858. 'num_functional_groups': 0,
  859. 'num_local_fragments': 0
  860. }
  861. })
  862. return important_substructures
  863. def analyze_full_dataset_substructures(self, test_csv_file, importance_threshold=0.3):
  864. print("\n" + "=" * 50)
  865. print("Full Dataset Substructure Analysis")
  866. print("=" * 50)
  867. test_df = pd.read_csv(test_csv_file)
  868. full_dataset_substructures = []
  869. successful_count = 0
  870. print(f"Analyzing substructures for all {len(test_df)} molecules in the dataset...")
  871. for i, row in test_df.iterrows():
  872. try:
  873. smiles = str(row['Smiles'])
  874. target_value = row.get('pchembl', 0)
  875. atom_features, edge_index = smiles_to_graph(smiles)
  876. data = Data(x=atom_features, edge_index=edge_index)
  877. explanation = self.explain_molecule(data)
  878. if explanation is not None:
  879. node_mask = explanation['node_mask']
  880. num_nodes = data.x.size(0)
  881. node_importances = self.process_node_importance(node_mask, num_nodes)
  882. important_atoms = []
  883. important_atom_indices = []
  884. for j in range(num_nodes):
  885. if node_importances[j] > importance_threshold:
  886. atom_symbol = self.get_atom_symbol_from_features(data.x[j])
  887. important_atoms.append((j, atom_symbol, node_importances[j]))
  888. important_atom_indices.append(j)
  889. if important_atoms:
  890. mol = Chem.MolFromSmiles(smiles)
  891. if mol is not None:
  892. all_substructures = self.substructure_identifier.identify_substructures_in_molecule(mol)
  893. functional_groups = self.substructure_identifier.get_functional_groups(mol)
  894. local_substructures = self.substructure_identifier.extract_important_substructures(
  895. mol, important_atom_indices, radius=2
  896. )
  897. relevant_substructures = {}
  898. for sub_name, sub_info in all_substructures.items():
  899. for match in sub_info['matches']:
  900. if any(atom_idx in match for atom_idx in important_atom_indices):
  901. if sub_name not in relevant_substructures:
  902. relevant_substructures[sub_name] = {
  903. 'pattern': sub_info['pattern'],
  904. 'matches': [],
  905. 'importance_scores': []
  906. }
  907. match_importance = []
  908. for atom_idx in match:
  909. if atom_idx in important_atom_indices:
  910. idx_in_list = important_atom_indices.index(atom_idx)
  911. match_importance.append(important_atoms[idx_in_list][2])
  912. if match_importance:
  913. avg_importance = np.mean(match_importance)
  914. relevant_substructures[sub_name]['matches'].append(match)
  915. relevant_substructures[sub_name]['importance_scores'].append(avg_importance)
  916. edge_index = data.edge_index.cpu().numpy()
  917. important_edges = []
  918. for j in range(edge_index.shape[1]):
  919. if (edge_index[0, j] in important_atom_indices and
  920. edge_index[1, j] in important_atom_indices):
  921. important_edges.append((edge_index[0, j], edge_index[1, j]))
  922. full_dataset_substructures.append({
  923. 'smiles': smiles,
  924. 'important_atoms': important_atoms,
  925. 'important_edges': important_edges,
  926. 'prediction': explanation['prediction'].item() if explanation[
  927. 'prediction'] is not None else 0,
  928. 'target_value': target_value,
  929. 'num_important_atoms': len(important_atoms),
  930. 'known_substructures': relevant_substructures,
  931. 'functional_groups': functional_groups,
  932. 'local_substructures': local_substructures,
  933. 'substructure_summary': {
  934. 'num_known_substructures': len(relevant_substructures),
  935. 'num_functional_groups': len(functional_groups),
  936. 'num_local_fragments': len(local_substructures)
  937. }
  938. })
  939. successful_count += 1
  940. except Exception as e:
  941. continue
  942. if (i + 1) % 100 == 0:
  943. print(f" Progress: {i + 1}/{len(test_df)}, successful: {successful_count}")
  944. print(f"Full dataset analysis completed: {successful_count} molecules with important substructures")
  945. return full_dataset_substructures
  946. def plot_feature_importance_summary(self, importance_df, save_path=None):
  947. if importance_df.empty:
  948. print("No importance data to plot")
  949. return
  950. custom_colors = ['#98CFE6', '#ADE7A8', '#F39F4E', '#EEB7D3', '#DBDAD3', '#FFDF97']
  951. try:
  952. atom_importance = importance_df.groupby('atom_type')['importance'].agg(['mean', 'std', 'count'])
  953. atom_importance = atom_importance.sort_values('mean', ascending=False)
  954. atom_counts = importance_df['atom_type'].value_counts()
  955. plt.figure(figsize=(14, 10))
  956. for i, (atom_type, stats) in enumerate(atom_importance.iterrows()):
  957. color_idx = i % len(custom_colors)
  958. plt.bar(atom_type, stats['mean'], yerr=stats['std'], capsize=5,
  959. alpha=0.8, color=custom_colors[color_idx], edgecolor='white', linewidth=1)
  960. plt.title('Average Atom Importance', fontsize=20, fontweight='bold', pad=25)
  961. plt.xlabel('Atom Type', fontsize=18, fontweight='bold')
  962. plt.ylabel('Average Importance', fontsize=18, fontweight='bold')
  963. plt.xticks(rotation=45, fontsize=16)
  964. plt.yticks(fontsize=16)
  965. plt.grid(axis='y', alpha=0.3)
  966. for i, (atom_type, stats) in enumerate(atom_importance.iterrows()):
  967. plt.text(i, stats['mean'] + stats['std'] + 0.01, f'{stats["mean"]:.3f}',
  968. ha='center', va='bottom', fontsize=12, fontweight='bold')
  969. plt.subplots_adjust(left=0.15, right=0.95, top=0.9, bottom=0.2)
  970. if save_path:
  971. plt.savefig(f"{save_path}_subplot1_atom_importance.png", dpi=300, bbox_inches='tight')
  972. print(f"Subplot 1 saved: {save_path}_subplot1_atom_importance.png")
  973. plt.show()
  974. plt.figure(figsize=(14, 10))
  975. total_contribution = importance_df.groupby('atom_type')['importance'].sum().sort_values(ascending=False)
  976. cumulative_contribution = total_contribution.cumsum() / total_contribution.sum() * 100
  977. plt.plot(range(1, len(cumulative_contribution) + 1), cumulative_contribution.values,
  978. 'o-', linewidth=6, markersize=12, color=custom_colors[0], markerfacecolor='white',
  979. markeredgewidth=3, markeredgecolor=custom_colors[0])
  980. plt.axhline(y=80, color='red', linestyle='--', alpha=0.7, label='80% Contribution', linewidth=4)
  981. plt.xlabel('Atom Type Rank', fontsize=18, fontweight='bold')
  982. plt.ylabel('Cumulative Contribution (%)', fontsize=18, fontweight='bold')
  983. plt.title('Cumulative Importance Contribution', fontsize=20, fontweight='bold', pad=25)
  984. plt.xticks(fontsize=16)
  985. plt.yticks(fontsize=16)
  986. plt.grid(True, alpha=0.3)
  987. plt.legend(fontsize=16, loc='lower right')
  988. for i, (atom_type, contrib) in enumerate(cumulative_contribution.items()):
  989. if i < 8:
  990. plt.annotate(atom_type, (i + 1, contrib), xytext=(0, 15),
  991. textcoords='offset points', fontsize=14, fontweight='bold',
  992. ha='center', va='bottom')
  993. plt.subplots_adjust(left=0.15, right=0.95, top=0.9, bottom=0.15)
  994. if save_path:
  995. plt.savefig(f"{save_path}_subplot2_cumulative_contribution.png", dpi=300, bbox_inches='tight')
  996. print(f"Subplot 2 saved: {save_path}_subplot2_cumulative_contribution.png")
  997. plt.show()
  998. plt.figure(figsize=(14, 12))
  999. pie_colors = [custom_colors[i % len(custom_colors)] for i in range(len(atom_counts))]
  1000. def autopct_func_atoms(pct, idx=0):
  1001. if idx < len(atom_counts) - 5:
  1002. return f'{pct:.1f}%'
  1003. else:
  1004. return ''
  1005. counter = [0]
  1006. def autopct_with_counter_atoms(pct):
  1007. result = autopct_func_atoms(pct, counter[0])
  1008. counter[0] += 1
  1009. return result
  1010. wedges, texts, autotexts = plt.pie(atom_counts.values, autopct=autopct_with_counter_atoms,
  1011. colors=pie_colors, startangle=90,
  1012. textprops={'fontsize': 16, 'fontweight': 'bold'})
  1013. plt.title('Atom Type Distribution', fontsize=20, fontweight='bold', pad=25)
  1014. for autotext in autotexts:
  1015. autotext.set_fontsize(16)
  1016. autotext.set_color('black')
  1017. autotext.set_fontweight('bold')
  1018. plt.legend(wedges, atom_counts.index, title="Atom Types",
  1019. title_fontsize=16, fontsize=14,
  1020. loc="center left", bbox_to_anchor=(1.1, 0.5))
  1021. plt.subplots_adjust(left=0.1, right=0.75, top=0.9, bottom=0.1)
  1022. if save_path:
  1023. plt.savefig(f"{save_path}_subplot3_atom_distribution.png", dpi=300, bbox_inches='tight')
  1024. print(f"Subplot 3 saved: {save_path}_subplot3_atom_distribution.png")
  1025. plt.show()
  1026. plt.figure(figsize=(14, 10))
  1027. atom_types_for_box = []
  1028. importance_values_for_box = []
  1029. for atom_type in atom_importance.index[:10]:
  1030. atom_data = importance_df[importance_df['atom_type'] == atom_type]['importance']
  1031. atom_types_for_box.extend([atom_type] * len(atom_data))
  1032. importance_values_for_box.extend(atom_data.values)
  1033. box_df = pd.DataFrame({
  1034. 'atom_type': atom_types_for_box,
  1035. 'importance': importance_values_for_box
  1036. })
  1037. import seaborn as sns
  1038. sns.set_context("paper", font_scale=1.5)
  1039. sns.boxplot(data=box_df, x='atom_type', y='importance',
  1040. palette=custom_colors[:len(atom_importance.index[:10])])
  1041. plt.title('Importance Distribution by Atom Type', fontsize=20, fontweight='bold', pad=25)
  1042. plt.xlabel('Atom Type', fontsize=18, fontweight='bold')
  1043. plt.ylabel('Importance Score', fontsize=18, fontweight='bold')
  1044. plt.xticks(rotation=45, fontsize=16)
  1045. plt.yticks(fontsize=16)
  1046. plt.grid(axis='y', alpha=0.3)
  1047. plt.subplots_adjust(left=0.15, right=0.95, top=0.9, bottom=0.2)
  1048. if save_path:
  1049. plt.savefig(f"{save_path}_subplot4_importance_distribution.png", dpi=300, bbox_inches='tight')
  1050. print(f"Subplot 4 saved: {save_path}_subplot4_importance_distribution.png")
  1051. plt.show()
  1052. sns.reset_defaults()
  1053. print(f"\nAll four subplots have been generated and saved separately!")
  1054. except Exception as e:
  1055. print(f"Error plotting feature importance: {e}")
  1056. def visualize_all_substructures(self, important_substructures, full_dataset_substructures,
  1057. base_save_path="substructure_analysis"):
  1058. print("\nGenerating substructure visualizations...")
  1059. print(" - Generating substructure importance summary (4 separate plots)...")
  1060. self.substructure_visualizer.visualize_substructure_importance_summary(
  1061. important_substructures,
  1062. save_path=f"{base_save_path}_importance_summary"
  1063. )
  1064. print(" - Generating highlighted molecules from full dataset...")
  1065. selected_molecules = self.substructure_visualizer.visualize_substructure_in_molecules(
  1066. full_dataset_substructures,
  1067. num_examples=6,
  1068. save_path=f"{base_save_path}_highlighted_molecules.png"
  1069. )
  1070. print(" - Generating substructure heatmap...")
  1071. self.substructure_visualizer.create_substructure_heatmap(
  1072. important_substructures,
  1073. save_path=f"{base_save_path}_heatmap.png"
  1074. )
  1075. print("Substructure visualization completed!")
  1076. return selected_molecules
  1077. def load_best_model(model_path='best_model.pth'):
  1078. try:
  1079. checkpoint = torch.load(model_path, map_location='cpu')
  1080. gat_graphsage_model = GAT_GraphSAGE(n_output=1, num_features_xd=35)
  1081. gat_graphsage_model.load_state_dict(checkpoint['gat_graphsage_model_state_dict'])
  1082. scaler = checkpoint.get('scaler', None)
  1083. print("Model loaded successfully!")
  1084. return gat_graphsage_model, scaler
  1085. except Exception as e:
  1086. print(f"Error loading model: {e}")
  1087. return None, None
  1088. def smiles_to_graph(smiles):
  1089. mol = Chem.MolFromSmiles(smiles)
  1090. if mol is None:
  1091. raise ValueError(f"Invalid SMILES string: {smiles}")
  1092. num_atoms = mol.GetNumAtoms()
  1093. atom_features = []
  1094. for atom in mol.GetAtoms():
  1095. results = one_of_k_encoding_unk(atom.GetSymbol(), ['C', 'N', 'O', 'S', 'F', 'P', 'Cl', 'Br', 'I', 'Unknown']) + \
  1096. one_of_k_encoding_unk(atom.GetDegree(), [0, 1, 2, 3, 4, 5, 6]) + \
  1097. one_of_k_encoding_unk(atom.GetImplicitValence(), [0, 1, 2, 3, 4, 5, 6]) + \
  1098. one_of_k_encoding_unk(atom.GetHybridization(), [
  1099. Chem.rdchem.HybridizationType.SP, Chem.rdchem.HybridizationType.SP2,
  1100. Chem.rdchem.HybridizationType.SP3, Chem.rdchem.HybridizationType.SP3D,
  1101. Chem.rdchem.HybridizationType.SP3D2
  1102. ]) + [atom.GetIsAromatic()] + \
  1103. one_of_k_encoding_unk(atom.GetTotalNumHs(), [0, 1, 2, 3, 4])
  1104. atom_feats = np.array(results).astype(np.float32)
  1105. atom_features.append(atom_feats)
  1106. atom_features = torch.tensor(atom_features, dtype=torch.float)
  1107. adj_matrix = torch.zeros((num_atoms, num_atoms), dtype=torch.float)
  1108. for bond in mol.GetBonds():
  1109. start, end = bond.GetBeginAtomIdx(), bond.GetEndAtomIdx()
  1110. adj_matrix[start, end] = 1.0
  1111. adj_matrix[end, start] = 1.0
  1112. edge_index = adj_matrix.nonzero(as_tuple=False).t().long()
  1113. return atom_features, edge_index
  1114. def quick_importance_analysis_all(test_csv_file, model_path):
  1115. print("Stage 1: Quick analysis of all molecules...")
  1116. print("-" * 50)
  1117. test_df = pd.read_csv(test_csv_file)
  1118. gat_model, scaler = load_best_model(model_path)
  1119. explainable_model = ExplainableGATGraphSAGE(gat_model)
  1120. explainer = MolecularExplainer(explainable_model)
  1121. molecule_info = []
  1122. successful_count = 0
  1123. for i, row in test_df.iterrows():
  1124. try:
  1125. smiles = str(row['Smiles'])
  1126. atom_features, edge_index = smiles_to_graph(smiles)
  1127. data = Data(x=atom_features, edge_index=edge_index)
  1128. explanation = explainer.simple_gradient_explanation(data)
  1129. if explanation:
  1130. molecule_info.append({
  1131. 'index': i,
  1132. 'smiles': smiles,
  1133. 'prediction': explanation['prediction'].item(),
  1134. 'avg_importance': explanation['node_mask'].mean().item(),
  1135. 'max_importance': explanation['node_mask'].max().item(),
  1136. 'std_importance': explanation['node_mask'].std().item(),
  1137. 'num_atoms': data.x.size(0),
  1138. 'target': row.get('pchembl', 0)
  1139. })
  1140. successful_count += 1
  1141. except Exception as e:
  1142. continue
  1143. if (i + 1) % 100 == 0:
  1144. print(f" Quick analysis progress: {i + 1}/961, successful: {successful_count}")
  1145. print(f"Stage 1 completed: Successfully analyzed {successful_count} molecules")
  1146. return molecule_info
  1147. def stratified_sample_by_column(df, column, target_count):
  1148. try:
  1149. df_copy = df.copy()
  1150. df_copy['quartile'] = pd.qcut(df_copy[column], q=5, labels=False, duplicates='drop')
  1151. sampled_indices = []
  1152. samples_per_quartile = target_count // 5
  1153. for q in df_copy['quartile'].unique():
  1154. if pd.isna(q):
  1155. continue
  1156. quartile_data = df_copy[df_copy['quartile'] == q]
  1157. if len(quartile_data) > 0:
  1158. sample_count = min(samples_per_quartile, len(quartile_data))
  1159. sampled = quartile_data.sample(n=sample_count, random_state=42)
  1160. sampled_indices.extend(sampled['index'].tolist())
  1161. return sampled_indices
  1162. except Exception as e:
  1163. print(f"Stratified sampling failed, using random sampling: {e}")
  1164. return df.sample(n=min(target_count, len(df)), random_state=42)['index'].tolist()
  1165. def select_representative_molecules(molecule_info, target_count=200):
  1166. print("\nStage 2: Selecting representative molecules for detailed analysis...")
  1167. print("-" * 50)
  1168. df = pd.DataFrame(molecule_info)
  1169. if len(df) < target_count:
  1170. print(f"Available molecules ({len(df)}) < target count ({target_count}), will analyze all available")
  1171. return df['index'].tolist()
  1172. selected_indices = []
  1173. print(" - Stratified sampling by prediction values...")
  1174. pred_samples = stratified_sample_by_column(df, 'prediction', int(target_count * 0.4))
  1175. selected_indices.extend(pred_samples)
  1176. print(" - Stratified sampling by average importance...")
  1177. remaining_df = df[~df['index'].isin(selected_indices)]
  1178. if len(remaining_df) > 0:
  1179. imp_samples = stratified_sample_by_column(remaining_df, 'avg_importance', int(target_count * 0.3))
  1180. selected_indices.extend(imp_samples)
  1181. print(" - Stratified sampling by molecule size...")
  1182. remaining_df = df[~df['index'].isin(selected_indices)]
  1183. if len(remaining_df) > 0:
  1184. size_samples = stratified_sample_by_column(remaining_df, 'num_atoms', int(target_count * 0.2))
  1185. selected_indices.extend(size_samples)
  1186. print(" - Random sampling for remaining molecules...")
  1187. remaining_df = df[~df['index'].isin(selected_indices)]
  1188. remaining_count = target_count - len(selected_indices)
  1189. if remaining_count > 0 and len(remaining_df) > 0:
  1190. random_samples = remaining_df.sample(n=min(remaining_count, len(remaining_df)),
  1191. random_state=42)['index'].tolist()
  1192. selected_indices.extend(random_samples)
  1193. print(f"Selected {len(selected_indices)} representative molecules for detailed analysis")
  1194. selected_df = df[df['index'].isin(selected_indices)]
  1195. print(f"\nSampling statistics:")
  1196. print(f" Prediction range: {selected_df['prediction'].min():.3f} - {selected_df['prediction'].max():.3f}")
  1197. print(f" Importance range: {selected_df['avg_importance'].min():.3f} - {selected_df['avg_importance'].max():.3f}")
  1198. print(f" Molecule size range: {selected_df['num_atoms'].min()} - {selected_df['num_atoms'].max()} atoms")
  1199. return selected_indices
  1200. def perform_detailed_analysis(test_csv_file, model_path, selected_indices, importance_threshold=0.3):
  1201. print("\nStage 3: Detailed substructure analysis of selected molecules...")
  1202. print("-" * 50)
  1203. test_df = pd.read_csv(test_csv_file)
  1204. gat_model, scaler = load_best_model(model_path)
  1205. explainable_model = ExplainableGATGraphSAGE(gat_model)
  1206. explainer = MolecularExplainer(explainable_model)
  1207. explanations_data = []
  1208. successful_count = 0
  1209. for i, idx in enumerate(selected_indices):
  1210. row = test_df.iloc[idx]
  1211. smiles = str(row['Smiles'])
  1212. try:
  1213. atom_features, edge_index = smiles_to_graph(smiles)
  1214. data = Data(x=atom_features, edge_index=edge_index)
  1215. explanation = explainer.explain_molecule(data)
  1216. if explanation is not None:
  1217. explanations_data.append({
  1218. 'smiles': smiles,
  1219. 'explanation': explanation,
  1220. 'original_target': row.get('pchembl', 0),
  1221. 'original_index': idx
  1222. })
  1223. successful_count += 1
  1224. if (i + 1) % 20 == 0:
  1225. print(f" Detailed analysis progress: {i + 1}/{len(selected_indices)}, successful: {successful_count}")
  1226. except Exception as e:
  1227. print(f" Molecule {idx} processing failed: {e}")
  1228. continue
  1229. print(f"Stage 3 completed: Successfully explained {successful_count} molecules")
  1230. print("\nPerforming feature importance analysis...")
  1231. importance_df = explainer.analyze_feature_importance(explanations_data)
  1232. print("Identifying important substructures from sampled data...")
  1233. important_substructures = explainer.find_important_substructures(
  1234. explanations_data, importance_threshold
  1235. )
  1236. print("Analyzing full dataset for qualified molecules...")
  1237. full_dataset_substructures = explainer.analyze_full_dataset_substructures(
  1238. test_csv_file, importance_threshold
  1239. )
  1240. return {
  1241. 'explanations_data': explanations_data,
  1242. 'importance_df': importance_df,
  1243. 'important_substructures': important_substructures,
  1244. 'full_dataset_substructures': full_dataset_substructures,
  1245. 'explainer': explainer
  1246. }
  1247. def combine_quick_and_detailed_results(quick_results, detailed_results):
  1248. print("\nStage 4: Combining and summarizing analysis results...")
  1249. print("-" * 50)
  1250. quick_df = pd.DataFrame(quick_results)
  1251. global_stats = {
  1252. 'total_molecules_analyzed': len(quick_df),
  1253. 'prediction_range': (quick_df['prediction'].min(), quick_df['prediction'].max()),
  1254. 'prediction_mean': quick_df['prediction'].mean(),
  1255. 'prediction_std': quick_df['prediction'].std(),
  1256. 'avg_importance_range': (quick_df['avg_importance'].min(), quick_df['avg_importance'].max()),
  1257. 'avg_importance_mean': quick_df['avg_importance'].mean(),
  1258. 'molecule_size_range': (quick_df['num_atoms'].min(), quick_df['num_atoms'].max()),
  1259. 'molecule_size_mean': quick_df['num_atoms'].mean()
  1260. }
  1261. combined_results = {
  1262. 'global_statistics': global_stats,
  1263. 'quick_analysis_results': quick_results,
  1264. 'detailed_analysis_results': detailed_results,
  1265. 'summary': {
  1266. 'total_molecules': len(quick_df),
  1267. 'detailed_molecules': len(detailed_results['explanations_data']),
  1268. 'identified_substructures': len(detailed_results['important_substructures']),
  1269. 'full_dataset_substructures': len(detailed_results['full_dataset_substructures']),
  1270. 'analysis_completeness': len(detailed_results['explanations_data']) / len(quick_df) * 100
  1271. }
  1272. }
  1273. return combined_results
  1274. def hybrid_analysis_strategy(test_csv_file, model_path, target_detailed_count=200, importance_threshold=0.3):
  1275. print("=" * 70)
  1276. print("Hybrid Analysis Strategy: Molecular Model Explainability Analysis")
  1277. print("=" * 70)
  1278. print(f"Data file: {test_csv_file}")
  1279. print(f"Model file: {model_path}")
  1280. print(f"Target detailed analysis count: {target_detailed_count}")
  1281. print(f"Importance threshold: {importance_threshold}")
  1282. print("=" * 70)
  1283. quick_results = quick_importance_analysis_all(test_csv_file, model_path)
  1284. if not quick_results:
  1285. print("Quick analysis failed, cannot continue")
  1286. return None
  1287. representative_indices = select_representative_molecules(
  1288. quick_results,
  1289. target_count=target_detailed_count
  1290. )
  1291. detailed_results = perform_detailed_analysis(
  1292. test_csv_file,
  1293. model_path,
  1294. representative_indices,
  1295. importance_threshold
  1296. )
  1297. final_results = combine_quick_and_detailed_results(quick_results, detailed_results)
  1298. generate_comprehensive_report(final_results)
  1299. return final_results
  1300. def generate_comprehensive_report(results):
  1301. print("\n" + "=" * 70)
  1302. print("Comprehensive Analysis Report")
  1303. print("=" * 70)
  1304. global_stats = results['global_statistics']
  1305. summary = results['summary']
  1306. print(f"\n[Global Statistics - Based on {global_stats['total_molecules_analyzed']} molecules]")
  1307. print("-" * 50)
  1308. print(f"Prediction distribution:")
  1309. print(f" Range: {global_stats['prediction_range'][0]:.3f} - {global_stats['prediction_range'][1]:.3f}")
  1310. print(f" Mean: {global_stats['prediction_mean']:.3f} ± {global_stats['prediction_std']:.3f}")
  1311. print(f"\nImportance distribution:")
  1312. print(f" Range: {global_stats['avg_importance_range'][0]:.3f} - {global_stats['avg_importance_range'][1]:.3f}")
  1313. print(f" Mean: {global_stats['avg_importance_mean']:.3f}")
  1314. print(f"\nMolecule size distribution:")
  1315. print(f" Range: {global_stats['molecule_size_range'][0]} - {global_stats['molecule_size_range'][1]} atoms")
  1316. print(f" Mean: {global_stats['molecule_size_mean']:.1f} atoms")
  1317. detailed = results['detailed_analysis_results']
  1318. print(f"\n[Detailed Analysis Results - Based on {summary['detailed_molecules']} representative molecules]")
  1319. print("-" * 50)
  1320. if not detailed['importance_df'].empty:
  1321. print("Atom importance statistics:")
  1322. atom_stats = detailed['importance_df'].groupby('atom_type')['importance'].agg(['mean', 'std', 'count'])
  1323. atom_stats = atom_stats.sort_values('mean', ascending=False)
  1324. for atom_type, stats in atom_stats.head(8).iterrows():
  1325. print(
  1326. f" {atom_type}: Average importance {stats['mean']:.3f} (±{stats['std']:.3f}), appeared {stats['count']} times")
  1327. important_substructures = detailed['important_substructures']
  1328. full_dataset_substructures = detailed['full_dataset_substructures']
  1329. if important_substructures:
  1330. print(f"\nFound {len(important_substructures)} molecules containing important substructures (from sample)")
  1331. print(
  1332. f"Found {len(full_dataset_substructures)} molecules containing important substructures (from full dataset)")
  1333. all_substructure_types = {}
  1334. all_functional_groups = Counter()
  1335. for struct in important_substructures:
  1336. for sub_name, sub_info in struct['known_substructures'].items():
  1337. if sub_name not in all_substructure_types:
  1338. all_substructure_types[sub_name] = {
  1339. 'count': 0,
  1340. 'total_importance': 0,
  1341. 'molecules': []
  1342. }
  1343. all_substructure_types[sub_name]['count'] += len(sub_info['matches'])
  1344. all_substructure_types[sub_name]['total_importance'] += sum(sub_info['importance_scores'])
  1345. all_substructure_types[sub_name]['molecules'].append(struct['smiles'][:30])
  1346. for fg in struct['functional_groups']:
  1347. all_functional_groups[fg['name']] += fg['count']
  1348. print(f"\nMost common important substructures (Top 10):")
  1349. sorted_substructures = sorted(all_substructure_types.items(),
  1350. key=lambda x: x[1]['count'], reverse=True)
  1351. for i, (sub_name, info) in enumerate(sorted_substructures[:10], 1):
  1352. avg_importance = info['total_importance'] / info['count'] if info['count'] > 0 else 0
  1353. print(f" {i:2d}. {sub_name:15s}: appeared {info['count']:3d} times, avg importance {avg_importance:.3f}")
  1354. print(f"\nMost common functional groups (Top 10):")
  1355. for i, (fg_name, count) in enumerate(all_functional_groups.most_common(10), 1):
  1356. print(f" {i:2d}. {fg_name:20s}: {count:3d} times")
  1357. print(f"\n[Analysis Completeness]")
  1358. print("-" * 30)
  1359. print(f"Total molecules: {summary['total_molecules']}")
  1360. print(f"Detailed analysis molecules: {summary['detailed_molecules']}")
  1361. print(f"Full dataset substructure analysis: {summary['full_dataset_substructures']}")
  1362. print(f"Analysis coverage: {summary['analysis_completeness']:.1f}%")
  1363. print(f"Identified important substructures: {summary['identified_substructures']}")
  1364. if not detailed['importance_df'].empty:
  1365. print(f"\nGenerating visualizations...")
  1366. try:
  1367. detailed['explainer'].plot_feature_importance_summary(
  1368. detailed['importance_df'],
  1369. 'enhanced_feature_importance_analysis'
  1370. )
  1371. selected_highlighted_molecules = detailed['explainer'].visualize_all_substructures(
  1372. detailed['important_substructures'],
  1373. detailed['full_dataset_substructures'],
  1374. base_save_path="substructure_analysis"
  1375. )
  1376. print(" - Generating selected molecule visualizations from full dataset (True y only)...")
  1377. qualified_molecules = []
  1378. for struct in full_dataset_substructures:
  1379. if (struct.get('target_value', 0) > 6 and
  1380. any(max(sub_info['importance_scores']) > 0.5
  1381. for sub_info in struct['known_substructures'].values()
  1382. if sub_info['importance_scores'])):
  1383. qualified_molecules.append(struct)
  1384. if len(qualified_molecules) >= 6:
  1385. break
  1386. print(f" Found {len(qualified_molecules)} qualified molecules from full dataset")
  1387. for i, struct_data in enumerate(qualified_molecules):
  1388. try:
  1389. smiles = struct_data['smiles']
  1390. target_value = struct_data.get('target_value')
  1391. atom_features, edge_index = smiles_to_graph(smiles)
  1392. data = Data(x=atom_features, edge_index=edge_index)
  1393. explanation = detailed['explainer'].explain_molecule(data)
  1394. if explanation is not None:
  1395. print(f" Generating visualization for selected molecule {i + 1}...")
  1396. detailed['explainer'].visualize_selected_molecule(
  1397. explanation, smiles,
  1398. target_value=target_value,
  1399. title=f'Selected Molecule {i + 1}',
  1400. save_path=f'selected_molecule_{i + 1}.png'
  1401. )
  1402. except Exception as e:
  1403. print(f" Error generating visualization for molecule {i + 1}: {e}")
  1404. except Exception as e:
  1405. print(f"Error generating visualizations: {e}")
  1406. print(f"\n" + "=" * 70)
  1407. print("Analysis completed!")
  1408. print("=" * 70)
  1409. print(f"Generated files:")
  1410. print(f"- enhanced_feature_importance_analysis_subplot1_atom_importance.png: Average Atom Importance")
  1411. print(f"- enhanced_feature_importance_analysis_subplot2_cumulative_contribution.png: Cumulative contribution")
  1412. print(f"- enhanced_feature_importance_analysis_subplot3_atom_distribution.png: Atom type distribution")
  1413. print(f"- enhanced_feature_importance_analysis_subplot4_importance_distribution.png: Importance distribution")
  1414. print(f"- substructure_analysis_importance_summary_subplot1_frequency.png: Substructure frequency")
  1415. print(f"- substructure_analysis_importance_summary_subplot2_importance.png: Substructure importance")
  1416. print(f"- substructure_analysis_importance_summary_subplot3_functional_groups.png: Functional group distribution")
  1417. print(f"- substructure_analysis_importance_summary_subplot4_scatter.png: Importance vs frequency scatter")
  1418. print(
  1419. f"- substructure_analysis_highlighted_molecules.png: Highlighted substructure molecules (Full Dataset: y > 6, Importance > 0.5)")
  1420. print(f"- substructure_analysis_heatmap.png: Substructure-molecule heatmap (Top 40 molecules)")
  1421. print(
  1422. f"- selected_molecule_1.png to selected_molecule_6.png: Selected molecules (Full Dataset: y > 6, Importance > 0.5, True y only, No predict values)")
  1423. def display_analysis_results(results):
  1424. if results is None:
  1425. print("No analysis results to display")
  1426. return
  1427. print("\n" + "=" * 60)
  1428. print("Detailed Analysis Results View")
  1429. print("=" * 60)
  1430. quick_df = pd.DataFrame(results['quick_analysis_results'])
  1431. print(f"\n=== Global prediction distribution (based on {len(quick_df)} molecules) ===")
  1432. print(quick_df['prediction'].describe())
  1433. print(f"\n=== Importance distribution ===")
  1434. print(quick_df['avg_importance'].describe())
  1435. substructures = results['detailed_analysis_results']['important_substructures']
  1436. full_dataset_substructures = results['detailed_analysis_results']['full_dataset_substructures']
  1437. print(f'\n=== Important substructure analysis ===')
  1438. print(f'Found {len(substructures)} molecules containing important substructures (from sample)')
  1439. print(f'Found {len(full_dataset_substructures)} molecules containing important substructures (from full dataset)')
  1440. all_substructure_types = {}
  1441. for struct in substructures:
  1442. for sub_name, sub_info in struct['known_substructures'].items():
  1443. if sub_name not in all_substructure_types:
  1444. all_substructure_types[sub_name] = 0
  1445. all_substructure_types[sub_name] += len(sub_info['matches'])
  1446. if all_substructure_types:
  1447. print(f"\nMost common chemical substructures:")
  1448. sorted_subs = sorted(all_substructure_types.items(), key=lambda x: x[1], reverse=True)
  1449. for sub_name, count in sorted_subs[:10]:
  1450. print(f" {sub_name}: {count} times")
  1451. importance_df = results['detailed_analysis_results']['importance_df']
  1452. if not importance_df.empty:
  1453. print(f"\n=== Atom importance analysis ===")
  1454. atom_stats = importance_df.groupby('atom_type')['importance'].agg(['mean', 'std', 'count'])
  1455. atom_stats = atom_stats.sort_values('mean', ascending=False)
  1456. print("Atom type importance ranking:")
  1457. for atom_type, stats in atom_stats.head(8).iterrows():
  1458. print(
  1459. f" {atom_type}: Average importance {stats['mean']:.3f} (±{stats['std']:.3f}), appeared {stats['count']} times")
  1460. print(f"\n=== Specific molecule examples (from full dataset) ===")
  1461. qualified_count = 0
  1462. for struct in full_dataset_substructures[:10]:
  1463. if (struct.get('target_value', 0) > 6 and
  1464. any(max(sub_info['importance_scores']) > 0.5
  1465. for sub_info in struct['known_substructures'].values()
  1466. if sub_info['importance_scores'])):
  1467. qualified_count += 1
  1468. print(f"\nQualified Molecule {qualified_count}:")
  1469. print(f" SMILES: {struct['smiles'][:60]}...")
  1470. print(f" Prediction: {struct['prediction']:.4f}")
  1471. print(f" True y value: {struct.get('target_value', 'N/A')}")
  1472. print(f" Important atoms: {struct['num_important_atoms']}")
  1473. if struct['known_substructures']:
  1474. print(f" Found substructures:")
  1475. for sub_name, sub_info in list(struct['known_substructures'].items())[:3]:
  1476. avg_imp = np.mean(sub_info['importance_scores']) if sub_info['importance_scores'] else 0
  1477. print(f" • {sub_name}: average importance {avg_imp:.3f}")
  1478. if qualified_count >= 3:
  1479. break
  1480. print(
  1481. f"\nTotal qualified molecules found in full dataset: {sum(1 for struct in full_dataset_substructures if struct.get('target_value', 0) > 6 and any(max(sub_info['importance_scores']) > 0.5 for sub_info in struct['known_substructures'].values() if sub_info['importance_scores']))}")
  1482. if __name__ == "__main__":
  1483. results = hybrid_analysis_strategy(
  1484. test_csv_file='D:\\pycharm\\gutingle\\pythonProject2\\突出核蛋白\\test_data.csv',
  1485. model_path='D:\\pycharm\\gutingle\\pythonProject2\\突出核蛋白\\best_model.pth',
  1486. target_detailed_count=200,
  1487. importance_threshold=0.3
  1488. )
  1489. if results is not None:
  1490. print("\nAnalysis results saved in 'results' variable")
  1491. print("You can access them through:")
  1492. print("- results['global_statistics']: Global statistics")
  1493. print("- results['quick_analysis_results']: Quick analysis results")
  1494. print("- results['detailed_analysis_results']: Detailed analysis results")
  1495. print("- results['summary']: Analysis summary")
  1496. display_analysis_results(results)
  1497. print("\nYou can use the following commands for further analysis:")
  1498. print("# View specific atom type importance")
  1499. print("importance_df = results['detailed_analysis_results']['importance_df']")
  1500. print("print(importance_df[importance_df['atom_type'] == 'N'].describe())")
  1501. print("\n# View molecules with highest predictions")
  1502. print("quick_df = pd.DataFrame(results['quick_analysis_results'])")
  1503. print("top_predictions = quick_df.nlargest(10, 'prediction')")
  1504. print("print(top_predictions[['smiles', 'prediction', 'avg_importance']])")

gnnexplainer.py at commit 76ec685, no license · at the source

Overview

Authors: Tingle Gu1, Zixu Ran2, Wenyin Li3, Xudong Guo2, Bo Li4, Fuyi Li2,5, Cangzhi Jia1
ORCID iDs: Fuyi Li, Cangzhi Jia
  1. School of Science, Dalian Maritime University, No. 1 Linghai Road, Dalian 116026, Liao Ning, China
  2. College of Information Engineering, Northwest A&F University, No. 3 Taicheng Road, Yangling 712100, Shanxi, China
  3. The First Clinical College, Liaoning University of Traditional Chinese Medicine, No. 79, Chongshen East Road, Huanggu District, Shenyang 110847, Liaoning, China
  4. Department of Dermatology, Dalian Dermatosis Hospital, ChangJiang Road 788, Dalian 116021, Liaoning, China
  5. South Australian immunoGENomics Cancer Institute (SAiGENCI), College of Health, Adelaide University, Adelaide 5005, Australia
Journal: Briefings in bioinformatics, volume 27, issue 2, article bbag118
Dates: received 5 November 2025; accepted 20 February 2026; published online 23 March 2026; in print March 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1093/bib/bbag118 · PMID 41870129 · PMCID PMC13006971 · OpenAlex W7140079345
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: computational modeling (no new data) (modality), human (organism), Parkinson's (population), cellular / molecular (subfield)
Methods: Machine learning
Keywords: α-synuclein, composite regularization, QSAR, graph contextual attention, dual-channel feature fusion
MeSH: alpha-Synuclein*, Drug Evaluation, Preclinical, Graph Neural Networks, Humans, Ligands, Molecular Docking Simulation (* major topic)
Topic: Computational Drug Discovery Methods (Computational Theory and Mathematics, Computer Science), according to OpenAlex
Funding: National Natural Science Foundation of China (62071079, 62202388); Hainan Normal University, Ministry of Education (JSKX202203)
Citations: cited by 1 paper (Europe PMC); 24 references in the paper

Abstract

The pathological aggregation of α-synuclein (α-syn) constitutes a pivotal hallmark in the progression of neurodegenerative disorders, including Parkinson’s disease, underscoring the imperative need for identifying site-specific ligands. This study presents, for the first time, an advanced deep learning framework specifically designed for the prediction of molecular properties associated with α-syn. The framework integrates graph-based contextual attention mechanisms, structural feature aggregation protocols, and dual-channel feature integration, complemented by a composite regularization strategy that synergizes mean squared error minimization, Kullback–Leibler divergence–induced latent space regularization, and L2 norm penalization, thereby delivering outstanding predictive accuracy on the independent test dataset with MSE of 0.1812. Mechanistic insights derived from GNNExplainer analysis and molecular docking studies (PDB: 6A6B) elucidated that aromatic ring systems (benzene ring significance: 0.737) and hydrogen bond donor groups (amino group significance: 0.438) play critical roles in mediating high-affinity ligand–receptor interactions through π–π stacking within the hydrophobic pocket formed by Val82 and Ala89 residues, as well as directed hydrogen bonding involving catalytic residues Ser42 and Lys45. These findings not only enhance the understanding of inhibitor mechanisms but also establish a novel framework for the preliminary screening of small-molecule therapeutics, thereby laying a rigorous groundwork for structure-guided drug optimization and rational molecular design.

Reproduced under the paper's license (CC BY), from the paper cited above.

Repository

Its files are read in the Code ↔ Paper reader above, with 6 matches between paragraphs and lines of code.

JiaCZ-Computational-Biology/M-GAT-GraphSAGE

License: none: the authors keep all their rights
State: the link answers, verified on 30 September 2026
Evidence: files inventoried
Commit: 76ec685148bad7b3a0362e816cc24b70f4946ddc, 4 November 2025
Languages: Python (43)
Size: 44 files, 43 scripts
Software Heritage: not archived
Found in: “Data availability”
Holds: README
Not found: license file, CITATION.cff, environment file, tests, continuous integration, documentation
Tools: NumPy (43 files), pandas (43 files), RDKit (43 files), PyTorch Geometric (41 files), PyTorch (41 files), scikit-learn (27 files), SciPy (17 files), Matplotlib (3 files), seaborn (2 files), LightGBM (1 file), NetworkX (1 file), Pillow (1 file), XGBoost (1 file)
Availability: 1 check, the latest on 30 September 2026: the link answers
  • 30 September 2026: the link answers
44 files

The paper's code and data availability statement is in the Data section.

Tracing map

Proposed by the machine: these links were found in the paper and verified at the source, without human review. The map will receive a Zenodo DOI once one of the paper's authors has validated it with their ORCID.

What the map holds:

  • 1 repository of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
  • 43 scripts, each with its path and the digest of its content;
  • 6 matches between paragraphs of the paper and lines of the code (method lexical-v1);
  • neither the text of the paper nor the code itself.

Its JSON (tracing-map.json) is deposited on Zenodo with its DOI once the map is validated.

Data

No dataset and no data link were found in the paper.

Data availability

The source code and the package are available at https://github.com/JiaCZ-Computational-Biology/M-GAT-GraphSAGE.

Reproduced under the paper's license (CC BY), from the paper cited above.

Versions

The history of this record: each version stored by the harvester or made by a correction of its authors or of the maintainers of its code, and what changed in its facts. The texts of the paper (its abstract, its availability statements) are not part of it; versions that changed only those are not listed.

Version 1, 30 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 7 authors, 5 keywords, 6 MeSH terms, 2 funders, 21 references.

Cite

This paper

Gu, T., Ran, Z., Li, W., Guo, X., Li, B., Li, F., & Jia, C. (2026). Drug screening for α-synuclein aggregation inhibitors via multimodal graph neural network. Briefings in bioinformatics, 27(2), bbag118. https://doi.org/10.1093/bib/bbag118

BibTeX

@article{gu2026drug,
author = {Gu, Tingle and Ran, Zixu and Li, Wenyin and Guo, Xudong and Li, Bo and Li, Fuyi and Jia, Cangzhi},
title = {{Drug screening for α-synuclein aggregation inhibitors via multimodal graph neural network}},
journal = {Briefings in bioinformatics},
year = {2026},
month = mar,
volume = {27},
number = {2},
pages = {bbag118},
publisher = {Oxford University Press},
issn = {1467-5463},
doi = {10.1093/bib/bbag118},
url = {https://doi.org/10.1093/bib/bbag118},
pmid = {41870129},
pmcid = {PMC13006971}
}

RIS

TY - JOUR
AU - Gu, Tingle
AU - Ran, Zixu
AU - Li, Wenyin
AU - Guo, Xudong
AU - Li, Bo
AU - Li, Fuyi
AU - Jia, Cangzhi
TI - Drug screening for α-synuclein aggregation inhibitors via multimodal graph neural network
T2 - Briefings in bioinformatics
J2 - Brief Bioinform
PY - 2026
DA - 2026/03/01
VL - 27
IS - 2
SP - bbag118
SN - 1467-5463
PB - Oxford University Press
DO - 10.1093/bib/bbag118
UR - https://doi.org/10.1093/bib/bbag118
LA - en
ER -

CSL-JSON

{
"id": "10.1093/bib/bbag118",
"type": "article-journal",
"title": "Drug screening for α-synuclein aggregation inhibitors via multimodal graph neural network",
"container-title": "Briefings in bioinformatics",
"author": [
{
"family": "Gu",
"given": "Tingle"
},
{
"family": "Ran",
"given": "Zixu"
},
{
"family": "Li",
"given": "Wenyin"
},
{
"family": "Guo",
"given": "Xudong"
},
{
"family": "Li",
"given": "Bo"
},
{
"family": "Li",
"given": "Fuyi"
},
{
"family": "Jia",
"given": "Cangzhi"
}
],
"container-title-short": "Brief Bioinform",
"volume": "27",
"issue": "2",
"page": "bbag118",
"DOI": "10.1093/bib/bbag118",
"PMID": "41870129",
"PMCID": "PMC13006971",
"ISSN": "1467-5463",
"publisher": "Oxford University Press",
"URL": "https://doi.org/10.1093/bib/bbag118",
"language": "en",
"issued": {
"date-parts": [
[
2026,
3,
1
]
]
}
}

The tracing map gets a citation of its own once an author has validated it and it has a DOI.

Similar papers

The papers with a page that share the most with this one: the tools found in their code, their categories, datasets, cited references and authors, the rarest counting most.

[1] doi:10.1038/s42003-026-10957-8 [code]
Brain defence by the extracellular matrix protein Cochlin.
Journal: Communications biology
In common: RDKit, XGBoost, NetworkX, 8 other tools, cellular / molecular
[2] doi:10.7554/elife.110588 [code]
Opening the black box toward a modular approach to spike sorting.
Journal: eLife
In common: LightGBM, XGBoost, NetworkX, 8 other tools
[3] doi:10.1016/j.isci.2026.116055 [code]
Mapping the transcriptional diversity of calcium signaling in the mouse and human brain.
Journal: iScience
In common: LightGBM, PyTorch Geometric, NetworkX, 7 other tools
[4] doi:10.1371/journal.pone.0345854 [code]
Shedding light on neural learning to rank models for anticancer drug prioritization.
Journal: PloS one
In common: RDKit, PyTorch Geometric, NetworkX, 7 other tools
[5] doi:10.1093/nar/gkag706 [code]
scDifformer: diffusion-based post-training for virtual cell modeling across large-scale single-cell data.
Journal: Nucleic acids research
In common: RDKit, PyTorch Geometric, NetworkX, 7 other tools
[6] doi:10.1186/s13244-026-02365-7 [code]
Super-resolution MRI and 2.5D deep learning for intratumoral-peritumoral radiomics in preoperative prediction of rectal cancer perineural invasion.
Journal: Insights into imaging
In common: LightGBM, XGBoost, Pillow, 7 other tools
[7] doi:10.1038/s41586-026-10670-w [code]
Zero-shot design of drug-binding proteins via neural iterative selection-expansion.
Journal: Nature
In common: RDKit, PyTorch Geometric, PyTorch, 6 other tools, computational modeling (no new data), cellular / molecular
[8] doi:10.1038/s41398-026-03965-z [code]
Disentangling individual heterogeneity reveals robust network and molecular signatures of major depressive disorder with suicidal ideation.
Journal: Translational psychiatry
In common: PyTorch Geometric, NetworkX, Pillow, 7 other tools, cellular / molecular
[9] doi:10.1002/hbm.70469 [code]
VarCoNet: A Variability-Aware Self-Supervised Framework for Functional Connectome Extraction From Resting-State fMRI.
Journal: Human brain mapping
In common: PyTorch Geometric, NetworkX, Pillow, 7 other tools
[10] doi:10.1371/journal.pcbi.1014555 [code]
Body surface potential driven personalisation of electrophysiological digital twins in hypertrophic cardiomyopathy.
Journal: PLoS computational biology
In common: XGBoost, NetworkX, Pillow, 7 other tools

Contribute

The authors of this paper can claim it, correct its record and validate its tracing map, and the maintainers of its code (its owner, or a public member of its organization) correct what it says of their repository; anyone signed in can ask for its removal. Every request goes to OSCR's own machine, which answers it; your account page follows them.

Sign in with ORCID to claim this paper as one of its authors, correct its record or validate its tracing map: when the paper's metadata lists your ORCID iD, you are recognized at once. Maintainers of its code: sign in with GitHub, then claim the repository on your account page.

Request its removal

To ask OSCR to remove this record, the copies of its authors' scripts or its tracing map, use the removal request page: signed in, you say who you are, what to remove and why, then review and confirm the request. Published rules decide every request (how).

Discussion, reproductions, activity

Discussion: questions and error reports about this paper and its code, from signed-in readers and its authors. It opens with sign-in.

Reproductions: reports from readers who ran the authors' code: what they reproduced, with which environment, commit and data. It opens with sign-in.

Activity: what happens around this paper: new versions of its record, its map's validation, discussions and reproductions. It opens with sign-in.