CycPepGNN / train.py
devansh0703's picture
Upload folder using huggingface_hub
d15fd98 verified
Raw
History Blame Contribute Delete
13.9 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_add_pool, global_mean_pool, BatchNorm
from torch_geometric.nn import MLP as PyGMLP
from rdkit import Chem, RDLogger
from rdkit.Chem import AllChem
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 = {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.get(rid), 'id': rid})
n4d = sum(1 for d in data if d['d4'])
print(f'Loaded {len(data)} PAMPA entries ({n4d} with 4D)')
return data
# ─── Graph construction ─────────────────────────────────────────
ATOM_TYPES = [5,6,7,8,9,15,16,17,35,53]
def mol_to_pyg(mol):
mol = Chem.AddHs(mol)
atoms = list(mol.GetAtoms())
n = len(atoms)
x = []
for a in atoms:
feat = []
# Atom type one-hot
t = [0]*len(ATOM_TYPES)
if a.GetAtomicNum() in ATOM_TYPES:
t[ATOM_TYPES.index(a.GetAtomicNum())] = 1
feat.extend(t)
feat += [a.GetDegree()/4.0, a.GetTotalNumHs()/3.0, a.GetFormalCharge()/1.0,
int(a.IsInRing()), int(a.GetIsAromatic())]
x.append(feat)
x = torch.tensor(x, dtype=torch.float)
edge_index, edge_attr = [], []
for i in range(n):
for j in range(i+1, n):
bond = mol.GetBondBetweenAtoms(i, j)
if bond is not None:
bt = bond.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)]
edge_index.extend([[i,j],[j,i]])
edge_attr.extend([e, e])
edge_index = torch.tensor(edge_index, dtype=torch.long).T if edge_index else torch.zeros((2,0), dtype=torch.long)
edge_attr = torch.tensor(edge_attr, dtype=torch.float) if edge_attr else torch.zeros((0,4), dtype=torch.float)
return Data(x=x, edge_index=edge_index, edge_attr=edge_attr)
# ─── Features ───────────────────────────────────────────────────
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 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)
# ─── Dataset ────────────────────────────────────────────────────
class CycPepDataset(Dataset):
def __init__(self, data, t_mean, t_std, use_4d=True):
self.samples = []
for d in data:
d4 = extract_d4(d['d4'])
if use_4d and d4 is None: continue
if d4 is None: d4 = torch.zeros(len(D4_FEAT_KEYS))
pyg = mol_to_pyg(d['mol'])
desc = extract_vec(d, DESC_KEYS)
pc = extract_vec(d, PHYSCHEM_KEYS)
target = (float(d['PAMPA']) - t_mean) / t_std
self.samples.append((pyg, desc, pc, d4, target))
def __len__(self): return len(self.samples)
def __getitem__(self, i): return self.samples[i]
def collate_fn(batch):
pygs, descs, pcs, d4s, targets = zip(*batch)
batch_pyg = Batch.from_data_list(list(pygs))
return (batch_pyg,
torch.stack(descs), torch.stack(pcs), 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), pc_dim=len(PHYSCHEM_KEYS),
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)
nn1 = PyGMLP([hidden, hidden, hidden], batch_norm=True)
nn2 = PyGMLP([hidden, hidden, hidden], batch_norm=True)
nn3 = PyGMLP([hidden, hidden, hidden], batch_norm=True)
nn4 = PyGMLP([hidden, hidden, hidden], batch_norm=True)
self.convs = nn.ModuleList([
GINConv(nn1, train_eps=True),
GINConv(nn2, train_eps=True),
GINConv(nn3, train_eps=True),
GINConv(nn4, train_eps=True),
])
self.bns = nn.ModuleList([BatchNorm(hidden) for _ in range(4)])
self.graph_proj = nn.Linear(hidden * 4, hidden)
self.desc_net = nn.Sequential(nn.Linear(desc_dim, 64), nn.GELU(), nn.LayerNorm(64))
self.pc_net = nn.Sequential(nn.Linear(pc_dim, 32), nn.GELU(), nn.LayerNorm(32))
self.d4_net = nn.Sequential(nn.Linear(d4_dim, 16), nn.GELU(), nn.LayerNorm(16))
fusion = hidden + 64 + 32 + 16
self.head = nn.Sequential(
nn.Linear(fusion, 256), nn.GELU(), nn.Dropout(0.15),
nn.Linear(256, 128), nn.GELU(), nn.Dropout(0.1),
nn.Linear(128, 1))
def forward(self, pyg_data, desc, pc, 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 = conv(x, pyg_data.edge_index)
x = bn(x)
x = F.relu(x)
xs.append(x)
# Readout: concat all layer outputs
h_graph = torch.cat([global_mean_pool(x, pyg_data.batch) for x in xs], -1)
# Re-project to hidden
h_graph = self.graph_proj(h_graph)
h_desc = self.desc_net(desc)
h_pc = self.pc_net(pc)
h_d4 = self.d4_net(d4)
h = torch.cat([h_graph, h_desc, h_pc, h_d4], 1)
return self.head(h).squeeze(-1)
# ─── Training ────────────────────────────────────────────────────
def train_epoch(model, loader, opt, device):
model.train()
total = 0
for batch in loader:
pyg = batch[0].to(device)
desc, pc, d4, t = [x.to(device) for x in batch[1:]]
opt.zero_grad()
loss = F.mse_loss(model(pyg, desc, pc, d4), t)
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 5.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:
pyg = batch[0].to(device)
desc, pc, d4, t = [x.to(device) for x in batch[1:]]
preds.append(model(pyg, desc, pc, d4).cpu())
targets.append(t.cpu())
preds = torch.cat(preds); targets = torch.cat(targets)
mse = F.mse_loss(preds, targets).item()
mae = F.l1_loss(preds, targets).item()
r2 = 1 - mse / targets.var().item() if targets.var().item() > 0 else 0
return mse, mae, r2
# ─── Main ────────────────────────────────────────────────────────
def main():
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
print(f'Device: {device}')
data = load_data()
all_pampa = torch.tensor([float(d['PAMPA']) for d in data])
t_mean, t_std = all_pampa.mean(), all_pampa.std()
print(f'Target: mean={t_mean:.3f} std={t_std:.3f}')
def run_experiment(name, use_4d, hidden=256, epochs=200):
ds = CycPepDataset(data, t_mean, t_std, use_4d=use_4d)
n = len(ds)
indices = list(range(n))
random.seed(42); random.shuffle(indices)
tr, va = int(0.8*n), int(0.1*n)
te = n - tr - va
tr_i, va_i, te_i = indices[:tr], indices[tr:tr+va], indices[tr+va:]
tr_ds = torch.utils.data.Subset(ds, tr_i)
va_ds = torch.utils.data.Subset(ds, va_i)
te_ds = torch.utils.data.Subset(ds, te_i)
bs = min(64, n//10)
tr_ld = DataLoader(tr_ds, bs, shuffle=True, collate_fn=collate_fn)
va_ld = DataLoader(va_ds, bs, shuffle=False, collate_fn=collate_fn)
te_ld = DataLoader(te_ds, bs, shuffle=False, collate_fn=collate_fn)
model = CycPepGNN(hidden=hidden).to(device)
n_p = sum(p.numel() for p in model.parameters())
print(f'{name}: {n:,} samples, {n_p:,} params')
opt = torch.optim.AdamW(model.parameters(), lr=5e-4, weight_decay=1e-5)
sched = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=epochs)
best_val = float('inf')
best_te = None
patience = 0
t0 = time.time()
for ep in range(epochs):
loss = train_epoch(model, tr_ld, opt, device)
vm, vma, vr2 = evaluate(model, va_ld, device)
sched.step()
vm_u = vm * t_std.item()**2
if (ep+1) % 20 == 0 or ep == 0:
print(f' E{ep+1:3d} loss={loss:.4f} val_mse={vm_u:.4f} val_r2={vr2:.4f}')
if vm < best_val:
best_val = vm
best_te = evaluate(model, te_ld, device)
patience = 0
else:
patience += 1
if patience >= 30:
print(f' Early stop at E{ep+1}')
break
elapsed = time.time() - t0
te_m_u = best_te[0] * t_std.item()**2
te_ma_u = best_te[1] * t_std.item()
print(f' TEST: MSE={te_m_u:.4f} MAE={te_ma_u:.4f} RΒ²={best_te[2]:.4f} time={elapsed:.0f}s')
return best_te
results = []
# Full dataset (no 4D)
res = run_experiment('GNN_Full', use_4d=False, hidden=256)
results.append(('GNN_Full', *[res[0]*t_std.item()**2, res[1]*t_std.item(), res[2]]))
# With 4D
res2 = run_experiment('GNN_4D', use_4d=True, hidden=256)
results.append(('GNN_4D', *[res2[0]*t_std.item()**2, res2[1]*t_std.item(), res2[2]]))
print('\n' + '='*60)
print('SUMMARY:')
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'{"MultiCycPermea":<20} {"0.160":<10} {"0.280":<10} {"~0.75":<10}')
if __name__ == '__main__':
main()