import csv, io, requests, random, time, math import numpy as np import torch import torch.nn as nn import torch.nn.functional as F from torch.utils.data import Dataset, DataLoader from torch_geometric.data import Data, Batch from torch_geometric.nn import GINConv, global_mean_pool, BatchNorm from torch_geometric.nn import MLP as PyGMLP from rdkit import Chem, RDLogger RDLogger.logger().setLevel(RDLogger.ERROR) # ─── Data ─────────────────────────────────────────────────────── def load_data(): r = requests.get('https://raw.githubusercontent.com/akiyamalab/cycpeptmp/main/data/CycPeptMPDB_Peptide_All.csv') rows = list(csv.DictReader(io.StringIO(r.content.decode('utf-8-sig')))) r4 = requests.get('https://zenodo.org/records/18754430/files/CycPeptMPDB-4D.csv') d4_map = {int(rr['CycPeptMPDB_ID']): rr for rr in csv.DictReader(io.StringIO(r4.content.decode('utf-8')))} data = [] for row in rows: if not row['PAMPA']: continue mol = Chem.MolFromSmiles(row['SMILES']) if mol is None: continue rid = int(row['CycPeptMPDB_ID']) data.append({**row, 'mol': mol, 'd4': d4_map.get(rid)}) n4d = sum(1 for d in data if d['d4']) print(f'Loaded {len(data)} PAMPA entries ({n4d} with 4D)') return data # ─── Featurization ───────────────────────────────────────────── ATOM_TYPES = [5,6,7,8,9,15,16,17,35,53] AMINO_ACIDS = list('ACDEFGHIKLMNPQRSTVWY') # 20 standard def featurize_mol(mol): mol = Chem.AddHs(mol) atoms = list(mol.GetAtoms()); n = len(atoms) x = [] for a in atoms: t = [0]*len(ATOM_TYPES) if a.GetAtomicNum() in ATOM_TYPES: t[ATOM_TYPES.index(a.GetAtomicNum())] = 1 x.append(t + [a.GetDegree()/4.0, a.GetTotalNumHs()/3.0, a.GetFormalCharge()/1.0, int(a.IsInRing()), int(a.GetIsAromatic())]) x = torch.tensor(x, dtype=torch.float) ei, ea = [], [] for i in range(n): for j in range(i+1, n): b = mol.GetBondBetweenAtoms(i, j) if b is not None: bt = b.GetBondType() e = [int(bt == Chem.rdchem.BondType.SINGLE), int(bt == Chem.rdchem.BondType.DOUBLE), int(bt == Chem.rdchem.BondType.TRIPLE), int(bt == Chem.rdchem.BondType.AROMATIC)] ei.extend([[i,j],[j,i]]); ea.extend([e, e]) ei = torch.tensor(ei, dtype=torch.long).T if ei else torch.zeros((2,0), dtype=torch.long) ea = torch.tensor(ea, dtype=torch.float) if ea else torch.zeros((0,4), dtype=torch.float) return Data(x=x, edge_index=ei, edge_attr=ea) D4_FEAT_KEYS = ['Water_avgRMSD_All','Water_avgRMSD_BackBone','Desolvation_Free_Energy', 'Water_3D_SASA','Water_3D_NPSA','Water_3D_PSA','Hexane_avgRMSD_All', 'Hexane_avgRMSD_BackBone','Hexane_3D_SASA','Hexane_3D_NPSA','Hexane_3D_PSA'] DESC_KEYS = ['MolLogP','MolWt','TPSA','FractionCSP3','NumHAcceptors','NumHDonors', 'NumRotatableBonds','RingCount','HeavyAtomCount','LabuteASA','BertzCT','BalabanJ', 'Kappa1','Kappa2','Kappa3','MolMR','qed','HallKierAlpha', 'NumAliphaticRings','NumAromaticRings','NumSaturatedRings', 'MinPartialCharge','MaxPartialCharge','FpDensityMorgan1','FpDensityMorgan2','FpDensityMorgan3'] def safe_float(v): if v is None: return 0.0 try: v = float(v) return 0.0 if math.isnan(v) or math.isinf(v) else v except: return 0.0 def extract_vec(row, keys): return torch.tensor([safe_float(row.get(k)) for k in keys], dtype=torch.float) def extract_d4(d4r): if d4r is None: return None return extract_vec(d4r, D4_FEAT_KEYS) def seq_to_onehot(seq): """One-hot encode amino acid sequence (monomer-level feature).""" vec = torch.zeros(len(AMINO_ACIDS)) for ch in str(seq).upper() if seq else '': if ch in AMINO_ACIDS: vec[AMINO_ACIDS.index(ch)] += 1 return vec / max(vec.sum(), 1) def enumerate_smiles(mol, n=5): """Generate n random SMILES for augmentation.""" smiles_set = set() for _ in range(n * 3): # try more to get unique s = Chem.MolToSmiles(mol, doRandom=True, canonical=False) if Chem.MolFromSmiles(s) is not None: smiles_set.add(s) if len(smiles_set) >= n: break return list(smiles_set) if smiles_set else [Chem.MolToSmiles(mol)] # ─── Dataset ──────────────────────────────────────────────────── class CycPepDataset(Dataset): def __init__(self, data, t_mean, t_std, augment=1): self.samples = [] for d in data: d4 = extract_d4(d['d4']) if d4 is None: d4 = torch.zeros(len(D4_FEAT_KEYS)) seq = d.get('Sequence', '') seq_oh = seq_to_onehot(seq) desc = extract_vec(d, DESC_KEYS) pyg = featurize_mol(d['mol']) target = (float(d['PAMPA']) - t_mean) / t_std if augment > 1: smiles_list = enumerate_smiles(d['mol'], augment) for s in smiles_list: m = Chem.MolFromSmiles(s) if m is not None: self.samples.append((featurize_mol(m), desc.clone(), seq_oh.clone(), d4.clone(), target)) else: self.samples.append((pyg, desc, seq_oh, d4, target)) def __len__(self): return len(self.samples) def __getitem__(self, i): return self.samples[i] def collate_fn(batch): pygs, descs, seqs, d4s, targets = zip(*batch) return (Batch.from_data_list(list(pygs)), torch.stack(descs), torch.stack(seqs), torch.stack(d4s), torch.tensor(targets, dtype=torch.float)) # ─── Model ────────────────────────────────────────────────────── class CycPepGNN(nn.Module): def __init__(self, node_dim=len(ATOM_TYPES)+5, edge_dim=4, desc_dim=len(DESC_KEYS), seq_dim=len(AMINO_ACIDS), d4_dim=len(D4_FEAT_KEYS), hidden=256): super().__init__() self.node_emb = nn.Linear(node_dim, hidden) self.edge_emb = nn.Linear(edge_dim, hidden) convs = [] for _ in range(5): mlp = PyGMLP([hidden, hidden, hidden], norm='batch_norm') convs.append(GINConv(mlp, train_eps=True)) self.convs = nn.ModuleList(convs) self.bns = nn.ModuleList([BatchNorm(hidden) for _ in range(5)]) self.graph_proj = nn.Linear(hidden, hidden) self.desc_net = nn.Sequential(nn.LayerNorm(desc_dim), nn.Linear(desc_dim, 32), nn.GELU()) self.seq_net = nn.Sequential(nn.LayerNorm(seq_dim), nn.Linear(seq_dim, 32), nn.GELU()) self.d4_net = nn.Sequential(nn.LayerNorm(d4_dim), nn.Linear(d4_dim, 16), nn.GELU()) fusion = hidden + 32 + 32 + 16 self.head = nn.Sequential( nn.Linear(fusion, hidden//2), nn.GELU(), nn.Dropout(0.2), nn.Linear(hidden//2, hidden//4), nn.GELU(), nn.Dropout(0.1), nn.Linear(hidden//4, 1)) def forward(self, pyg_data, desc, seq, d4): x = F.relu(self.node_emb(pyg_data.x)) e = F.relu(self.edge_emb(pyg_data.edge_attr)) xs = [] for conv, bn in zip(self.convs, self.bns): x = F.relu(bn(conv(x, pyg_data.edge_index))) xs.append(x) h_g = self.graph_proj(global_mean_pool(x, pyg_data.batch)) h_d = self.desc_net(desc) h_s = self.seq_net(seq) h_4 = self.d4_net(d4) return self.head(torch.cat([h_g, h_d, h_s, h_4], 1)).squeeze(-1) # ─── Training ──────────────────────────────────────────────────── def train_epoch(model, loader, opt, device): model.train(); total = 0 for batch in loader: g = batch[0].to(device) desc, seq, d4, t = [x.to(device) for x in batch[1:]] opt.zero_grad() loss = F.mse_loss(model(g, desc, seq, d4), t) loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 3.0) opt.step() total += loss.item() * t.size(0) return total / len(loader.dataset) @torch.no_grad() def evaluate(model, loader, device): model.eval(); preds, targets = [], [] for batch in loader: g = batch[0].to(device) desc, seq, d4, t = [x.to(device) for x in batch[1:]] preds.append(model(g, desc, seq, d4).cpu()); targets.append(t.cpu()) p = torch.cat(preds); t = torch.cat(targets) mse = F.mse_loss(p, t).item() return mse, F.l1_loss(p, t).item(), 1 - mse / t.var().item() if t.var().item() > 0 else 0 # ─── Main ──────────────────────────────────────────────────────── def main(): device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') print(f'Device: {device}') data = load_data() ap = torch.tensor([float(d['PAMPA']) for d in data]) tm, ts = ap.mean(), ap.std() print(f'Target: mean={tm:.3f} std={ts:.3f}') results = [] for aug in [1]: name = f'GIN_aug{aug}' ds = CycPepDataset(data, tm, ts, augment=aug) n = len(ds); idx = list(range(n)) random.seed(42); random.shuffle(idx) tr, va = int(0.8*n), int(0.1*n); te = n - tr - va tr_i, va_i, te_i = idx[:tr], idx[tr:tr+va], idx[tr+va:] bs = min(128, n//10) tr_l = DataLoader(torch.utils.data.Subset(ds, tr_i), bs, shuffle=True, collate_fn=collate_fn) va_l = DataLoader(torch.utils.data.Subset(ds, va_i), bs, shuffle=False, collate_fn=collate_fn) te_l = DataLoader(torch.utils.data.Subset(ds, te_i), bs, shuffle=False, collate_fn=collate_fn) model = CycPepGNN(hidden=256).to(device) np_ = sum(p.numel() for p in model.parameters()) print(f'\n{name}: {n} samples, {np_:,} params') opt = torch.optim.AdamW(model.parameters(), lr=3e-4, weight_decay=1e-4) sc = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=150) bv = float('inf'); bt = None; pt_ = 0; t0 = time.time() for ep in range(150): loss = train_epoch(model, tr_l, opt, device) vm, vma, vr2 = evaluate(model, va_l, device) sc.step() if (ep+1) % 15 == 0 or ep == 0: print(f' E{ep+1:3d} loss={loss:.4f} val_mse={vm*ts**2:.4f} val_r2={vr2:.4f}') if vm < bv: bv = vm; bt = evaluate(model, te_l, device); pt_ = 0 else: pt_ += 1 if pt_ >= 25: break te_m_u = bt[0]*ts**2; te_ma_u = bt[1]*ts elapsed = time.time() - t0 print(f' TEST: MSE={te_m_u:.4f} MAE={te_ma_u:.4f} R²={bt[2]:.4f} time={elapsed:.0f}s') results.append((name, te_m_u, te_ma_u, bt[2])) print('\n' + '='*60) print(f'{"Model":<20} {"MSE":<10} {"MAE":<10} {"R²":<10}') print('-'*60) for r in results: print(f'{r[0]:<20} {r[1]:<10.4f} {r[2]:<10.4f} {r[3]:<10.4f}') print('-'*60) print(f'{"MSF-CPMP (SOTA)":<20} {"0.092":<10} {"0.242":<10} {"~0.88":<10}') print(f'{"CycPeptMP":<20} {"0.271":<10} {"0.355":<10} {"0.780":<10}') if __name__ == '__main__': main()