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()