Download train2.py from devansh0703/CycPepGNN: direct link, hf CLI and curl.
- Browser
- Download file 9.94 kB
-
https://huggingface.co/devansh0703/CycPepGNN/resolve/main/train2.py
- Command line
-
hf download hf://devansh0703/CycPepGNN/train2.py
-
curl -L -o train2.py https://huggingface.co/devansh0703/CycPepGNN/resolve/main/train2.py
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})') | |
| 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) | |
| 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() | |