| 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) |
|
|
| |
| print('Loading 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')))} |
|
|
| ATOM_TYPES = [5,6,7,8,9,15,16,17,35,53] |
| AMINO_ACIDS = list('ACDEFGHIKLMNPQRSTVWY') |
| PHYSCHEM_KEYS = ['MolLogP','MolWt','TPSA','FractionCSP3','NumHAcceptors','NumHDonors', |
| 'NumRotatableBonds','RingCount','HeavyAtomCount','LabuteASA', |
| 'HallKierAlpha','Kappa1','Kappa2','Kappa3','BertzCT','BalabanJ'] |
| DESC_KEYS = ['BCUT2D_MWHI','BCUT2D_MWLOW','BCUT2D_CHGHI','BCUT2D_CHGLO','BCUT2D_LOGPHI', |
| 'BCUT2D_LOGPLOW','BCUT2D_MRHI','BCUT2D_MRLOW','BalabanJ','BertzCT', |
| 'Chi0','Chi0n','Chi0v','Chi1','Chi1n','Chi1v','Chi2n','Chi2v','Chi3n','Chi3v','Chi4n','Chi4v', |
| 'HallKierAlpha','Kappa1','Kappa2','Kappa3','LabuteASA', |
| 'PEOE_VSA1','PEOE_VSA10','PEOE_VSA11','PEOE_VSA12','PEOE_VSA13','PEOE_VSA14', |
| 'PEOE_VSA2','PEOE_VSA3','PEOE_VSA4','PEOE_VSA5','PEOE_VSA6','PEOE_VSA7','PEOE_VSA8','PEOE_VSA9', |
| 'SMR_VSA1','SMR_VSA10','SMR_VSA2','SMR_VSA3','SMR_VSA4','SMR_VSA5','SMR_VSA6','SMR_VSA7','SMR_VSA8','SMR_VSA9', |
| 'SlogP_VSA1','SlogP_VSA10','SlogP_VSA11','SlogP_VSA12','SlogP_VSA2','SlogP_VSA3','SlogP_VSA4', |
| 'SlogP_VSA5','SlogP_VSA6','SlogP_VSA7','SlogP_VSA8','SlogP_VSA9', |
| 'EState_VSA1','EState_VSA10','EState_VSA11','EState_VSA2','EState_VSA3','EState_VSA4', |
| 'EState_VSA5','EState_VSA6','EState_VSA7','EState_VSA8','EState_VSA9', |
| 'VSA_EState1','VSA_EState10','VSA_EState2','VSA_EState3','VSA_EState4','VSA_EState5', |
| 'VSA_EState6','VSA_EState7','VSA_EState8','VSA_EState9', |
| 'MolLogP','MolWt','TPSA','FractionCSP3','HeavyAtomCount','NHOHCount','NOCount', |
| 'NumAliphaticCarbocycles','NumAliphaticHeterocycles','NumAliphaticRings', |
| 'NumAromaticCarbocycles','NumAromaticHeterocycles','NumAromaticRings', |
| 'NumHAcceptors','NumHDonors','NumHeteroatoms','NumRotatableBonds','RingCount', |
| 'NumSaturatedCarbocycles','NumSaturatedHeterocycles','NumSaturatedRings', |
| 'MolMR','qed','MaxAbsEStateIndex','MaxEStateIndex','MinEStateIndex', |
| 'MaxPartialCharge','MinPartialCharge','MaxAbsPartialCharge','MinAbsPartialCharge', |
| 'FpDensityMorgan1','FpDensityMorgan2','FpDensityMorgan3', |
| 'ExactMolWt','HeavyAtomMolWt','NumValenceElectrons'] |
| 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'] |
|
|
| def sf(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 ev(r, ks): return torch.tensor([sf(r.get(k)) for k in ks], dtype=torch.float) |
|
|
| def seq_onehot(seq): |
| v = torch.zeros(len(AMINO_ACIDS)) |
| for c in str(seq).upper() if seq else '': |
| if c in AMINO_ACIDS: v[AMINO_ACIDS.index(c)] += 1 |
| return v / max(v.sum(), 1) |
|
|
| def mol_to_pyg(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(), 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]) |
| return Data(x=x, edge_index=torch.tensor(ei,dtype=torch.long).T if ei else torch.zeros((2,0),dtype=torch.long), |
| edge_attr=torch.tensor(ea,dtype=torch.float) if ea else torch.zeros((0,4),dtype=torch.float)) |
|
|
| |
| all_data = [] |
| for row in rows: |
| if not row['PAMPA']: continue |
| mol = Chem.MolFromSmiles(row['SMILES']) |
| if mol is None: continue |
| d4 = d4_map.get(int(row['CycPeptMPDB_ID'])) |
| all_data.append({**row, 'mol': mol, 'd4': d4, 'has_4d': d4 is not None}) |
|
|
| pa = torch.tensor([float(d['PAMPA']) for d in all_data]) |
| tm, ts = pa.mean(), pa.std() |
| print(f'Total {len(all_data)} PAMPA samples, target mean={tm:.3f} std={ts:.3f}') |
| print(f'With 4D: {sum(1 for d in all_data if d["has_4d"])}') |
|
|
| |
| class CycPepGNN(nn.Module): |
| def __init__(self, nd=len(ATOM_TYPES)+5, ed=4, dd=len(DESC_KEYS), pd=len(PHYSCHEM_KEYS), |
| sd=len(AMINO_ACIDS), d4d=len(D4_FEAT_KEYS), h=256, use_4d=True, use_seq=True): |
| super().__init__() |
| self.use_4d = use_4d; self.use_seq = use_seq |
| self.ne = nn.Linear(nd, h); self.ee = nn.Linear(ed, h) |
| c = nn.ModuleList() |
| for _ in range(4): |
| c.append(GINConv(PyGMLP([h,h,h], norm='batch_norm'), train_eps=True)) |
| self.cs = c; self.bns = nn.ModuleList([BatchNorm(h) for _ in range(4)]) |
| self.gp = nn.Linear(h*4, h) |
| self.dn = nn.Sequential(nn.Linear(dd, 64), nn.GELU(), nn.LayerNorm(64)) |
| self.pn = nn.Sequential(nn.Linear(pd, 32), nn.GELU(), nn.LayerNorm(32)) |
| f = h + 64 + 32 |
| if use_seq: self.sn = nn.Sequential(nn.Linear(sd, 32), nn.GELU(), nn.LayerNorm(32)); f += 32 |
| if use_4d: self.d4n = nn.Sequential(nn.Linear(d4d, 16), nn.GELU(), nn.LayerNorm(16)); f += 16 |
| self.head = nn.Sequential( |
| nn.Linear(f, h//2), nn.GELU(), nn.Dropout(0.15), |
| nn.Linear(h//2, h//4), nn.GELU(), nn.Dropout(0.1), |
| nn.Linear(h//4, 1)) |
|
|
| def forward(self, pg, desc, pc, seq=None, d4=None): |
| x = F.relu(self.ne(pg.x)); xs = [] |
| for c, bn in zip(self.cs, self.bns): |
| x = F.relu(bn(c(x, pg.edge_index))); xs.append(x) |
| h = self.gp(torch.cat([global_mean_pool(x, pg.batch) for x in xs], -1)) |
| h = torch.cat([h, self.dn(desc), self.pn(pc)], 1) |
| if self.use_seq: h = torch.cat([h, self.sn(seq)], 1) |
| if self.use_4d: h = torch.cat([h, self.d4n(d4)], 1) |
| return self.head(h).squeeze(-1) |
|
|
| |
| class CycPepDataset(Dataset): |
| def __init__(self, data, tm, ts, use_4d=True, augment=1, only_4d=False): |
| self.samples = [] |
| for d in data: |
| if only_4d and not d['has_4d']: continue |
| d4v = ev(d['d4'], D4_FEAT_KEYS) if d['d4'] else torch.zeros(len(D4_FEAT_KEYS)) |
| seqq = seq_onehot(d.get('Sequence','')) |
| desc = ev(d, DESC_KEYS); pc = ev(d, PHYSCHEM_KEYS) |
| tgt = (float(d['PAMPA']) - tm) / ts |
| self.add_sample(d['mol'], desc, pc, seqq, d4v, tgt, augment) |
| def add_sample(self, mol, desc, pc, seq, d4, tgt, aug): |
| self.samples.append((mol_to_pyg(mol), desc, pc, seq, d4, tgt)) |
| if aug > 1: |
| seen = set() |
| for _ in range(aug * 5): |
| s = Chem.MolToSmiles(mol, doRandom=True, canonical=False) |
| if s in seen: continue; seen.add(s) |
| m2 = Chem.MolFromSmiles(s) |
| if m2 is not None and len(seen) <= aug: |
| self.samples.append((mol_to_pyg(m2), desc, pc, seq, d4, tgt)) |
| if len(seen) >= aug: break |
| def __len__(self): return len(self.samples) |
| def __getitem__(self, i): return self.samples[i] |
|
|
| def coll(batch): |
| pg, d, pc, sq, d4, t = zip(*batch) |
| return (Batch.from_data_list(list(pg)), torch.stack(d), torch.stack(pc), |
| torch.stack(sq), torch.stack(d4), torch.tensor(t, dtype=torch.float)) |
|
|
| |
| def run(name, use_4d, use_seq, augment, only_4d, data, tm, ts, seeds=[42]): |
| ds = CycPepDataset(data, tm, ts, use_4d=use_4d, augment=augment, only_4d=only_4d) |
| n = len(ds); idx = list(range(n)); random.seed(42); random.shuffle(idx) |
| tr, va = int(0.8*n), int(0.1*n) |
| tr_i, va_i, te_i = idx[:tr], idx[tr:tr+va], idx[tr+va:] |
|
|
| res = [] |
| for seed in seeds: |
| torch.manual_seed(seed); random.seed(seed) |
| model = CycPepGNN(use_4d=use_4d, use_seq=use_seq).cuda() |
| np_ = sum(p.numel() for p in model.parameters()) |
| opt = torch.optim.AdamW(model.parameters(), lr=5e-4, weight_decay=1e-5) |
| sc = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=200) |
| bs = min(128, n//10) |
|
|
| tr_l = DataLoader(torch.utils.data.Subset(ds, tr_i), bs, shuffle=True, collate_fn=coll) |
| va_l = DataLoader(torch.utils.data.Subset(ds, va_i), bs, shuffle=False, collate_fn=coll) |
| te_l = DataLoader(torch.utils.data.Subset(ds, te_i), bs, shuffle=False, collate_fn=coll) |
|
|
| bv = float('inf'); bt = None; pt = 0 |
| for ep in range(200): |
| model.train() |
| for pg, d, pc, sq, d4, tgt in tr_l: |
| pg = pg.cuda(); d, pc, sq, d4, tgt = [x.cuda() for x in (d, pc, sq, d4, tgt)] |
| opt.zero_grad() |
| kw = {}; kw2 = {} |
| if use_seq: kw['seq'] = sq; kw2['sq'] = sq |
| if use_4d: kw['d4'] = d4; kw2['d4'] = d4 |
| F.mse_loss(model(pg, d, pc, **kw), tgt).backward() |
| torch.nn.utils.clip_grad_norm_(model.parameters(), 5.0); opt.step() |
| model.eval(); pv, tv = [], [] |
| for pg, d, pc, sq, d4, tgt in va_l: |
| pg = pg.cuda(); d, pc, sq, d4, tgt = [x.cuda() for x in (d, pc, sq, d4, tgt)] |
| with torch.no_grad(): |
| kw = {}; kw2 = {} |
| if use_seq: kw['seq'] = sq |
| if use_4d: kw['d4'] = d4 |
| pv.append(model(pg, d, pc, **kw).cpu()); tv.append(tgt.cpu()) |
| vm = F.mse_loss(torch.cat(tv), torch.cat(pv)).item() |
| sc.step() |
| if vm < bv: bv = vm; pt = 0 |
| else: pt += 1 |
| if pt >= 25: break |
| |
| model.eval(); ps, ts_ = [], [] |
| for pg, d, pc, sq, d4, tgt in te_l: |
| pg = pg.cuda(); d, pc, sq, d4, tgt = [x.cuda() for x in (d, pc, sq, d4, tgt)] |
| with torch.no_grad(): |
| kw = {}; kw2 = {} |
| if use_seq: kw['seq'] = sq |
| if use_4d: kw['d4'] = d4 |
| ps.append(model(pg, d, pc, **kw).cpu()); ts_.append(tgt.cpu()) |
| p = torch.cat(ps); t = torch.cat(ts_) |
| te_m = F.mse_loss(p, t).item() * ts**2 |
| te_ma = F.l1_loss(p, t).item() * ts |
| te_r2 = 1 - F.mse_loss(p, t).item() / t.var().item() if t.var().item() > 0 else 0 |
| res.append((te_m, te_ma, te_r2)) |
| print(f' [{name}] seed={seed}: MSE={te_m:.4f} MAE={te_ma:.4f} RΒ²={te_r2:.4f} (n={n})') |
| return res |
|
|
| |
| print('\n====== EXPERIMENTS ======\n') |
|
|
| |
| |
| r_baseline = run('4Dsub(no4D)', use_4d=False, use_seq=False, augment=1, only_4d=True, data=all_data, tm=tm, ts=ts) |
| r_with4d = run('4Dsub(+4D)', use_4d=True, use_seq=False, augment=1, only_4d=True, data=all_data, tm=tm, ts=ts) |
| r_with4d_seq = run('4Dsub(+4D+Seq)', use_4d=True, use_seq=True, augment=1, only_4d=True, data=all_data, tm=tm, ts=ts) |
| r_fullpower = run('4Dsub(+4D+Seq+Aug)', use_4d=True, use_seq=True, augment=3, only_4d=True, data=all_data, tm=tm, ts=ts) |
|
|
| print('\n' + '='*70) |
| print(f'{"Experiment":<30} {"MSE":>8} {"MAE":>8} {"RΒ²":>8} {"n":>6}') |
| print('='*70) |
| for name, res in [('4Dsub(no4D)',r_baseline),('4Dsub(+4D)',r_with4d), |
| ('4Dsub(+4D+Seq)',r_with4d_seq),('4Dsub(+4D+Seq+Aug)',r_fullpower)]: |
| m = np.mean([r[0] for r in res]) |
| ma = np.mean([r[1] for r in res]) |
| r2m = np.mean([r[2] for r in res]) |
| print(f'{name:<30} {m:>8.4f} {ma:>8.4f} {r2m:>8.4f}') |
| print('='*70) |
| print(f'{"MSF-CPMP (SOTA)":<30} {"0.092":>8} {"0.242":>8} {"~0.88":>8}') |
| print(f'{"CycPeptMP":<30} {"0.271":>8} {"0.355":>8} {"0.780":>8}') |
|
|