OSCR

HYG-mol: An Interpretable Multimodal Hypergraph Framework for Molecular Property Prediction.

Code ↔ Paper

11 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 11 matches
  1. [1] § Results › Experimental setting ↔ src/main.py, lines 20–74 · score 0.78 · FreeSolv, scaffold splitting, random seeds, stratified, threshold, QM7
  2. [2] § Methods › Molecular hypergraph construction and representation › Hyperedge representation ↔ src/data/dataset.py, lines 627–719 · score 0.67 · substructure matching, SMARTS patterns, alerts, amides, amines, esters
  3. [3] § Methods › Molecular hypergraph construction and representation › Multimodal hypernode representation ↔ src/data/dataset.py, lines 137–199 · score 0.66 · Gasteiger charge, formal charge, electrons, traditional, aromaticity, ChemBERTa
  4. [4] § Methods › Molecular hypergraph construction and representation › Multimodal hypernode representation ↔ src/data/dataset.py, lines 256–316 · score 0.63 · tokens mapped, atom mapping, ChemBERTa model, tokenized, SMILES, mol
  5. [5] § Results › Experimental setting ↔ src/utils/utils.py, lines 20–79 · score 0.62 · Murcko scaffold splitting, class distribution, threshold, subsets, seeds, 10 %
  6. [6] § Results › Outlier analysis ↔ src/data/dataset.py, lines 627–719 · score 0.60 · c1ccc2c, OC, cc1, Cl, halogen, motifs
  7. [7] § Methods › Molecular hypergraph construction and representation › Hyperedge representation ↔ src/data/dataset.py, lines 34–135 · score 0.59 · isolated atoms, ring connections, RDKit, Node features, substructures, hyperedges
  8. [8] § Results › Interpretability analysis ↔ src/data/dataset.py, lines 759–896 · score 0.56 · carbonyl, donor, bridging, hydrophobic, acceptor, amide
  9. [9] § Results › Experimental setting ↔ src/training/trainer.py, lines 275–389 · score 0.56 · squared error, absolute error, MAE, RMSE, regression, metrics
  10. [10] § Results › Comparison experiments ↔ src/main.py, lines 20–74 · score 0.55 · FreeSolv, random seeds, QM7, Lipophilicity, QM8, regression
  11. [11] § Methods ↔ src/data/dataset.py, lines 34–135 · score 0.51 · ring connections, RDKit, matching, substructures, nodes, Atoms

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,787 lines · 72 KB · no license · 7 matches

  1. # src/data/dataset.py
  2. import os
  3. import warnings
  4. import numpy as np
  5. import pandas as pd
  6. import torch
  7. from .. import config
  8. from torch_geometric.data import Data, Dataset
  9. from tqdm import tqdm
  10. from rdkit import Chem
  11. from rdkit.Chem import Draw, AllChem, rdMolDescriptors, Descriptors, Crippen, MolSurf
  12. from rdkit.Chem.Scaffolds import MurckoScaffold
  13. from transformers import AutoTokenizer, AutoModel
  14. from rdkit import RDLogger
  15. RDLogger.DisableLog('rdApp.*')
  16. class MoleculeHypergraph:
  17. def __init__(self, feature_type):
  18. self.feature_type = feature_type
  19. self.node_features = []
  20. self.hyperedge_index = None
  21. self.mol = None
  22. self.hyperedges = []
  23. self.hyperedge_labels = []
  24. if self.feature_type != "traditional_only":
  25. self.tokenizer = AutoTokenizer.from_pretrained(config.CHEMBERTA_MODEL_NAME)
  26. self.chemberta_model = AutoModel.from_pretrained(config.CHEMBERTA_MODEL_NAME).to('cpu')
  27. else:
  28. self.tokenizer = None
  29. self.chemberta_model = None
  30. def build_from_mol(self, mol):
  31. self.mol = mol
  32. num_atoms = mol.GetNumAtoms()
  33. functional_groups, fg_labels = self.get_functional_groups(mol)
  34. rings, ring_labels = self.get_ring_structures(mol)
  35. special_structures, special_labels = self.get_special_structures(mol)
  36. hyperedges = []
  37. hyperedge_labels = []
  38. def add_valid_hyperedge(edge, label):
  39. valid_nodes = [node_id for node_id in edge if 0 <= node_id < num_atoms]
  40. if valid_nodes:
  41. is_duplicate = False
  42. sorted_valid_nodes_tuple = tuple(sorted(valid_nodes))
  43. for existing_edge_tuple, existing_label in zip(map(lambda e: tuple(sorted(e)), hyperedges),
  44. hyperedge_labels):
  45. if existing_edge_tuple == sorted_valid_nodes_tuple and existing_label == label:
  46. is_duplicate = True
  47. break
  48. if not is_duplicate:
  49. hyperedges.append(valid_nodes)
  50. hyperedge_labels.append(label)
  51. for edge, label in zip(functional_groups, fg_labels):
  52. add_valid_hyperedge(edge, label)
  53. for edge, label in zip(rings, ring_labels):
  54. add_valid_hyperedge(edge, label)
  55. for edge, label in zip(special_structures, special_labels):
  56. add_valid_hyperedge(edge, label)
  57. ring_atoms = set()
  58. atom_to_ring = {}
  59. for ring_id, ring_node_indices in enumerate(rings):
  60. ring_atoms.update(ring_node_indices)
  61. for atom_idx in ring_node_indices:
  62. atom_to_ring.setdefault(atom_idx, set()).add(ring_id)
  63. for bond in mol.GetBonds():
  64. begin_idx = bond.GetBeginAtom().GetIdx()
  65. end_idx = bond.GetEndAtom().GetIdx()
  66. if begin_idx >= num_atoms or end_idx >= num_atoms:
  67. continue
  68. begin_in_ring = begin_idx in ring_atoms
  69. end_in_ring = end_idx in ring_atoms
  70. if begin_in_ring and end_in_ring:
  71. common_rings = atom_to_ring.get(begin_idx, set()) & atom_to_ring.get(end_idx, set())
  72. if not common_rings:
  73. add_valid_hyperedge([begin_idx, end_idx], 'RingConnector')
  74. elif begin_in_ring != end_in_ring:
  75. add_valid_hyperedge([begin_idx, end_idx], 'RingSubstituent')
  76. all_atoms_in_hyperedges = set()
  77. for he in hyperedges:
  78. all_atoms_in_hyperedges.update(he)
  79. excluded_atoms = set(range(num_atoms)) - all_atoms_in_hyperedges
  80. for atom_idx in list(excluded_atoms):
  81. atom = mol.GetAtomWithIdx(atom_idx)
  82. neighbor_indices = [n.GetIdx() for n in atom.GetNeighbors() if n.GetIdx() < num_atoms]
  83. if neighbor_indices:
  84. current_group = sorted([atom_idx] + neighbor_indices)
  85. is_new_group = True
  86. for existing_edge, existing_label in zip(hyperedges, hyperedge_labels):
  87. if existing_label == 'IsolatedAtomGroup' and sorted(existing_edge) == current_group:
  88. is_new_group = False
  89. break
  90. if is_new_group:
  91. add_valid_hyperedge(current_group, 'IsolatedAtomGroup')
  92. else:
  93. add_valid_hyperedge([atom_idx], 'SingleAtom')
  94. try:
  95. murcko_mol_full = MurckoScaffold.GetScaffoldForMol(mol)
  96. if murcko_mol_full.GetNumAtoms() > 0:
  97. match_indices_full = mol.GetSubstructMatch(murcko_mol_full)
  98. if match_indices_full:
  99. add_valid_hyperedge(list(match_indices_full), 'MurckoScaffoldFull')
  100. if murcko_mol_full.GetNumAtoms() > 0:
  101. murcko_mol_generic = MurckoScaffold.MakeScaffoldGeneric(murcko_mol_full)
  102. if murcko_mol_generic.GetNumAtoms() > 0:
  103. match_indices_generic = mol.GetSubstructMatch(murcko_mol_generic)
  104. if match_indices_generic:
  105. add_valid_hyperedge(list(match_indices_generic), 'MurckoScaffoldCore')
  106. except ImportError:
  107. print(
  108. "Warning: RDKit MurckoScaffold module not available. Skipping Murcko scaffold hyperedges.")
  109. except Exception as e:
  110. print(f"Warning: Error processing Murcko scaffold for molecule {Chem.MolToSmiles(mol)}: {e}")
  111. self.build_node_features(mol)
  112. self.build_hyperedge_index(hyperedges, hyperedge_labels)
  113. return self
  114. def build_node_features(self, mol):
  115. num_atoms = mol.GetNumAtoms()
  116. model_max_length = 512
  117. if hasattr(self.chemberta_model, 'config') and hasattr(self.chemberta_model.config, 'max_position_embeddings'):
  118. model_max_length = self.chemberta_model.config.max_position_embeddings
  119. mol_id = Chem.MolToSmiles(mol)
  120. cache_key = f"{self.feature_type}_{mol_id}"
  121. if hasattr(self, 'feature_cache') and cache_key in self.feature_cache:
  122. self.node_features = self.feature_cache[cache_key]
  123. return self.node_features
  124. if not hasattr(self, 'feature_cache'):
  125. self.feature_cache = {}
  126. if self.feature_type == "traditional_only":
  127. additional_features = []
  128. mol_weight = Descriptors.MolWt(mol) / 500.0 if hasattr(Descriptors, 'MolWt') else 0.0
  129. logp = Crippen.MolLogP(mol) / 10.0 if hasattr(Crippen, 'MolLogP') else 0.0
  130. tpsa = MolSurf.TPSA(mol) / 100.0 if hasattr(MolSurf, 'TPSA') else 0.0
  131. num_rings = len(mol.GetRingInfo().AtomRings()) / 10.0 if hasattr(mol, 'GetRingInfo') else 0.0
  132. charges = [0.0] * num_atoms
  133. try:
  134. AllChem.ComputeGasteigerCharges(mol)
  135. charges = [atom.GetDoubleProp('_GasteigerCharge')
  136. if atom.HasProp('_GasteigerCharge') else 0.0
  137. for atom in mol.GetAtoms()]
  138. charges = [0.0 if (c is None or np.isnan(c) or np.isinf(c)) else c for c in charges]
  139. except:
  140. pass
  141. ring_info = mol.GetRingInfo()
  142. for atom_idx in range(num_atoms):
  143. atom = mol.GetAtomWithIdx(atom_idx)
  144. neighbors = [n.GetIdx() for n in atom.GetNeighbors()]
  145. base_features = [
  146. atom.GetAtomicNum() / 100.0,
  147. atom.GetDegree() / 4.0,
  148. int(atom.GetIsAromatic()),
  149. atom.GetFormalCharge() / 8.0,
  150. atom.GetNumRadicalElectrons() / 8.0,
  151. atom.GetChiralTag() / 10.0,
  152. atom.GetHybridization() / 6.0,
  153. atom.GetImplicitValence() / 8.0,
  154. atom.IsInRing() * 1.0,
  155. len(neighbors) / 8.0,
  156. ]
  157. hybridization_features = [
  158. int(str(atom.GetHybridization()) == "SP"),
  159. int(str(atom.GetHybridization()) == "SP2"),
  160. int(str(atom.GetHybridization()) == "SP3"),
  161. ]
  162. atom_type_features = [
  163. int(atom.GetAtomicNum() == 1), # H
  164. int(atom.GetAtomicNum() == 6), # C
  165. int(atom.GetAtomicNum() == 7), # N
  166. int(atom.GetAtomicNum() == 8), # O
  167. int(atom.GetAtomicNum() == 9), # F
  168. int(atom.GetAtomicNum() == 15), # P
  169. int(atom.GetAtomicNum() == 16), # S
  170. int(atom.GetAtomicNum() == 17), # Cl
  171. int(atom.GetAtomicNum() == 35), # Br
  172. int(atom.GetAtomicNum() == 53), # I
  173. ]
  174. bond_types = [0, 0, 0, 0]
  175. for neighbor in neighbors:
  176. bond = mol.GetBondBetweenAtoms(atom_idx, neighbor)
  177. if bond.GetBondType() == Chem.rdchem.BondType.SINGLE:
  178. bond_types[0] += 1
  179. elif bond.GetBondType() == Chem.rdchem.BondType.DOUBLE:
  180. bond_types[1] += 1
  181. elif bond.GetBondType() == Chem.rdchem.BondType.TRIPLE:
  182. bond_types[2] += 1
  183. elif bond.GetBondType() == Chem.rdchem.BondType.AROMATIC:
  184. bond_types[3] += 1
  185. if sum(bond_types) > 0:
  186. bond_types = [b / sum(bond_types) for b in bond_types]
  187. neighbor_features = []
  188. if neighbors:
  189. neighbor_atomic_nums = [mol.GetAtomWithIdx(n).GetAtomicNum() for n in neighbors]
  190. neighbor_features = [
  191. sum(1 for n in neighbor_atomic_nums if n == 6) / len(neighbors),
  192. sum(1 for n in neighbor_atomic_nums if n == 7) / len(neighbors),
  193. sum(1 for n in neighbor_atomic_nums if n == 8) / len(neighbors),
  194. sum(1 for n in neighbor_atomic_nums if n == 9 or n == 17 or n == 35 or n == 53) / len(neighbors)
  195. ]
  196. else:
  197. neighbor_features = [0.0] * 4
  198. additional_chemical_features = [
  199. charges[atom_idx],
  200. atom.GetTotalNumHs() / 4.0,
  201. int(ring_info.IsAtomInRingOfSize(atom_idx, 3)),
  202. int(ring_info.IsAtomInRingOfSize(atom_idx, 5)),
  203. int(ring_info.IsAtomInRingOfSize(atom_idx, 6)),
  204. ]
  205. global_mol_features = [mol_weight, logp, tpsa, num_rings]
  206. atom_features = (
  207. base_features +
  208. hybridization_features +
  209. atom_type_features +
  210. bond_types +
  211. neighbor_features +
  212. additional_chemical_features +
  213. global_mol_features
  214. )
  215. additional_features.append(atom_features)
  216. self.node_features = np.array(additional_features, dtype=np.float32)
  217. elif self.feature_type == "chemberta_only":
  218. original_smiles = Chem.MolToSmiles(mol)
  219. atom_mapped_mol = Chem.AddHs(mol)
  220. for atom in atom_mapped_mol.GetAtoms():
  221. atom.SetAtomMapNum(atom.GetIdx() + 1)
  222. inputs = self.tokenizer(
  223. original_smiles,
  224. return_tensors="pt",
  225. truncation=True,
  226. max_length=model_max_length,
  227. padding="max_length"
  228. )
  229. inputs = {k: v.to('cpu') for k, v in inputs.items()}
  230. with torch.no_grad():
  231. outputs = self.chemberta_model(**inputs)
  232. hidden_states = outputs.last_hidden_state.squeeze(0)
  233. tokens = self.tokenizer.convert_ids_to_tokens(inputs['input_ids'][0])
  234. token_to_char = {}
  235. smiles_tokens = self.tokenizer.tokenize(original_smiles)
  236. current_pos = 0
  237. for i, token in enumerate(smiles_tokens):
  238. token_to_char[i] = current_pos
  239. clean_token = token.replace('#', '').replace('Ġ', '')
  240. current_pos += len(clean_token)
  241. token_to_atom_map = {}
  242. current_atom_idx = 0
  243. atom_to_token_map = {}
  244. for token_idx, token in enumerate(tokens):
  245. if token == self.tokenizer.cls_token or \
  246. token == self.tokenizer.sep_token or \
  247. token == self.tokenizer.pad_token:
  248. continue
  249. if token.startswith('##'):
  250. if token_idx > 0 and token_idx - 1 in token_to_atom_map:
  251. token_to_atom_map[token_idx] = token_to_atom_map[token_idx - 1]
  252. continue
  253. atom_tokens = ['C', 'c', 'N', 'n', 'O', 'o', 'S', 's', 'P', 'p', 'F', 'Cl', 'Br', 'I',
  254. '[cH]', '[nH]', '[oH]', '[sH]', '[CH]', '[CH2]', '[CH3]', '[NH]', '[NH2]',
  255. '[OH]', '[SH]', '[PH]']
  256. if any(token == t or token.startswith(t) for t in atom_tokens):
  257. if current_atom_idx < num_atoms:
  258. token_to_atom_map[token_idx] = current_atom_idx
  259. if current_atom_idx not in atom_to_token_map:
  260. atom_to_token_map[current_atom_idx] = []
  261. atom_to_token_map[current_atom_idx].append(token_idx)
  262. current_atom_idx += 1
  263. else:
  264. break
  265. feature_dim = hidden_states.shape[1]
  266. atom_features = torch.zeros((num_atoms, feature_dim), device='cpu')
  267. mapped_atom_indices = set()
  268. for atom_idx, token_indices in atom_to_token_map.items():
  269. if atom_idx < num_atoms:
  270. valid_tokens = [ti for ti in token_indices if ti < hidden_states.shape[0]]
  271. if valid_tokens:
  272. weights = torch.ones(len(valid_tokens))
  273. for i, ti in enumerate(valid_tokens):
  274. if not tokens[ti].startswith('##'):
  275. weights[i] = 2.0
  276. weights = weights / weights.sum()
  277. for i, ti in enumerate(valid_tokens):
  278. atom_features[atom_idx] += weights[i] * hidden_states[ti]
  279. mapped_atom_indices.add(atom_idx)
  280. else:
  281. print(f"Warning: No valid token indices for atom {atom_idx}")
  282. for atom_idx in range(num_atoms):
  283. if atom_idx not in mapped_atom_indices:
  284. atom = mol.GetAtomWithIdx(atom_idx)
  285. neighbors = [n.GetIdx() for n in atom.GetNeighbors() if n.GetIdx() in mapped_atom_indices]
  286. if neighbors:
  287. neighbor_features = atom_features[neighbors]
  288. atom_features[atom_idx] = torch.mean(neighbor_features, dim=0)
  289. mapped_atom_indices.add(atom_idx)
  290. unmapped_atoms = [i for i in range(num_atoms) if i not in mapped_atom_indices]
  291. if unmapped_atoms:
  292. valid_token_indices = [
  293. idx for idx, t in enumerate(tokens)
  294. if t != self.tokenizer.cls_token and \
  295. t != self.tokenizer.sep_token and \
  296. t != self.tokenizer.pad_token and \
  297. idx < hidden_states.shape[0]
  298. ]
  299. if valid_token_indices:
  300. global_features = torch.mean(hidden_states[valid_token_indices], dim=0)
  301. for atom_idx in unmapped_atoms:
  302. atom_features[atom_idx] = global_features
  303. else:
  304. print(f"Warning: No valid token features for global average in SMILES: {original_smiles[:50]}...")
  305. self.node_features = atom_features.cpu().numpy()
  306. else:
  307. original_smiles = Chem.MolToSmiles(mol)
  308. atom_mapped_mol = Chem.AddHs(mol)
  309. for atom in atom_mapped_mol.GetAtoms():
  310. atom.SetAtomMapNum(atom.GetIdx() + 1)
  311. inputs = self.tokenizer(
  312. original_smiles,
  313. return_tensors="pt",
  314. truncation=True,
  315. max_length=model_max_length,
  316. padding="max_length"
  317. )
  318. inputs = {k: v.to('cpu') for k, v in inputs.items()}
  319. with torch.no_grad():
  320. outputs = self.chemberta_model(**inputs)
  321. hidden_states = outputs.last_hidden_state.squeeze(0)
  322. tokens = self.tokenizer.convert_ids_to_tokens(inputs['input_ids'][0])
  323. token_to_atom_map = {}
  324. atom_to_token_map = {}
  325. current_atom_idx = 0
  326. for token_idx, token in enumerate(tokens):
  327. if token == self.tokenizer.cls_token or \
  328. token == self.tokenizer.sep_token or \
  329. token == self.tokenizer.pad_token:
  330. continue
  331. if token.startswith('##'):
  332. if token_idx > 0 and token_idx - 1 in token_to_atom_map:
  333. token_to_atom_map[token_idx] = token_to_atom_map[token_idx - 1]
  334. continue
  335. atom_tokens = ['C', 'c', 'N', 'n', 'O', 'o', 'S', 's', 'P', 'p', 'F', 'Cl', 'Br', 'I',
  336. '[cH]', '[nH]', '[oH]', '[sH]']
  337. if any(token == t or token.startswith(t) for t in atom_tokens):
  338. if current_atom_idx < num_atoms:
  339. token_to_atom_map[token_idx] = current_atom_idx
  340. if current_atom_idx not in atom_to_token_map:
  341. atom_to_token_map[current_atom_idx] = []
  342. atom_to_token_map[current_atom_idx].append(token_idx)
  343. current_atom_idx += 1
  344. else:
  345. break
  346. feature_dim = hidden_states.shape[1]
  347. atom_features_tensor = torch.zeros((num_atoms, feature_dim), device='cpu')
  348. mapped_atom_indices = set()
  349. for atom_idx, token_indices in atom_to_token_map.items():
  350. if atom_idx < num_atoms:
  351. valid_tokens = [ti for ti in token_indices if ti < hidden_states.shape[0]]
  352. if valid_tokens:
  353. for ti in valid_tokens:
  354. atom_features_tensor[atom_idx] += hidden_states[ti]
  355. atom_features_tensor[atom_idx] /= len(valid_tokens)
  356. mapped_atom_indices.add(atom_idx)
  357. for atom_idx in range(num_atoms):
  358. if atom_idx not in mapped_atom_indices:
  359. atom = mol.GetAtomWithIdx(atom_idx)
  360. neighbors = [n.GetIdx() for n in atom.GetNeighbors() if n.GetIdx() in mapped_atom_indices]
  361. if neighbors:
  362. neighbor_features = atom_features_tensor[neighbors]
  363. atom_features_tensor[atom_idx] = torch.mean(neighbor_features, dim=0)
  364. mapped_atom_indices.add(atom_idx)
  365. else:
  366. valid_token_indices = [
  367. idx for idx, t in enumerate(tokens)
  368. if t != self.tokenizer.cls_token and \
  369. t != self.tokenizer.sep_token and \
  370. t != self.tokenizer.pad_token and \
  371. idx < hidden_states.shape[0]
  372. ]
  373. if valid_token_indices:
  374. atom_features_tensor[atom_idx] = torch.mean(hidden_states[valid_token_indices], dim=0)
  375. else:
  376. print(
  377. f"Warning: No valid token features to average for fallback on atom {atom_idx} in SMILES: {original_smiles[:50]}...")
  378. mol_weight = Descriptors.MolWt(mol) / 500.0 if hasattr(Descriptors, 'MolWt') else 0.0
  379. logp = Crippen.MolLogP(mol) / 10.0 if hasattr(Crippen, 'MolLogP') else 0.0
  380. tpsa = MolSurf.TPSA(mol) / 100.0 if hasattr(MolSurf, 'TPSA') else 0.0
  381. charges = [0.0] * num_atoms
  382. try:
  383. AllChem.ComputeGasteigerCharges(mol)
  384. charges = [atom.GetDoubleProp('_GasteigerCharge')
  385. if atom.HasProp('_GasteigerCharge') else 0.0
  386. for atom in mol.GetAtoms()]
  387. charges = [0.0 if (c is None or np.isnan(c) or np.isinf(c)) else c for c in charges]
  388. except:
  389. pass
  390. additional_features_list = []
  391. for atom_idx in range(num_atoms):
  392. atom = mol.GetAtomWithIdx(atom_idx)
  393. neighbors = [n.GetIdx() for n in atom.GetNeighbors()]
  394. base_features = [
  395. atom.GetAtomicNum() / 100.0,
  396. atom.GetDegree() / 4.0,
  397. int(atom.GetIsAromatic()),
  398. atom.GetFormalCharge() / 8.0,
  399. atom.GetNumRadicalElectrons() / 8.0
  400. ]
  401. enhanced_features = [
  402. charges[atom_idx],
  403. int(atom.IsInRing()),
  404. atom.GetTotalNumHs() / 4.0,
  405. mol_weight,
  406. logp,
  407. tpsa
  408. ]
  409. atom_additional_features = base_features + enhanced_features
  410. additional_features_list.append(atom_additional_features)
  411. additional_features_tensor = torch.tensor(additional_features_list, dtype=torch.float, device='cpu')
  412. bert_weight = 0.6
  413. trad_weight = 0.4
  414. bert_mean = torch.mean(atom_features_tensor, dim=0, keepdim=True)
  415. bert_std = torch.std(atom_features_tensor, dim=0, keepdim=True) + 1e-8
  416. atom_features_tensor = (atom_features_tensor - bert_mean) / bert_std
  417. trad_mean = torch.mean(additional_features_tensor, dim=0, keepdim=True)
  418. trad_std = torch.std(additional_features_tensor, dim=0, keepdim=True) + 1e-8
  419. additional_features_tensor = (additional_features_tensor - trad_mean) / trad_std
  420. atom_features_tensor = bert_weight * atom_features_tensor
  421. additional_features_tensor = trad_weight * additional_features_tensor
  422. final_features = torch.cat([atom_features_tensor, additional_features_tensor], dim=1)
  423. self.node_features = final_features.cpu().numpy()
  424. if hasattr(self, 'use_nonlinear_features') and self.use_nonlinear_features:
  425. self.node_features = self._add_nonlinear_interactions(self.node_features)
  426. self.feature_cache[cache_key] = self.node_features
  427. assert self.node_features.shape[0] == num_atoms, \
  428. f"Feature matrix shape {self.node_features.shape} does not match number of atoms {num_atoms}"
  429. return self.node_features
  430. def _add_nonlinear_interactions(self, features):
  431. if features.shape[1] > 30:
  432. main_features = features[:, :10]
  433. else:
  434. main_features = features
  435. num_samples, num_features = main_features.shape
  436. interactions = []
  437. for i in range(min(5, num_features)):
  438. interactions.append(np.square(main_features[:, i]).reshape(-1, 1))
  439. count = 0
  440. for i in range(min(5, num_features)):
  441. for j in range(i + 1, min(5, num_features)):
  442. if count < 10:
  443. interactions.append((main_features[:, i] * main_features[:, j]).reshape(-1, 1))
  444. count += 1
  445. if interactions:
  446. interaction_features = np.hstack(interactions)
  447. return np.hstack([features, interaction_features])
  448. else:
  449. return features
  450. def build_hyperedge_index(self, hyperedges, hyperedge_labels):
  451. node_idx = []
  452. edge_idx = []
  453. num_nodes = self.mol.GetNumAtoms()
  454. valid_edges = []
  455. valid_labels = []
  456. for edge_id, (hyperedge, label) in enumerate(zip(hyperedges, hyperedge_labels)):
  457. valid_nodes = [node_id for node_id in hyperedge if 0 <= node_id < num_nodes]
  458. if len(valid_nodes) > 0:
  459. if len(valid_nodes) != len(hyperedge):
  460. print(
  461. f"Warning: Hyperedge {edge_id} contains invalid nodes, keeping only valid ones: {valid_nodes}")
  462. valid_edges.append(valid_nodes)
  463. valid_labels.append(label)
  464. else:
  465. print(f"Warning: Hyperedge {edge_id} has no valid nodes, discarding")
  466. if not valid_edges:
  467. print("Warning: No valid hyperedges found! Creating a default edge.")
  468. if num_nodes > 0:
  469. valid_edges = [[0]]
  470. valid_labels = ["Default"]
  471. else:
  472. raise ValueError("No valid hyperedges could be created for molecule with no atoms")
  473. self.hyperedges = valid_edges
  474. self.hyperedge_labels = valid_labels
  475. for edge_id, hyperedge in enumerate(valid_edges):
  476. for node_id in hyperedge:
  477. if 0 <= node_id < num_nodes:
  478. node_idx.append(node_id)
  479. edge_idx.append(edge_id)
  480. if not node_idx or not edge_idx:
  481. print("Warning: Empty hyperedge index after processing. Creating fallback index.")
  482. if num_nodes > 0:
  483. node_idx = [0]
  484. edge_idx = [0]
  485. else:
  486. raise ValueError("Cannot create valid hyperedge index for empty molecule")
  487. self.hyperedge_index = np.array([node_idx, edge_idx])
  488. def get_functional_groups(self, mol):
  489. fg_smarts = {
  490. 'Amine': '[NX3;H2,H1,H0;!$(NC=O)]',
  491. 'QuaternaryAmmonium': '[NX4+]',
  492. 'Alcohol': '[OX2H]',
  493. 'Phenol': 'c[OH]',
  494. 'CarboxylicAcid': 'C(=O)[OX1H0-,OX2H1]',
  495. 'Ester': 'C(=O)O[C;!$(C=O)]',
  496. 'Ketone': 'C(=O)[C;!$(C=O)]',
  497. 'Aldehyde': '[CX3H1](=O)[#6]',
  498. 'Ether': '[OD2]([#6])[#6]',
  499. 'Amide': 'C(=O)N',
  500. 'Halogen': '[F,Cl,Br,I]',
  501. 'Sulfonamide': 'S(=O)(=O)N',
  502. 'Thiol': '[#16X2H]',
  503. 'Disulfide': '[#16X2]-[#16X2]',
  504. 'Nitrate': '[N+](=O)[O-]',
  505. 'Cyano': '[C]#N',
  506. 'Azide': '[N]=[N+]=[N-]',
  507. 'Alkyne': '[CX2]#C',
  508. 'Alkene': 'C=C',
  509. 'Phosphate': '[#15](=O)(O)(O)O',
  510. 'Phosphonate': 'P(=O)(O)(O)[C]',
  511. 'Sulfate': 'S(=O)(=O)(O)(O)',
  512. 'Sulfonate': 'S(=O)(=O)O[C]',
  513. 'Sulfoxide': '[#16X3+1][#6]',
  514. 'SulfonicAcid': 'S(=O)(=O)[OH]',
  515. 'Isocyanate': 'N=C=O',
  516. 'Urea': 'N-C(=O)-N',
  517. 'Carbamate': 'O=C(O)N',
  518. 'Imine': '[CX2]=[NX2]',
  519. 'Thioether': '[#16X2][#6]',
  520. 'Epoxide': '[C@H1]1O[C@H1]1',
  521. 'Peroxide': '[OX2][OX2]',
  522. 'BoronicAcid': 'B(O)O',
  523. 'Anhydride': 'C(=O)OC(=O)',
  524. 'Thiocyanate': '[N-]=C=S',
  525. 'Isothiocyanate': 'N=C=S',
  526. 'Oxime': '[CX3](=NO)',
  527. 'Hydrazone': '[CX3](=NN)',
  528. 'Guanidine': 'N=C(N)N',
  529. 'Pyridine': 'n1ccccc1',
  530. 'Pyrazine': 'n1cnccn1',
  531. 'Pyrrole': 'n1cccc1',
  532. 'Imidazole': 'n1c[nH]cc1',
  533. 'Thiazole': 'c1ncsc1',
  534. 'NitroAromatic': '[N](=O)=O[c]',
  535. 'Quinone': 'O=C1C=CC(=O)C=C1',
  536. 'Piperidine': 'N1CCCCC1',
  537. 'Pyrrolidine': 'N1CCCC1',
  538. 'Morpholine': 'O1CCNCC1',
  539. 'TertiaryButylEster': 'C(=O)OC(C)(C)C',
  540. 'IsopropylEster': 'C(=O)O[CH](C)C',
  541. 'MethylEster': 'C(=O)OC',
  542. 'EthylEster': 'C(=O)OCC',
  543. 'PropylEster': 'C(=O)OCCC',
  544. 'HexylEster': 'C(=O)OCCCCCC',
  545. 'OctylEster': 'C(=O)OCCCCCCCC',
  546. 'Fluorine': '[F]',
  547. 'Trifluoromethyl': 'C(F)(F)F',
  548. 'Adamantane': 'C1C2CC3CC(C1)CC(C2)C3',
  549. 'Cyclohexyl': 'C1CCCCC1',
  550. 'Sulfone': 'S(=O)(=O)[C,N,O]',
  551. 'PhosphineOxide': 'P(=O)[C,N,O]',
  552. 'Benzene': 'c1ccccc1',
  553. 'Naphthalene': 'c1cccc2ccccc12',
  554. 'NitroAromatic': '[$(c-[N+](=O)[O-]),$(c-[N+]-[O-])]',
  555. 'Aniline': 'c-[NX3H2]',
  556. 'MichaelAcceptor_Enone': '[CX3]=[CX3]-[CX3](=[O,S,N])',
  557. 'Epoxide_Alert': '[OD1r3]1[#6r3][#6r3]1',
  558. 'Tetrahydroisoquinoline_Core_Alert': 'c1ccc2c(c1)CCNCC2',
  559. '4_Phenylpiperidine_Motif': 'c1ccccc1-C1CCNCC1',
  560. }
  561. functional_groups = []
  562. fg_labels = []
  563. for fg_name, smarts in fg_smarts.items():
  564. try:
  565. pattern = Chem.MolFromSmarts(smarts)
  566. if pattern is None:
  567. warnings.warn(f"Invalid SMARTS pattern: {fg_name} -> {smarts}")
  568. continue
  569. matches = mol.GetSubstructMatches(pattern)
  570. for match in matches:
  571. fg_atoms = set(match)
  572. functional_groups.append(list(fg_atoms))
  573. fg_labels.append(fg_name)
  574. except Exception as e:
  575. warnings.warn(f"Error in matching SMARTS pattern: {fg_name} -> {smarts}: {e}")
  576. continue
  577. return functional_groups, fg_labels
  578. def get_ring_structures(self, mol):
  579. rings = []
  580. labels = []
  581. ring_info = mol.GetRingInfo()
  582. for idxs in ring_info.AtomRings():
  583. is_aromatic = all([mol.GetAtomWithIdx(idx).GetIsAromatic() for idx in idxs])
  584. ring_size = len(idxs)
  585. ring_label = f"AromaticRing_{ring_size}" if is_aromatic else f"Ring_{ring_size}"
  586. rings.append(list(idxs))
  587. labels.append(ring_label)
  588. fused_rings = rdMolDescriptors.CalcNumSpiroAtoms(mol)
  589. if fused_rings > 0:
  590. labels.append("FusedRings")
  591. rings.append([atom.GetIdx() for atom in mol.GetAtoms() if atom.IsInRing()])
  592. return rings, labels
  593. def get_special_structures(self, mol):
  594. special_structures = []
  595. special_labels = []
  596. metal_atomic_numbers = [13, 12, 20, 26, 30, 29, 25, 24, 27, 28, 33, 80, 82, 50]
  597. metal_atoms = [atom.GetIdx() for atom in mol.GetAtoms() if atom.GetAtomicNum() in metal_atomic_numbers]
  598. if metal_atoms:
  599. special_structures.append(metal_atoms)
  600. special_labels.append("MetalAtoms")
  601. ssr = Chem.GetSymmSSSR(mol)
  602. atom_rings = [set(ring) for ring in ssr]
  603. spiro_atoms = set()
  604. for i in range(len(atom_rings)):
  605. for j in range(i + 1, len(atom_rings)):
  606. shared_atoms = atom_rings[i] & atom_rings[j]
  607. if len(shared_atoms) == 1:
  608. spiro_atoms.update(shared_atoms)
  609. if spiro_atoms:
  610. special_structures.append(list(spiro_atoms))
  611. special_labels.append("SpiroAtoms")
  612. return special_structures, special_labels
  613. def generate_enhanced_hyperedge_attributes(self):
  614. num_hyperedges = len(self.hyperedges)
  615. hyperedge_features = np.zeros((num_hyperedges, 5), dtype=np.float32)
  616. if num_hyperedges == 0:
  617. print(f"Warning: No hyperedges found for molecule. Creating default attribute matrix with shape [1, 5].")
  618. return np.zeros((1, 5), dtype=np.float32)
  619. for i, (edge, label) in enumerate(zip(self.hyperedges, self.hyperedge_labels)):
  620. if i >= hyperedge_features.shape[0]:
  621. print(
  622. f"Warning: Hyperedge index {i} exceeds feature matrix dimension {hyperedge_features.shape[0]}. Expanding matrix.")
  623. expanded_features = np.zeros((i + 1, 5), dtype=np.float32)
  624. expanded_features[:hyperedge_features.shape[0]] = hyperedge_features
  625. hyperedge_features = expanded_features
  626. current_label_handled = False
  627. if 'MurckoScaffoldCore' in label:
  628. hyperedge_features[i, 0] = 0.95
  629. current_label_handled = True
  630. elif 'MurckoScaffoldFull' in label:
  631. hyperedge_features[i, 0] = 0.90
  632. current_label_handled = True
  633. if not current_label_handled:
  634. if 'Ring' in label:
  635. if 'Aromatic' in label:
  636. hyperedge_features[i, 0] = 0.85
  637. else:
  638. hyperedge_features[i, 0] = 0.75
  639. elif any(fg in label for fg in ['Amine', 'Alcohol', 'Acid', 'Amide', 'Carbonyl', 'Ester', 'Ketone', 'Aldehyde']):
  640. hyperedge_features[i, 0] = 0.80
  641. elif any(fg in label for fg in ['Halogen', 'Cyano', 'Nitro', 'Sulfonamide', 'Thiol', 'Sulfone']):
  642. hyperedge_features[i, 0] = 0.70
  643. elif 'Connector' in label or 'Bridge' in label or 'Substituent' in label:
  644. hyperedge_features[i, 0] = 0.60
  645. elif 'IsolatedAtomGroup' in label or 'SingleAtom' in label:
  646. hyperedge_features[i, 0] = 0.30
  647. else:
  648. hyperedge_features[i, 0] = 0.50
  649. if 'Ring' in label:
  650. if 'Aromatic' in label:
  651. hyperedge_features[i, 0] = 0.9
  652. else:
  653. hyperedge_features[i, 0] = 0.7
  654. elif any(fg in label for fg in ['Amine', 'Alcohol', 'Acid', 'Amide', 'Carbonyl']):
  655. hyperedge_features[i, 0] = 0.8
  656. elif any(fg in label for fg in ['Halogen', 'Cyano', 'Nitro']):
  657. hyperedge_features[i, 0] = 0.75
  658. elif 'Connector' in label or 'Bridge' in label:
  659. hyperedge_features[i, 0] = 0.6
  660. else:
  661. hyperedge_features[i, 0] = 0.5
  662. hyperedge_size = len(edge)
  663. normalized_size = min(1.0, hyperedge_size / 10.0)
  664. hyperedge_features[i, 1] = normalized_size
  665. electron_feature = 0.0
  666. atom_count = 0
  667. for atom_idx in edge:
  668. if atom_idx < self.mol.GetNumAtoms():
  669. atom = self.mol.GetAtomWithIdx(atom_idx)
  670. atomic_num = atom.GetAtomicNum()
  671. if atomic_num in [7, 8, 9, 17, 35, 53]: # N, O, F, Cl, Br, I
  672. electron_feature += 0.8
  673. elif atomic_num == 6 and atom.GetIsAromatic():
  674. electron_feature += 0.6
  675. elif atomic_num == 6:
  676. electron_feature += 0.4
  677. elif atomic_num == 1:
  678. electron_feature += 0.2
  679. else:
  680. electron_feature += 0.5
  681. atom_count += 1
  682. hyperedge_features[i, 2] = electron_feature / max(1, atom_count)
  683. connectivity = 0.0
  684. for other_idx, other_edge in enumerate(self.hyperedges):
  685. if i != other_idx:
  686. if set(edge).intersection(set(other_edge)):
  687. connectivity += 1.0
  688. hyperedge_features[i, 3] = min(1.0, connectivity / max(1, len(self.hyperedges) / 2))
  689. pharmacophore_patterns = {
  690. 'HBondDonor': ['[OH]', '[NH]', '[NH2]'],
  691. 'HBondAcceptor': ['[O]', '[N;!$(N-*=O)]'],
  692. 'Hydrophobic': ['[C;!$(C=O);!$(C#N)]', '[c]'],
  693. 'Aromatic': ['c1ccccc1', 'c1ccncc1'],
  694. 'Charged': ['[+]', '[-]', '[N+]', '[O-]']
  695. }
  696. is_pharmacophore = False
  697. for pattern_list in pharmacophore_patterns.values():
  698. for pattern in pattern_list:
  699. patt = Chem.MolFromSmarts(pattern)
  700. if patt and self.mol.HasSubstructMatch(patt):
  701. matches = self.mol.GetSubstructMatches(patt)
  702. for match in matches:
  703. if any(atom_idx in edge for atom_idx in match):
  704. is_pharmacophore = True
  705. break
  706. if is_pharmacophore:
  707. break
  708. if is_pharmacophore:
  709. break
  710. hyperedge_features[i, 4] = 0.9 if is_pharmacophore else 0.4
  711. if hyperedge_features.shape[0] != num_hyperedges:
  712. print(
  713. f"Warning: Final hyperedge feature matrix shape {hyperedge_features.shape} doesn't match expected size {num_hyperedges}. Fixing.")
  714. corrected_features = np.zeros((num_hyperedges, 5), dtype=np.float32)
  715. min_size = min(hyperedge_features.shape[0], num_hyperedges)
  716. corrected_features[:min_size] = hyperedge_features[:min_size]
  717. hyperedge_features = corrected_features
  718. return hyperedge_features
  719. def get_adjacency_matrix(self):
  720. num_nodes = self.mol.GetNumAtoms()
  721. num_hyperedges = len(self.hyperedges)
  722. adjacency_matrix = np.zeros((num_nodes, num_hyperedges), dtype=int)
  723. for hyperedge_idx, hyperedge in enumerate(self.hyperedges):
  724. for node in hyperedge:
  725. adjacency_matrix[node][hyperedge_idx] = 1
  726. return adjacency_matrix
  727. class MoleculeData(Data):
  728. def __init__(self, x=None, edge_index=None, y=None, hyperedge_attr=None, smiles=None, **kwargs):
  729. super(MoleculeData, self).__init__(**kwargs)
  730. self.hyperedge_attr = None
  731. self.smiles = None
  732. if x is not None:
  733. if not isinstance(x, torch.Tensor):
  734. x = torch.FloatTensor(x)
  735. if x.size(0) == 0:
  736. print("Warning: Empty feature matrix provided")
  737. self.x = x
  738. if edge_index is not None:
  739. if not isinstance(edge_index, torch.Tensor):
  740. edge_index = torch.LongTensor(edge_index)
  741. if edge_index.dim() != 2 or edge_index.size(0) != 2:
  742. print(f"Warning: Unusual edge_index shape: {edge_index.shape}")
  743. self.edge_index = edge_index
  744. if y is not None:
  745. if not isinstance(y, torch.Tensor):
  746. y = torch.FloatTensor(y)
  747. self.y = y
  748. if hyperedge_attr is not None:
  749. if not isinstance(hyperedge_attr, torch.Tensor):
  750. hyperedge_attr = torch.FloatTensor(hyperedge_attr)
  751. self.hyperedge_attr = hyperedge_attr
  752. if smiles is not None:
  753. self.smiles = smiles
  754. @classmethod
  755. def from_cache(cls, filepath):
  756. data = torch.load(filepath)
  757. if not cls.verify_data(data):
  758. raise ValueError("Invalid cached data")
  759. return data
  760. @staticmethod
  761. def verify_data(data):
  762. if not isinstance(data, Data):
  763. print("Invalid data type")
  764. return False
  765. if not hasattr(data, 'x') or not hasattr(data, 'edge_index') or not hasattr(data, 'y'):
  766. print("Missing required attributes in data")
  767. return False
  768. if data.x is None:
  769. print("Data contains None x attribute")
  770. return False
  771. if data.edge_index is None:
  772. print("Data contains None edge_index attribute")
  773. return False
  774. if data.y is None:
  775. print("Data contains None y attribute")
  776. return False
  777. if data.x.size(0) == 0:
  778. print("Empty x in data")
  779. return False
  780. if data.edge_index.numel() == 0:
  781. print("Empty edge_index in data")
  782. return False
  783. if not torch.isfinite(data.x).all():
  784. print("Non-finite values in data.x")
  785. return False
  786. if not torch.isfinite(data.y).all():
  787. print("Non-finite values in data.y")
  788. return False
  789. if hasattr(data, 'hyperedge_attr') and data.hyperedge_attr is not None:
  790. if not torch.isfinite(data.hyperedge_attr).all():
  791. print("Non-finite values in data.hyperedge_attr")
  792. return False
  793. num_nodes = data.x.size(0)
  794. if data.edge_index.numel() > 0:
  795. if data.edge_index.size(0) != 2:
  796. print(f"Invalid edge_index shape: expected (2, N), got {data.edge_index.shape}")
  797. return False
  798. node_indices = data.edge_index[0]
  799. edge_indices = data.edge_index[1]
  800. if node_indices.min().item() < 0 or node_indices.max().item() >= num_nodes:
  801. print(
  802. f"Invalid node indices in edge_index: range [{node_indices.min().item()}, {node_indices.max().item()}], num_nodes={num_nodes}")
  803. return False
  804. if edge_indices.min().item() < 0:
  805. print(f"Invalid edge indices in edge_index: min={edge_indices.min().item()} is negative")
  806. return False
  807. if hasattr(data, 'hyperedge_attr') and data.hyperedge_attr is not None:
  808. num_hyperedges = data.hyperedge_attr.size(0)
  809. if edge_indices.max().item() >= num_hyperedges:
  810. print(
  811. f"Invalid edge indices in edge_index: max={edge_indices.max().item()}, num_hyperedges={num_hyperedges}")
  812. return False
  813. return True
  814. def get_hyperedge_attr(self):
  815. return self.hyperedge_attr
  816. class MoleculeDataset(Dataset):
  817. def __init__(self, data_path, dataset_name, task_type='classification',
  818. max_samples=None, feature_type='combined'):
  819. super(MoleculeDataset, self).__init__()
  820. self.feature_type = feature_type
  821. self.dataset_name = dataset_name.lower()
  822. self.task_type = task_type
  823. self.cache_dir = f"./processed_data/{dataset_name}/"
  824. self.metadata_file = os.path.join(self.cache_dir, "metadata.pt")
  825. self.is_multi_label = False
  826. os.makedirs(self.cache_dir, exist_ok=True)
  827. self.label_encoding = {
  828. 'positive': 1,
  829. 'negative': -1,
  830. 'missing': 0,
  831. 'original_to_encoded': {},
  832. 'encoded_to_original': {}
  833. }
  834. if self._try_load_cache():
  835. print("\nLoaded processed data from cache")
  836. return
  837. print(f"\nLoading dataset from {data_path}")
  838. df = pd.read_csv(data_path)
  839. print(f"Total records in CSV: {len(df)}")
  840. df.columns = df.columns.str.strip().str.lower()
  841. dataset_handler = getattr(self, f"_process_{self.dataset_name}", None)
  842. if dataset_handler:
  843. dataset_handler(df)
  844. else:
  845. raise ValueError(f"Unsupported dataset: {dataset_name}")
  846. if max_samples is not None:
  847. print(f"\nUsing only the first {max_samples} samples for testing.")
  848. self.smiles_list = self.smiles_list[:max_samples]
  849. self.labels = self.labels[:max_samples]
  850. self.valid_smiles = []
  851. self.valid_indices = []
  852. self.processed_data = []
  853. self.hyperedge_attrs = {}
  854. print("\nProcessing molecules...")
  855. success_count = 0
  856. error_count = 0
  857. empty_count = 0
  858. for i, smiles in enumerate(tqdm(self.smiles_list)):
  859. if pd.isna(smiles):
  860. empty_count += 1
  861. continue
  862. try:
  863. mol = self._robust_smiles_to_mol(smiles)
  864. if mol is None:
  865. error_count += 1
  866. continue
  867. hypergraph = MoleculeHypergraph(feature_type=self.feature_type)
  868. try:
  869. hypergraph.build_from_mol(mol)
  870. except Exception as e:
  871. error_count += 1
  872. print(f"\nError building hypergraph for molecule {i} with SMILES {smiles}: {str(e)}")
  873. continue
  874. if hypergraph.node_features is None or hypergraph.hyperedge_index is None:
  875. error_count += 1
  876. print(f"\nInvalid hypergraph for molecule {i} with SMILES {smiles}")
  877. continue
  878. x = torch.FloatTensor(hypergraph.node_features)
  879. edge_index = torch.LongTensor(hypergraph.hyperedge_index)
  880. num_nodes = x.size(0)
  881. if edge_index.shape[0] != 2:
  882. error_count += 1
  883. print(
  884. f"\nInvalid edge_index shape for molecule {i} with SMILES {smiles}: expected (2, N), got {edge_index.shape}")
  885. continue
  886. node_indices = edge_index[0]
  887. edge_indices = edge_index[1]
  888. if node_indices.max() >= num_nodes or node_indices.min() < 0:
  889. error_count += 1
  890. print(
  891. f"\nInvalid node indices in edge_index for molecule {i} with SMILES {smiles}: range [{node_indices.min()}, {node_indices.max()}], num_nodes={num_nodes}")
  892. continue
  893. num_hyperedges = len(hypergraph.hyperedges)
  894. if edge_indices.max() >= num_hyperedges or edge_indices.min() < 0:
  895. error_count += 1
  896. print(
  897. f"\nInvalid edge indices in edge_index for molecule {i} with SMILES {smiles}: range [{edge_indices.min()}, {edge_indices.max()}], num_hyperedges={num_hyperedges}")
  898. continue
  899. y = self.process_label(i)
  900. if not (torch.isfinite(x).all() and torch.isfinite(y).all()):
  901. error_count += 1
  902. print(f"\nNon-finite values in molecule {i}")
  903. continue
  904. hyperedge_attr = hypergraph.generate_enhanced_hyperedge_attributes()
  905. if hyperedge_attr is None:
  906. print(f"Warning: No hyperedges found for molecule {i}. Creating default attribute.")
  907. hyperedge_attr = np.zeros((1, 5), dtype=np.float32)
  908. if isinstance(hyperedge_attr, torch.Tensor):
  909. self.hyperedge_attrs[smiles] = hyperedge_attr.cpu().numpy()
  910. else:
  911. self.hyperedge_attrs[smiles] = hyperedge_attr
  912. hyperedge_attr_tensor = torch.FloatTensor(hyperedge_attr)
  913. data = MoleculeData(
  914. x=x,
  915. edge_index=edge_index,
  916. y=y,
  917. hyperedge_attr=hyperedge_attr_tensor,
  918. smiles=smiles
  919. )
  920. data.smiles = smiles
  921. self.processed_data.append(data)
  922. self.valid_smiles.append(smiles)
  923. self.valid_indices.append(i)
  924. success_count += 1
  925. except Exception as e:
  926. error_count += 1
  927. print(f"\nError processing molecule {i}: {str(e)}")
  928. continue
  929. if len(self.processed_data) == 0:
  930. raise ValueError("No valid molecules processed!")
  931. self._check_data_validity()
  932. self._save_cache()
  933. print("\nProcessing Summary:")
  934. print(f"Total molecules: {len(self.smiles_list)}")
  935. print(f"Successfully processed: {success_count}")
  936. print(f"Failed to process: {error_count}")
  937. print(f"Empty SMILES: {empty_count}")
  938. self._print_label_statistics()
  939. def _robust_smiles_to_mol(self, smiles):
  940. mol = Chem.MolFromSmiles(smiles)
  941. if mol is not None:
  942. return mol
  943. mol = Chem.MolFromSmiles(smiles, sanitize=False)
  944. if mol is None:
  945. return None
  946. try:
  947. mol.UpdatePropertyCache(strict=False)
  948. for atom in mol.GetAtoms():
  949. if atom.GetSymbol() == 'N':
  950. if atom.GetExplicitValence() >= 4 and atom.GetFormalCharge() == 0:
  951. atom.SetFormalCharge(1)
  952. Chem.SanitizeMol(mol)
  953. return mol
  954. except Exception:
  955. try:
  956. Chem.SanitizeMol(mol, sanitizeOps=Chem.SanitizeFlags.SANITIZE_ALL ^ Chem.SanitizeFlags.SANITIZE_PROPERTIES)
  957. return mol
  958. except:
  959. return None
  960. def _process_bace(self, df):
  961. """处理 BACE 数据集(单标签分类)"""
  962. required_cols = ['mol', 'class']
  963. if not all(col in df.columns for col in required_cols):
  964. raise ValueError(f"bace dataset must contain columns: {required_cols}")
  965. self.smiles_list = df['mol'].values
  966. self.labels = df['class'].values
  967. self.labels = np.where(self.labels == 0, -1, self.labels)
  968. self.label_cols = ['class']
  969. self.is_multi_label = False
  970. self.process_label = lambda idx: torch.FloatTensor([[float(self.labels[idx])]])
  971. print("\nLabel distribution in original data:")
  972. print(df['class'].value_counts(normalize=True))
  973. def _process_bbbp(self, df):
  974. """处理 BBBP 数据集(单标签分类)"""
  975. required_cols = ['smiles', 'p_np']
  976. if not all(col in df.columns for col in required_cols):
  977. raise ValueError(f"BBBP dataset must contain columns: {required_cols}")
  978. self.smiles_list = df['smiles'].values
  979. self.labels = df['p_np'].values
  980. self.labels = np.where(self.labels == 0, -1, self.labels)
  981. self.label_cols = ['p_np']
  982. self.is_multi_label = False
  983. self.process_label = lambda idx: torch.FloatTensor([[float(self.labels[idx])]])
  984. print("\nLabel distribution in original data:")
  985. print(df['p_np'].value_counts(normalize=True))
  986. def _process_sider(self, df):
  987. """处理 SIDER 数据集(多标签分类)"""
  988. if 'smiles' not in df.columns:
  989. raise ValueError("SIDER dataset must contain a 'smiles' column")
  990. self.smiles_list = df['smiles'].values
  991. label_cols = [col for col in df.columns if col != 'smiles']
  992. self.label_cols = label_cols
  993. self.labels = df[label_cols].values.astype(float)
  994. self.labels = np.where(self.labels == 0, -1, self.labels)
  995. self.is_multi_label = True
  996. self.process_label = lambda idx: torch.FloatTensor(self.labels[idx]).unsqueeze(0)
  997. print("\nLabel distribution in original data:")
  998. print(df[label_cols].apply(lambda x: pd.value_counts(x, normalize=True)))
  999. def _process_clintox(self, df):
  1000. """处理 ClinTox 数据集(单标签分类)"""
  1001. if 'smiles' not in df.columns:
  1002. raise ValueError("ClinTox dataset must contain a 'smiles' column")
  1003. label_cols = ['fda_approved', 'ct_tox']
  1004. required_cols = ['smiles'] + label_cols
  1005. if not all(col in df.columns for col in required_cols):
  1006. raise ValueError(f"ClinTox dataset must contain columns: {required_cols}")
  1007. self.smiles_list = df['smiles'].values
  1008. self.labels = df[label_cols].values.astype(float)
  1009. self.labels = np.where(self.labels == 0, -1, self.labels)
  1010. self.labels = np.nan_to_num(self.labels, nan=0)
  1011. self.label_cols = label_cols
  1012. self.is_multi_label = True
  1013. self.process_label = lambda idx: torch.FloatTensor(self.labels[idx]).unsqueeze(0)
  1014. def _process_tox21(self, df):
  1015. """处理 Tox21 数据集(多标签分类)"""
  1016. if 'smiles' not in df.columns:
  1017. raise ValueError("Tox21 dataset must contain a 'smiles' column")
  1018. self.smiles_list = df['smiles'].values
  1019. label_cols = [col for col in df.columns if col.lower() not in ['smiles', 'mol_id']]
  1020. self.label_cols = label_cols
  1021. self.labels = df[label_cols].values.astype(float)
  1022. self.labels = np.where(self.labels == 0, -1, self.labels)
  1023. self.labels = np.nan_to_num(self.labels, nan=0)
  1024. self.is_multi_label = True
  1025. self.process_label = lambda idx: torch.FloatTensor(self.labels[idx]).unsqueeze(0)
  1026. print("\nLabel distribution after processing:")
  1027. print(df[label_cols].apply(lambda x: pd.value_counts(x, normalize=True)))
  1028. def _process_toxcast(self, df):
  1029. """处理 ToxCast 数据集(多标签分类)"""
  1030. if 'smiles' not in df.columns:
  1031. raise ValueError("ToxCast dataset must contain a 'smiles' column")
  1032. self.smiles_list = df['smiles'].values
  1033. label_cols = [col for col in df.columns if col != 'smiles']
  1034. self.label_cols = label_cols
  1035. self.labels = df[label_cols].values.astype(float)
  1036. self.labels = np.where(self.labels == 0, -1, self.labels)
  1037. self.labels = np.nan_to_num(self.labels, nan=0)
  1038. self.is_multi_label = True
  1039. self.process_label = lambda idx: torch.FloatTensor(self.labels[idx]).unsqueeze(0)
  1040. print("\nLabel distribution in original data:")
  1041. print(df[label_cols].apply(lambda x: pd.value_counts(x, normalize=True)))
  1042. def _process_esol(self, df):
  1043. """处理 ESOL 数据集(单目标回归)"""
  1044. required_cols = ['smiles', 'measured log solubility in mols per litre']
  1045. if not all(col in df.columns for col in required_cols):
  1046. raise ValueError(f"ESOL dataset must contain columns: {required_cols}")
  1047. self.smiles_list = df['smiles'].values
  1048. self.labels = df['measured log solubility in mols per litre'].values.astype(float)
  1049. self.label_cols = ['log_solubility']
  1050. self.is_multi_label = False
  1051. self.task_type = 'regression'
  1052. self.process_label = lambda idx: torch.FloatTensor([[float(self.labels[idx])]])
  1053. print("\nESOL dataset statistics:")
  1054. print(f"Min solubility: {np.min(self.labels):.2f}")
  1055. print(f"Max solubility: {np.max(self.labels):.2f}")
  1056. print(f"Mean solubility: {np.mean(self.labels):.2f}")
  1057. print(f"Std solubility: {np.std(self.labels):.2f}")
  1058. def _process_freesolv(self, df):
  1059. """处理 FreeSolv 数据集(单目标回归)"""
  1060. required_cols = ['smiles', 'expt']
  1061. if not all(col in df.columns for col in required_cols):
  1062. raise ValueError(f"FreeSolv dataset must contain columns: {required_cols}")
  1063. self.smiles_list = df['smiles'].values
  1064. self.labels = df['expt'].values.astype(float)
  1065. self.label_cols = ['expt']
  1066. self.is_multi_label = False
  1067. self.task_type = 'regression'
  1068. self.process_label = lambda idx: torch.FloatTensor([[float(self.labels[idx])]])
  1069. print("\nFreeSolv dataset statistics:")
  1070. print(f"Min expt: {np.min(self.labels):.2f}")
  1071. print(f"Max expt: {np.max(self.labels):.2f}")
  1072. print(f"Mean expt: {np.mean(self.labels):.2f}")
  1073. print(f"Std expt: {np.std(self.labels):.2f}")
  1074. def _process_lipophilicity(self, df):
  1075. """处理 Lipophilicity 数据集(单目标回归)"""
  1076. required_cols = ['cmpd_chemblid', 'exp', 'smiles']
  1077. if not all(col in df.columns for col in required_cols):
  1078. raise ValueError(f"Lipophilicity dataset must contain columns: {required_cols}")
  1079. self.smiles_list = df['smiles'].values
  1080. self.labels = df['exp'].values.astype(float)
  1081. self.label_cols = ['exp']
  1082. self.is_multi_label = False
  1083. self.task_type = 'regression'
  1084. self.process_label = lambda idx: torch.FloatTensor([[float(self.labels[idx])]])
  1085. print("\nLipophilicity dataset statistics:")
  1086. print(f"Min exp: {np.min(self.labels):.2f}")
  1087. print(f"Max exp: {np.max(self.labels):.2f}")
  1088. print(f"Mean exp: {np.mean(self.labels):.2f}")
  1089. print(f"Std exp: {np.std(self.labels):.2f}")
  1090. def _process_qm8(self, df):
  1091. """处理 QM8 数据集(多目标回归)"""
  1092. required_cols = ['smiles']
  1093. if not all(col in df.columns for col in required_cols):
  1094. raise ValueError("QM8 dataset must contain a 'smiles' column")
  1095. self.smiles_list = df['smiles'].values
  1096. label_cols = [col for col in df.columns if col != 'smiles']
  1097. self.label_cols = label_cols
  1098. self.labels = df[label_cols].values.astype(float)
  1099. self.is_multi_label = True
  1100. self.task_type = 'regression'
  1101. self.process_label = lambda idx: torch.FloatTensor(self.labels[idx]).unsqueeze(0)
  1102. print("\nQM8 dataset statistics:")
  1103. print(f"Label columns: {label_cols}")
  1104. print(f"Min values: {np.min(self.labels, axis=0)}")
  1105. print(f"Max values: {np.max(self.labels, axis=0)}")
  1106. print(f"Mean values: {np.mean(self.labels, axis=0)}")
  1107. print(f"Std values: {np.std(self.labels, axis=0)}")
  1108. def _process_qm7(self, df):
  1109. """
  1110. 处理 QM7 数据集
  1111. """
  1112. if 'smiles' not in df.columns:
  1113. raise ValueError("QM7 dataset must contain a 'smiles' column")
  1114. possible_targets = ['u0_atom', 'u0', 'target', 'eat']
  1115. target_col = next((c for c in possible_targets if c in df.columns), None)
  1116. if target_col is None:
  1117. target_col = [c for c in df.columns if c != 'smiles'][-1]
  1118. print(f"Warning: Specific QM7 target column not found. Using '{target_col}' as target.")
  1119. self.smiles_list = df['smiles'].values
  1120. self.labels = df[target_col].values.astype(float)
  1121. self.label_cols = [target_col]
  1122. self.is_multi_label = False
  1123. self.task_type = 'regression'
  1124. self.process_label = lambda idx: torch.FloatTensor([[float(self.labels[idx])]])
  1125. print(f"\nQM7 dataset statistics ({target_col}):")
  1126. print(f"Total samples: {len(self.labels)}")
  1127. print(f"Min: {np.min(self.labels):.4f}")
  1128. print(f"Max: {np.max(self.labels):.4f}")
  1129. print(f"Mean: {np.mean(self.labels):.4f}")
  1130. print(f"Std: {np.std(self.labels):.4f}")
  1131. def _check_data_validity(self):
  1132. if len(self.processed_data) == 0:
  1133. print("Warning: No processed data to check validity.")
  1134. return
  1135. if self.task_type == 'classification':
  1136. invalid_tasks = []
  1137. if self.is_multi_label:
  1138. if not hasattr(self, 'label_cols') or not self.label_cols:
  1139. print("Warning: Cannot check validity, label_cols not defined.")
  1140. return
  1141. num_tasks = len(self.label_cols)
  1142. if num_tasks == 0:
  1143. print("Warning: Cannot check validity, num_tasks is zero.")
  1144. return
  1145. for task_idx in range(num_tasks):
  1146. task_labels = []
  1147. for data in self.processed_data:
  1148. if hasattr(data, 'y') and data.y is not None and data.y.dim() == 2 and data.y.shape[0] == 1 and data.y.shape[1] > task_idx:
  1149. try:
  1150. label_value = data.y[0, task_idx].item()
  1151. task_labels.append(label_value)
  1152. except IndexError:
  1153. print(f"Warning: IndexError accessing data.y[0, {task_idx}] for data with y shape {data.y.shape}. Skipping.")
  1154. except Exception as e:
  1155. print(f"Warning: Error accessing label for task {task_idx}: {e}. Skipping.")
  1156. if not task_labels:
  1157. print(f"Warning: No valid labels found for task {task_idx} ({self.label_cols[task_idx]}).")
  1158. invalid_tasks.append((self.label_cols[task_idx], []))
  1159. continue
  1160. valid_labels = [v for v in task_labels if not np.isnan(v)]
  1161. if not valid_labels:
  1162. print(f"Warning: All labels are NaN for task {task_idx} ({self.label_cols[task_idx]}) after filtering.")
  1163. invalid_tasks.append((self.label_cols[task_idx], []))
  1164. continue
  1165. unique = np.unique(valid_labels)
  1166. if len(unique) < 2:
  1167. invalid_tasks.append((self.label_cols[task_idx], unique))
  1168. else:
  1169. try:
  1170. label_values = []
  1171. for data in self.processed_data:
  1172. if hasattr(data, 'y') and data.y is not None and data.y.dim() == 2 and data.y.shape == (1, 1):
  1173. try:
  1174. label_value = data.y[0, 0].item()
  1175. label_values.append(label_value)
  1176. except IndexError:
  1177. print(f"Warning: IndexError accessing data.y[0, 0] for data with y shape {data.y.shape}. Skipping.")
  1178. except Exception as e:
  1179. print(f"Warning: Error accessing single label: {e}. Skipping.")
  1180. if not label_values:
  1181. print("Warning: No valid labels found for single-label task.")
  1182. invalid_tasks.append(('main_label', []))
  1183. else:
  1184. unique = np.unique(label_values)
  1185. if len(unique) < 2:
  1186. invalid_tasks.append(('main_label', unique))
  1187. except Exception as e:
  1188. print(f"Error analyzing single-label distribution: {str(e)}")
  1189. if invalid_tasks:
  1190. print("\nWARNING: Found tasks with only one class (or no valid labels):")
  1191. for task, classes in invalid_tasks:
  1192. if classes:
  1193. print(f" Task {task}: classes {classes}")
  1194. else:
  1195. print(f" Task {task}: No valid labels found.")
  1196. print("This may cause evaluation metrics like AUC/AUPR to fail or be unreliable!")
  1197. def get_original_label_value(self, encoded_value):
  1198. if hasattr(self, 'label_encoding') and 'encoded_to_original' in self.label_encoding:
  1199. return self.label_encoding['encoded_to_original'].get(encoded_value, encoded_value)
  1200. return encoded_value
  1201. def is_missing_label(self, encoded_value):
  1202. if hasattr(self, 'label_encoding'):
  1203. return encoded_value == self.label_encoding['missing']
  1204. return False
  1205. def _print_label_statistics(self):
  1206. try:
  1207. if self.is_multi_label and self.dataset_name in ['sider', 'tox21', 'toxcast']:
  1208. print("\nLabel distribution in processed data:")
  1209. for idx, col in enumerate(self.label_cols):
  1210. if idx < self.labels.shape[1]:
  1211. label_values = self.labels[:, idx]
  1212. unique, counts = np.unique(label_values, return_counts=True)
  1213. print(f"{col}:")
  1214. for u, c in zip(unique, counts):
  1215. print(f" Label {u}: {c} samples ({c / len(label_values) * 100:.2f}%)")
  1216. elif self.dataset_name in ['qm8', 'qm9']:
  1217. try:
  1218. processed_labels_array = []
  1219. for data in self.processed_data:
  1220. if data.y.numel() > 0:
  1221. processed_labels_array.append(data.y.cpu().numpy())
  1222. processed_labels_array = np.array(processed_labels_array)
  1223. print("\nProcessed label statistics (per dimension):")
  1224. for i, col in enumerate(self.label_cols):
  1225. if i < processed_labels_array.shape[1]:
  1226. col_values = processed_labels_array[:, i]
  1227. print(
  1228. f"{col}: Min={col_values.min():.4f}, Max={col_values.max():.4f}, Mean={col_values.mean():.4f}, Std={col_values.std():.4f}")
  1229. except Exception as e:
  1230. print(f"Error computing label statistics: {str(e)}")
  1231. else:
  1232. try:
  1233. processed_labels = []
  1234. for data in self.processed_data:
  1235. if data.y.numel() >= 1:
  1236. processed_labels.append(data.y[0].item())
  1237. unique_labels = np.unique(processed_labels)
  1238. print("\nLabel distribution in processed data:")
  1239. for label in unique_labels:
  1240. count = sum(1 for y in processed_labels if y == label)
  1241. print(f"Label {label}: {count} samples ({count / len(processed_labels) * 100:.2f}%)")
  1242. except Exception as e:
  1243. print(f"Error computing label distribution: {str(e)}")
  1244. except Exception as e:
  1245. print(f"Error in label statistics calculation: {str(e)}")
  1246. def _save_hyperedge_attrs(self):
  1247. for data in self.processed_data:
  1248. if hasattr(data, 'smiles') and data.smiles and hasattr(data,
  1249. 'hyperedge_attr') and data.hyperedge_attr is not None:
  1250. if isinstance(data.hyperedge_attr, torch.Tensor):
  1251. self.hyperedge_attrs[data.smiles] = data.hyperedge_attr.cpu().numpy()
  1252. else:
  1253. self.hyperedge_attrs[data.smiles] = data.hyperedge_attr
  1254. hyperedge_attrs_file = os.path.join(self.cache_dir, "hyperedge_attrs.pt")
  1255. torch.save(self.hyperedge_attrs, hyperedge_attrs_file)
  1256. print(f"Saved hyperedge attributes for {len(self.hyperedge_attrs)} molecules")
  1257. def _get_cache_filename(self, idx):
  1258. return os.path.join(self.cache_dir, f"mol_{idx}.pt")
  1259. def _save_metadata(self):
  1260. has_hyperedge_attr = False
  1261. if self.processed_data and len(self.processed_data) > 0:
  1262. sample_item = self.processed_data[0]
  1263. has_hyperedge_attr = hasattr(sample_item, 'hyperedge_attr') and sample_item.hyperedge_attr is not None
  1264. metadata = {
  1265. 'dataset_name': self.dataset_name,
  1266. 'task_type': self.task_type,
  1267. 'num_classes': self._get_num_classes(),
  1268. 'label_cols': self.label_cols,
  1269. 'is_multi_label': self.is_multi_label,
  1270. 'valid_indices': self.valid_indices,
  1271. 'has_hyperedge_attr': has_hyperedge_attr,
  1272. 'label_encoding': getattr(self, 'label_encoding', {})
  1273. }
  1274. torch.save(metadata, self.metadata_file)
  1275. def _load_metadata(self):
  1276. if os.path.exists(self.metadata_file):
  1277. try:
  1278. metadata = torch.load(self.metadata_file)
  1279. if metadata.get('dataset_name') == self.dataset_name:
  1280. self.label_cols = metadata.get('label_cols', [])
  1281. self.is_multi_label = metadata.get('is_multi_label', False)
  1282. self.valid_indices = metadata.get('valid_indices', [])
  1283. self.task_type = metadata.get('task_type', self.task_type)
  1284. self.has_hyperedge_attr = metadata.get('has_hyperedge_attr', False)
  1285. self.label_encoding = metadata.get('label_encoding', {
  1286. 'positive': 1,
  1287. 'negative': -1,
  1288. 'missing': 0,
  1289. 'original_to_encoded': {},
  1290. 'encoded_to_original': {}
  1291. })
  1292. return True
  1293. return False
  1294. except Exception as e:
  1295. print(f"Error loading metadata: {str(e)}")
  1296. return False
  1297. return False
  1298. def _try_load_cache(self):
  1299. if not self._load_metadata():
  1300. return False
  1301. cache_files = [f for f in os.listdir(self.cache_dir) if
  1302. f.endswith(".pt") and not f == "metadata.pt" and not f == "hyperedge_attrs.pt"]
  1303. if not cache_files:
  1304. return False
  1305. print("\nChecking cached data...")
  1306. self.hyperedge_attrs = {}
  1307. hyperedge_attrs_file = os.path.join(self.cache_dir, "hyperedge_attrs.pt")
  1308. if os.path.exists(hyperedge_attrs_file):
  1309. try:
  1310. self.hyperedge_attrs = torch.load(hyperedge_attrs_file)
  1311. print(f"Loaded hyperedge attributes dictionary for {len(self.hyperedge_attrs)} molecules")
  1312. except Exception as e:
  1313. print(f"Error loading hyperedge attributes: {e}")
  1314. self.hyperedge_attrs = {}
  1315. self.processed_data = []
  1316. self.valid_smiles = []
  1317. success = 0
  1318. failures = 0
  1319. max_idx = -1
  1320. for f in cache_files:
  1321. try:
  1322. idx = int(f.split('_')[1].split('.')[0])
  1323. max_idx = max(max_idx, idx)
  1324. except:
  1325. continue
  1326. for idx in tqdm(range(max_idx + 1)):
  1327. cache_file = self._get_cache_filename(idx)
  1328. if os.path.exists(cache_file):
  1329. try:
  1330. data = torch.load(cache_file)
  1331. if not hasattr(data, 'x') or data.x is None or data.x.shape[0] == 0:
  1332. print(f"Skipping invalid cache file (bad x): {cache_file}")
  1333. failures += 1
  1334. continue
  1335. if not hasattr(data, 'edge_index') or data.edge_index is None:
  1336. print(f"Skipping invalid cache file (bad edge_index): {cache_file}")
  1337. failures += 1
  1338. continue
  1339. if not hasattr(data, 'y') or data.y is None:
  1340. print(f"Skipping invalid cache file (bad y): {cache_file}")
  1341. failures += 1
  1342. continue
  1343. smiles = getattr(data, 'smiles', None)
  1344. hyperedge_attr = getattr(data, 'hyperedge_attr', None)
  1345. if hyperedge_attr is None and smiles is not None and smiles in self.hyperedge_attrs:
  1346. hyperedge_attr_np = self.hyperedge_attrs[smiles]
  1347. hyperedge_attr = torch.FloatTensor(hyperedge_attr_np)
  1348. print(f"Recovered hyperedge attributes for molecule with SMILES {smiles}")
  1349. cleaned_data = MoleculeData(
  1350. x=data.x.clone(),
  1351. edge_index=data.edge_index.clone(),
  1352. y=data.y.clone(),
  1353. hyperedge_attr=hyperedge_attr.clone() if hyperedge_attr is not None else None,
  1354. smiles=smiles
  1355. )
  1356. self.processed_data.append(cleaned_data)
  1357. if smiles is not None:
  1358. self.valid_smiles.append(smiles)
  1359. success += 1
  1360. except Exception as e:
  1361. print(f"Error loading cache file {cache_file}: {str(e)}")
  1362. failures += 1
  1363. print(f"Loaded {success}/{len(cache_files)} valid cached entries ({failures} failures)")
  1364. return success > 0
  1365. def _save_cache(self):
  1366. print("\nSaving processed data to cache...")
  1367. os.makedirs(self.cache_dir, exist_ok=True)
  1368. for idx, data in enumerate(tqdm(self.processed_data)):
  1369. try:
  1370. smiles = getattr(data, 'smiles', None)
  1371. hyperedge_attr = getattr(data, 'hyperedge_attr', None)
  1372. data_cpu = MoleculeData(
  1373. x=data.x.cpu(),
  1374. edge_index=data.edge_index.cpu(),
  1375. y=data.y.cpu(),
  1376. hyperedge_attr=hyperedge_attr.cpu() if hyperedge_attr is not None else None,
  1377. smiles=smiles
  1378. )
  1379. torch.save(data_cpu, self._get_cache_filename(idx))
  1380. except Exception as e:
  1381. print(f"Error saving {idx}: {str(e)}")
  1382. self._save_hyperedge_attrs()
  1383. self._save_metadata()
  1384. print(f"Saved metadata to {self.metadata_file}")
  1385. def len(self):
  1386. return len(self.processed_data)
  1387. def get(self, idx):
  1388. return self.processed_data[idx]
  1389. @property
  1390. def num_node_features(self):
  1391. if len(self.processed_data) == 0:
  1392. raise ValueError("Dataset is empty!")
  1393. return self.processed_data[0].x.size(1)
  1394. def _get_num_classes(self):
  1395. if self.task_type == 'regression':
  1396. if self.is_multi_label:
  1397. return len(self.label_cols)
  1398. else:
  1399. return 1
  1400. elif self.is_multi_label:
  1401. return len(self.label_cols)
  1402. else:
  1403. return 1
  1404. @property
  1405. def num_classes(self):
  1406. if len(self.processed_data) == 0:
  1407. raise ValueError("Dataset is empty!")
  1408. return self._get_num_classes()

