CycPepGNN / train_comprehensive.py
devansh0703's picture
Upload folder using huggingface_hub
d15fd98 verified
Raw
History Blame Contribute Delete
13.4 kB
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 ───────────────────────────────────────────────────────
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))
# Build dataset
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"])}')
# ─── Model ──────────────────────────────────────────────────────
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)
# ─── Dataset ────────────────────────────────────────────────────
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))
# ─── Training ────────────────────────────────────────────────────
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
# Test
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
# ─── Experiments ────────────────────────────────────────────────
print('\n====== EXPERIMENTS ======\n')
# Required comparsion: 4D subset WITHOUT vs WITH 4D features
# This is the key result: showing 4D helps on the same subset
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}')