CoolFace
Modelpublic

rk-random/PACT-Net

sourceHugging Facemitupdated 9mo agoView on Hugging Face
1likes
polyatomic_featurize.py161 linesDownload Raw Back to data
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