dataset.py at commit 171920e, no license · at the source

Overview

Authors: Jiani Ma1, Qi Yang1, Lin Zhang1, Hui Liu1, Yuanting Zheng2,3
  1. School of Information and Control Engineering, China University of Mining and Technology, Xuzhou 221116, P. R. China
  2. National Key Laboratory of Agricultural Microbiology, College of Veterinary Medicine, Huazhong Agricultural University, 430070 Wuhan, Hubei, P. R. China
  3. Faculty of Science, Melbourne Veterinary School, The University of Melbourne, Melbourne, Victoria, Australia
Journal: Computational and structural biotechnology journal, volume 35, issue 1, article 0036
Dates: received 19 December 2025; accepted 15 March 2026; published online 9 April 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.34133/csbj.0036 · PMID 41971950 · PMCID PMC13062488 · OpenAlex W7137359194
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: cellular / molecular (subfield)
Methods: Preprocessing, Machine learning
Journal subjects: General
Topic: Computational Drug Discovery Methods (Computational Theory and Mathematics, Computer Science), according to OpenAlex
Citations: not cited yet (Europe PMC); 42 references in the paper

Abstract

Accurate molecular property prediction is a central task in drug discovery, yet existing methods often struggle to simultaneously capture higher-order molecular structures and chemically grounded semantic information. Graph neural networks are limited to pairwise atomic interactions, while Simplified Molecular Input Line Entry System-based language models lack explicit structural grounding, leading to incomplete structure–property representations. In this work, we propose HYG-mol, an interpretable molecular property prediction framework that integrates hypergraph-based structural modeling with multimodal chemical semantics. Molecules are represented as hypergraphs in which chemically meaningful substructures, such as functional groups and ring systems, are explicitly encoded as hyperedges, enabling direct modeling of higher-order structural dependencies. Chemical semantic information derived from a pretrained language model is fused with physicochemical descriptors at the atomic level. A hypergraph attention network is employed to capture cross-scale interactions and to identify substructures relevant to the prediction task. Extensive evaluations on MoleculeNet benchmark datasets demonstrate that HYG-mol consistently outperforms state-of-the-art baseline methods across both classification and regression tasks. Ablation and interpretability analyses further validate the effectiveness of the proposed representation and reveal strong correspondence between model-identified substructures and chemically meaningful motifs. Overall, HYG-mol provides a unified and interpretable framework for molecular property prediction by explicitly grounding chemical semantics in higher-order structural representations.

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 11 matches between paragraphs and lines of code.

