CycPepGNN / train_final.py
devansh0703's picture
Upload folder using huggingface_hub
d15fd98 verified
Raw
History Blame Contribute Delete
11.7 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 ───────────────────────────────────────────────────────
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()