rk-random/PACT-Net
1
1import torch2import numpy as np3from rdkit import Chem4import networkx as nx5from collections import defaultdict6from torch_geometric.data import Data7from polyatomic_complexes.src.complexes.abstract_complex import AbstractComplex8from polyatomic_complexes.src.complexes import PolyatomicGeometrySMILE9 10 11def compressed_topsignal_graph_from_smiles(12 smile: str, y_val: int, topk_lap: int = 513) -> Data | None:14 try:15 # 1) Abstract complex16 pg = PolyatomicGeometrySMILE(smile=smile, mode="abstract")17 ac = pg.smiles_to_geom_complex()18 assert isinstance(ac, AbstractComplex)19 20 # 2) RDKit molecule21 mol = Chem.MolFromSmiles(smile) # type: ignore22 if mol is None:23 return None24 25 # 3) Node features: chain0 value + RDKit descriptors26 chains = ac.get_raw_k_chains()27 chain0 = chains.get("chain_0", [])28 atom_types = [6, 7, 8, 15, 16, 17]29 hyb_types = [30 Chem.rdchem.HybridizationType.SP,31 Chem.rdchem.HybridizationType.SP2,32 Chem.rdchem.HybridizationType.SP3,33 ]34 node_feats = []35 for atom in mol.GetAtoms():36 idx = atom.GetIdx()37 # fallback if chain0 shorter than atom count38 c0 = float(chain0[idx]) if idx < len(chain0) else 0.039 feats = [c0]40 feats += one_hot(atom.GetAtomicNum(), atom_types)41 feats += one_hot(atom.GetHybridization(), hyb_types)42 feats += [43 float(atom.GetDegree()),44 float(atom.GetIsAromatic()),45 float(atom.GetFormalCharge()),46 ]47 node_feats.append(feats)48 x = torch.tensor(node_feats, dtype=torch.float32)49 n = x.size(0) # use number of atoms for all subsequent node counts50 51 # 4) Edges: abstract bonds + RDKit fallback52 sk = ac.get_skeleta().get("molecule_skeleta", [[]])[0]53 zero = next((lst for dim, lst in sk if dim == "0"), [])54 node_ids = [next(iter(fz))[0] for fz in zero]55 atom_map = defaultdict(list)56 for i, nid in enumerate(node_ids):57 symbol = nid.split("_")[0]58 atom_map[symbol].append(i)59 60 edge_index_list, edge_attr_list = [], []61 bond_types = [62 Chem.rdchem.BondType.SINGLE,63 Chem.rdchem.BondType.DOUBLE,64 Chem.rdchem.BondType.TRIPLE,65 Chem.rdchem.BondType.AROMATIC,66 ]67 for a1, a2, (btype, order) in ac.get_bonds():68 bt_val = getattr(Chem.rdchem.BondType, btype, None)69 for i in atom_map.get(a1, []):70 for j in atom_map.get(a2, []):71 if i < n and j < n:72 edge_index_list += [[i, j], [j, i]]73 attr = one_hot(bt_val, bond_types) + [float(order), 0.0]74 edge_attr_list += [attr, attr]75 if not edge_index_list:76 for bond in mol.GetBonds():77 i, j = bond.GetBeginAtomIdx(), bond.GetEndAtomIdx()78 edge_index_list += [[i, j], [j, i]]79 attr = one_hot(bond.GetBondType(), bond_types)80 attr += [float(bond.GetIsConjugated()), float(bond.IsInRing())]81 edge_attr_list += [attr, attr]82 83 edge_index = torch.tensor(edge_index_list, dtype=torch.long).t().contiguous()84 edge_attr = torch.tensor(edge_attr_list, dtype=torch.float32)85 86 # 5) Topology features: centrality + SPD87 G = nx.Graph()88 G.add_nodes_from(range(n))89 G.add_edges_from(edge_index_list)90 cent = nx.closeness_centrality(G)91 spd = dict(nx.all_pairs_shortest_path_length(G))92 cent_vec = [cent.get(i, 0.0) for i in range(n)]93 spd_vec = [94 sum(d.values()) / max(len(d), 1) for d in (spd.get(i, {}) for i in range(n))95 ]96 cent_t = torch.tensor(cent_vec, dtype=torch.float32).view(n, 1)97 spd_t = torch.tensor(spd_vec, dtype=torch.float32).view(n, 1)98 x = torch.cat([x, cent_t, spd_t], dim=1)99 100 # print("MANAGED TO CONCAT?")101 102 # 6) Graph-level features: chain stats + laplacians103 g_stats, lap_feats = [], []104 for k, arr in chains.items():105 if k == "chain_0":106 continue107 a = np.array(arr, dtype=np.float32)108 g_stats += [a.mean(), a.std()]109 110 # print("COMPUTED GRAP STATS")111 112 for grp in ac.get_laplacians().get("molecule_laplacians", []):113 recs = grp if isinstance(grp, list) else [grp]114 for _, mat in recs:115 # use dense eigen solver to avoid ARPACK issues116 M = np.array(mat, dtype=np.float32)117 # compute eigenvalues of symmetric Laplacian118 try:119 eigs = np.linalg.eigvalsh(M)120 except Exception:121 eigs = np.zeros(M.shape[0], dtype=np.float32)122 # take smallest non-zero eigenvalues (skip the first zero)123 nonzero = eigs[eigs > 1e-6]124 vals = nonzero[:topk_lap] if len(nonzero) >= topk_lap else nonzero125 # pad to exactly topk_lap126 if len(vals) < topk_lap:127 vals = np.pad(vals, (0, topk_lap - len(vals)))128 lap_feats += list(vals)129 130 # --- 7) Spectral k-chains stats ---131 spectral = ac.get_spectral_k_chains()132 spec_feats = []133 for arr in spectral.values():134 a = np.array(arr, dtype=np.float32)135 spec_feats += [a.mean(), a.std()]136 137 # --- 8) Betti numbers (components & cycles) ---138 b0 = nx.number_connected_components(G)139 b1 = sum(140 len(nx.cycle_basis(G.subgraph(comp))) for comp in nx.connected_components(G)141 )142 143 # --- 9) Assemble graph_feats ---144 all_feats = g_stats + lap_feats + spec_feats + [float(b0), float(b1)]145 graph_feats = torch.tensor(all_feats, dtype=torch.float32)146 147 # print("managed to feat?")148 149 data = Data(x=x, edge_index=edge_index, edge_attr=edge_attr)150 data.graph_feats = graph_feats151 data.y = torch.tensor([y_val], dtype=torch.float)152 # print(f"SUCCESS for : {smile}")153 return data154 except Exception as e:155 # print(f"Failed {smile}: {e}")156 return None157 158 159def one_hot(val, choices):160 return [1.0 if val == c else 0.0 for c in choices]161 