sutera777/HYG-mol

License: none: the authors keep all their rights
State: the link answers, verified on 29 September 2026
Evidence: files inventoried
Commit: 171920ec65c262afd5beaf00af3eb67f706e60b8, 15 March 2026
Languages: Python (15)
Size: 27 files, 15 scripts
Software Heritage: not archived
Found in: “Data Availability”
Holds: README, environment (requirements.txt)
Not found: license file, CITATION.cff, tests, continuous integration, documentation
Tools: PyTorch (7 files), NumPy (6 files), RDKit (4 files), PyTorch Geometric (3 files), scikit-learn (2 files), Matplotlib (1 file), pandas (1 file), Hugging Face Transformers (1 file)
Availability: 1 check, the latest on 29 September 2026: the link answers
  • 29 September 2026: the link answers
16 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;
  • 15 scripts, each with its path and the digest of its content;
  • 11 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 data used to support the findings of this study are available from the corresponding authors upon reasonable request. The molecular property data were derived from the MoleculeNet benchmark dataset (https://moleculenet.org/), which is publicly accessible. The code for the HYG-mol framework and related experimental scripts are deposited in GitHub (https://github.com/sutera777/HYG-mol).

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, 29 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 5 authors, 1 funder, 26 references.

Cite

This paper

Ma, J., Yang, Q., Zhang, L., Liu, H., & Zheng, Y. (2026). HYG-mol: An Interpretable Multimodal Hypergraph Framework for Molecular Property Prediction. Computational and structural biotechnology journal, 35(1), 0036. https://doi.org/10.34133/csbj.0036

BibTeX

@article{ma2026hyg,
author = {Ma, Jiani and Yang, Qi and Zhang, Lin and Liu, Hui and Zheng, Yuanting},
title = {{HYG-mol: An Interpretable Multimodal Hypergraph Framework for Molecular Property Prediction}},
journal = {Computational and structural biotechnology journal},
year = {2026},
month = apr,
volume = {35},
number = {1},
pages = {0036},
publisher = {AAAS Science Partner Journal Program},
issn = {2001-0370},
doi = {10.34133/csbj.0036},
url = {https://doi.org/10.34133/csbj.0036},
pmid = {41971950},
pmcid = {PMC13062488}
}

RIS

TY - JOUR
AU - Ma, Jiani
AU - Yang, Qi
AU - Zhang, Lin
AU - Liu, Hui
AU - Zheng, Yuanting
TI - HYG-mol: An Interpretable Multimodal Hypergraph Framework for Molecular Property Prediction
T2 - Computational and structural biotechnology journal
J2 - Comput Struct Biotechnol J
PY - 2026
DA - 2026/04/09
VL - 35
IS - 1
SP - 0036
SN - 2001-0370
PB - AAAS Science Partner Journal Program
DO - 10.34133/csbj.0036
UR - https://doi.org/10.34133/csbj.0036
LA - en
ER -

CSL-JSON

{
"id": "10.34133/csbj.0036",
"type": "article-journal",
"title": "HYG-mol: An Interpretable Multimodal Hypergraph Framework for Molecular Property Prediction",
"container-title": "Computational and structural biotechnology journal",
"author": [
{
"family": "Ma",
"given": "Jiani"
},
{
"family": "Yang",
"given": "Qi"
},
{
"family": "Zhang",
"given": "Lin"
},
{
"family": "Liu",
"given": "Hui"
},
{
"family": "Zheng",
"given": "Yuanting"
}
],
"container-title-short": "Comput Struct Biotechnol J",
"volume": "35",
"issue": "1",
"page": "0036",
"DOI": "10.34133/csbj.0036",
"PMID": "41971950",
"PMCID": "PMC13062488",
"ISSN": "2001-0370",
"publisher": "AAAS Science Partner Journal Program",
"URL": "https://doi.org/10.34133/csbj.0036",
"language": "en",
"issued": {
"date-parts": [
[
2026,
4,
9
]
]
}
}

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.34133/csbj.0184 [code]
Cross-Species Multitask Learning with Molecular and ADME Descriptors for Liver Microsomal Metabolic Stability.
Journal: Computational and structural biotechnology journal
In common: RDKit, PyTorch Geometric, PyTorch, 4 other tools, cellular / molecular, 3 references
[2] doi:10.1021/acsomega.5c09368 [code]
Structure-Based and AI-Assisted Identification of AGPS Inhibitors for Glioma via Integrated Docking, Molecular Dynamics, and Binding Affinity Screening.
Journal: ACS omega
In common: RDKit, PyTorch Geometric, Hugging Face Transformers, 5 other tools
[3] doi:10.1093/bib/bbag118 [code]
Drug screening for α-synuclein aggregation inhibitors via multimodal graph neural network.
Journal: Briefings in bioinformatics
In common: RDKit, PyTorch Geometric, PyTorch, 4 other tools, cellular / molecular, 1 reference
[4] doi:10.1002/pro.70695 [code]
MIF-MAPMS: Enhancing identification of myelin autoantigenic peptides in multiple sclerosis through multimodal information fusion.
Journal: Protein science : a publication of the Protein Society
In common: RDKit, PyTorch, scikit-learn, 3 other tools, cellular / molecular, 2 references
[5] doi:10.3390/molecules31111972 [code]
A Bond-Level Sequence Framework for Molecular Representation Learning with Structural Constraints.
Journal: Molecules (Basel, Switzerland)
In common: RDKit, PyTorch, scikit-learn, 2 other tools, cellular / molecular, 2 references
[6] doi:10.3390/ijms27156614 [code]
Candidalysin Inhibits &lt;i&gt;Porphyromonas gingivalis&lt;/i&gt; Lipoprotein-Induced IL-1β Production in BV-2 Microglia via Hydrophobic Microbial Interactions.
Journal: International journal of molecular sciences
In common: RDKit, PyTorch Geometric, PyTorch, 4 other tools, cellular / molecular
[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, 4 other tools, cellular / molecular
[8] doi:10.1038/s41598-026-53415-5 [code]
Computational design and immunoinformatics validation of a T cell multi-epitope vaccine targeting glioblastoma stem cells.
Journal: Scientific reports
In common: RDKit, PyTorch Geometric, PyTorch, 4 other tools, cellular / molecular
[9] doi:10.1038/s41586-026-10391-0 [code]
Cell-type-targeted mitochondrial transplantation rescues cell degeneration.
Journal: Nature
In common: RDKit, PyTorch Geometric, PyTorch, 4 other tools, cellular / molecular
[10] doi:10.1021/acs.biochem.5c00596 [code]
Cargo Recognition of Nesprin-2 by the Dynein Adapter Bicaudal D2 for a Nuclear Positioning Pathway That Is Important for Brain Development.
Journal: Biochemistry
In common: RDKit, PyTorch Geometric, PyTorch, 4 other tools, cellular / molecular

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.