CycPepGNN / train2.py
devansh0703's picture
Upload folder using huggingface_hub
d15fd98 verified
Raw
History Blame Contribute Delete
9.94 kB
import csv, io, requests, random, time, math, sys
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 rdkit import Chem, RDLogger
from rdkit.Chem import AllChem
from transformers import AutoTokenizer, AutoModel
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
# ─── Features ───────────────────────────────────────────────────
PHYSCHEM_KEYS = ['MolLogP','MolWt','TPSA','FractionCSP3','NumHAcceptors','NumHDonors',
'NumRotatableBonds','RingCount','HeavyAtomCount','LabuteASA',
'HallKierAlpha','Kappa1','Kappa2','Kappa3','BertzCT','BalabanJ']
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)
CHEM_DESC = [
'MolLogP','MolWt','TPSA','FractionCSP3','NumHAcceptors','NumHDonors',
'NumRotatableBonds','RingCount','HeavyAtomCount','LabuteASA',
'NumAliphaticRings','NumAromaticRings','NumSaturatedRings',
'qed','BertzCT','BalabanJ','HallKierAlpha',
'MinPartialCharge','MaxPartialCharge','MinAbsPartialCharge','MaxAbsPartialCharge',
'NumValenceElectrons','NHOHCount','NOCount',
'Kappa1','Kappa2','Kappa3','MolMR',
'FpDensityMorgan1','FpDensityMorgan2','FpDensityMorgan3',
]
def compute_morgan(mol, bits=2048):
fp = AllChem.GetMorganFingerprintAsBitVect(mol, 3, nBits=bits)
return torch.tensor(fp, dtype=torch.float)
def extract_desc(row):
return extract_vec(row, CHEM_DESC)
# Pretrained model for SMILES
print('Loading ChemBERTa-2 tokenizer/model...')
tok = AutoTokenizer.from_pretrained('seyonec/PubChem10M_SMILES_BPE_450k')
chemberta = AutoModel.from_pretrained('seyonec/PubChem10M_SMILES_BPE_450k')
chemberta.eval()
for p in chemberta.parameters():
p.requires_grad = False
chem_dim = 768
print(f'Model loaded (dim={chem_dim})')
@torch.no_grad()
def smiles_embed(smiles):
inputs = tok(smiles, return_tensors='pt', padding=True, truncation=True, max_length=128)
if torch.cuda.is_available():
inputs = {k: v.cuda() for k, v in inputs.items()}
chemberta.cuda()
outputs = chemberta(**inputs)
emb = outputs.last_hidden_state[:,0,:] # CLS token
return emb.cpu()
# ─── Dataset ────────────────────────────────────────────────────
class CycPepDataset(Dataset):
def __init__(self, data, t_mean, t_std):
self.samples = []
for d in data:
d4 = extract_d4(d['d4'])
if d4 is None: d4 = torch.zeros(len(D4_FEAT_KEYS))
fp = compute_morgan(d['mol'])
desc = extract_desc(d)
target = (float(d['PAMPA']) - t_mean) / t_std
self.samples.append((d['SMILES'], fp, desc, d4, target))
# Precompute ChemBERTa embeddings
all_smiles = [s[0] for s in self.samples]
self.chem_embs = []
bs = 64
for i in range(0, len(all_smiles), bs):
batch_smiles = all_smiles[i:i+bs]
emb = smiles_embed(batch_smiles)
self.chem_embs.append(emb)
self.chem_embs = torch.cat(self.chem_embs, 0)
print(f'Precomputed ChemBERTa embeddings: {self.chem_embs.shape}')
# Replace SMILES with embeddings
self.samples = [(self.chem_embs[i], fp, desc, d4, t)
for i, (_, fp, desc, d4, t) in enumerate(self.samples)]
def __len__(self): return len(self.samples)
def __getitem__(self, i): return self.samples[i]
def collate_fn(batch):
chem, fp, desc, d4, targets = zip(*batch)
return (torch.stack(chem), torch.stack(fp), torch.stack(desc),
torch.stack(d4), torch.tensor(targets, dtype=torch.float))
# ─── Model ──────────────────────────────────────────────────────
class CycPepModel(nn.Module):
def __init__(self, chem_dim=768, fp_dim=2048, desc_dim=len(CHEM_DESC),
d4_dim=len(D4_FEAT_KEYS), hidden=256):
super().__init__()
self.chem_net = nn.Sequential(nn.LayerNorm(chem_dim), nn.Linear(chem_dim, 128), nn.GELU())
self.fp_net = nn.Sequential(nn.LayerNorm(fp_dim), nn.Linear(fp_dim, 128), nn.GELU())
self.desc_net = nn.Sequential(nn.LayerNorm(desc_dim), nn.Linear(desc_dim, 64), nn.GELU())
self.d4_net = nn.Sequential(nn.LayerNorm(d4_dim), nn.Linear(d4_dim, 32), nn.GELU())
fusion = 128 + 128 + 64 + 32
self.head = nn.Sequential(
nn.Linear(fusion, hidden), nn.GELU(), nn.Dropout(0.3),
nn.Linear(hidden, hidden//2), nn.GELU(), nn.Dropout(0.2),
nn.Linear(hidden//2, 1))
def forward(self, chem, fp, desc, d4):
h = torch.cat([self.chem_net(chem), self.fp_net(fp),
self.desc_net(desc), self.d4_net(d4)], 1)
return self.head(h).squeeze(-1)
# ─── Training ────────────────────────────────────────────────────
def train_epoch(model, loader, opt, device):
model.train()
total = 0
for chem, fp, desc, d4, t in loader:
chem, fp, desc, d4, t = [x.to(device) for x in (chem, fp, desc, d4, t)]
opt.zero_grad()
loss = F.mse_loss(model(chem, fp, desc, 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 chem, fp, desc, d4, t in loader:
chem, fp, desc, d4 = [x.to(device) for x in (chem, fp, desc, d4)]
preds.append(model(chem, fp, desc, 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
def run(data, name, epochs=150):
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
all_pampa = torch.tensor([float(d['PAMPA']) for d in data])
t_mean, t_std = all_pampa.mean(), all_pampa.std()
ds = CycPepDataset(data, t_mean, t_std)
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:]
bs = 64
tr_ld = DataLoader(torch.utils.data.Subset(ds, tr_i), bs, shuffle=True, collate_fn=collate_fn)
va_ld = DataLoader(torch.utils.data.Subset(ds, va_i), bs, shuffle=False, collate_fn=collate_fn)
te_ld = DataLoader(torch.utils.data.Subset(ds, te_i), bs, shuffle=False, collate_fn=collate_fn)
model = CycPepModel().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=3e-4, weight_decay=1e-4)
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) % 15 == 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: 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 te_m_u, te_ma_u, best_te[2]
def main():
data = load_data()
results = []
results.append(run(data, 'ChemBERTa+FP+Desc+4D'))
print('\n' + '='*50)
for r in results:
print(f' MSE={r[0]:.4f} MAE={r[1]:.4f} RΒ²={r[2]:.4f}')
print(f' MSF-CPMP: MSE=0.092 MAE=0.242 RΒ²~0.88')
if __name__ == '__main__':
main()