ahmedmiloudi/BioTechLabAI
1
1"""2Real ADMET Predictor using ADMET-AI pretrained models3ADMET-AI is the actual pip-installable package with 41 ADMET endpoints4"""5 6import pandas as pd7import numpy as np8from typing import List, Dict, Optional9import logging10from rdkit import Chem11from rdkit.Chem import Descriptors12 13logger = logging.getLogger(__name__)14 15 16# ============================================================================17# ADMET-AI Predictor (REAL pip package with pretrained models)18# ============================================================================19 20class ADMETAIPredictor:21 """22 Real ADMET predictions using ADMET-AI pretrained models23 Package: pip install admet-ai24 41 ADMET endpoints - NO TRAINING NEEDED!25 """26 27 def __init__(self):28 """Initialize ADMET-AI predictor"""29 self.available = False30 self.predictor = None31 32 try:33 from admet_ai import ADMETModel34 35 # Load pretrained model (downloads automatically first time)36 logger.info("Loading ADMET-AI pretrained models...")37 self.predictor = ADMETModel()38 self.predictor = ADMETModel(drugbank_path=None) 39 self.available = True40 logger.info("✓ ADMET-AI loaded successfully (41 pretrained endpoints)")41 42 except ImportError:43 logger.warning("ADMET-AI not available. Install: pip install admet-ai")44 self.available = False45 except Exception as e:46 logger.error(f"Failed to initialize ADMET-AI: {e}")47 self.available = False48 49 def predict2(self, smiles_list: List[str]) -> pd.DataFrame:50 """51 Predict ADMET properties using ADMET-AI pretrained models52 53 Returns DataFrame with 41 ADMET endpoints including:54 - Absorption (Caco-2, HIA, Pgp, etc.)55 - Distribution (BBB, PPB, VDss)56 - Metabolism (CYP inhibition)57 - Excretion (Clearance, Half-life)58 - Toxicity (hERG, AMES, DILI, etc.)59 """60 if not self.available or self.predictor is None:61 logger.error("ADMET-AI not available")62 return self._fallback_predictions(smiles_list)63 64 try:65 # ADMET-AI prediction (returns DataFrame with all 41 endpoints)66 logger.info(f"Running ADMET-AI predictions on {len(smiles_list)} compounds...")67 # ADMET-AI prediction - FIXED FOR BATCH PREDICTIONS68 69 70 # Process each SMILES individually71 # all_predictions = []72 # for smile in smiles_list:73 # pred = self.predictor.predict(smiles=smile)74 # all_predictions.append(pred)75 76 #predictions = pd.concat([self.predictor.predict(smiles=s) for s in smiles_list])77 78 predictions = self.predictor.predict(smiles=smiles_list)79 80 # Add SMILES column if not present81 if 'SMILES' not in predictions.columns:82 predictions.insert(0, 'SMILES', smiles_list[:len(predictions)])83 84 # Add derived metrics85 predictions['Lipinski_Pass'] = predictions.apply(86 lambda row: self._check_lipinski(row.get('SMILES', row.name)), axis=187 )88 89 # Simplify BBB to categorical90 if 'BBB_Martins' in predictions.columns:91 predictions['BBB_Class'] = predictions['BBB_Martins'].apply(92 lambda x: 'Permeant' if x > 0.5 else 'Non-permeant'93 )94 95 # Calculate overall druglikeness96 predictions['Druglikeness_Score'] = predictions.apply(97 lambda row: self._calculate_druglikeness(row), axis=198 )99 100 # Rename columns for clarity101 column_mapping = {102 'BBB_Martins': 'BBB_Permeability',103 'Caco2_Wang': 'Caco2_Permeability',104 'HIA_Hou': 'Intestinal_Absorption',105 'Pgp_Broccatelli': 'Pgp_Substrate',106 'Clearance_Hepatocyte_AZ': 'Clearance',107 'Half_Life_Obach': 'Half_Life',108 'hERG': 'hERG_Blocker',109 'AMES': 'AMES_Toxicity',110 'DILI': 'DILI_Risk',111 }112 113 predictions.rename(columns=column_mapping, inplace=True)114 115 logger.info(f"✓ ADMET-AI predictions completed for {len(predictions)} compounds")116 return predictions117 118 except Exception as e:119 logger.error(f"ADMET-AI prediction failed: {e}")120 return self._fallback_predictions(smiles_list)121 122 123 def predict(self, smiles_list: List[str]) -> pd.DataFrame:124 """125 Predict ADMET properties using ADMET-AI pretrained models126 """127 if not self.available or self.predictor is None:128 logger.error("ADMET-AI not available")129 return self._fallback_predictions(smiles_list)130 131 try:132 logger.info(f"Running ADMET-AI predictions on {len(smiles_list)} compounds...")133 134 # Get predictions from ADMET-AI135 raw_predictions = self.predictor.predict(smiles=smiles_list)136 137 print(f"DEBUG 1: raw_predictions type: {type(raw_predictions)}")138 139 # Handle different return types from ADMET-AI140 if isinstance(raw_predictions, dict):141 print(f"DEBUG 2: Got dict, converting to DataFrame")142 predictions = pd.DataFrame([raw_predictions])143 else:144 print(f"DEBUG 2: Got DataFrame, shape: {raw_predictions.shape}")145 predictions = raw_predictions.reset_index(drop=True)146 147 print(f"DEBUG 3: predictions shape: {predictions.shape}")148 print(f"DEBUG 3: predictions columns: {predictions.columns.tolist()}")149 150 # Add SMILES column if not present151 if 'SMILES' not in predictions.columns:152 predictions.insert(0, 'SMILES', smiles_list[:len(predictions)])153 154 # DEBUG: Check BBB_Martins before processing155 if 'BBB_Martins' in predictions.columns:156 print(f"DEBUG 4: BBB_Martins type: {type(predictions['BBB_Martins'].iloc[0])}")157 print(f"DEBUG 4: BBB_Martins value: {predictions['BBB_Martins'].iloc[0]}")158 print(f"DEBUG 4: BBB_Martins has len? {hasattr(predictions['BBB_Martins'].iloc[0], '__len__')}")159 160 # FIX: Handle arrays in the DataFrame161 # Convert any array columns to scalars162 for col in predictions.columns:163 # Check if column contains arrays (numpy arrays)164 if not predictions.empty and col != 'SMILES':165 sample = predictions[col].iloc[0] if len(predictions) > 0 else None166 if sample is not None and hasattr(sample, '__len__') and not isinstance(sample, str):167 print(f"DEBUG 5: Converting array column: {col}")168 # Extract first element from arrays169 predictions[col] = predictions[col].apply(170 lambda x: x[0] if hasattr(x, '__len__') and len(x) > 0 else x171 )172 173 # Add derived metrics174 print(f"DEBUG 6: Adding Lipinski_Pass")175 predictions['Lipinski_Pass'] = predictions['SMILES'].apply(176 lambda smile: self._check_lipinski(smile)177 )178 179 # Simplify BBB to categorical180 if 'BBB_Martins' in predictions.columns:181 print(f"DEBUG 7: Adding BBB_Class")182 predictions['BBB_Class'] = predictions['BBB_Martins'].apply(183 lambda x: 'Permeant' if float(x) > 0.5 else 'Non-permeant'184 )185 186 # Calculate overall druglikeness187 print(f"DEBUG 8: Calculating Druglikeness_Score")188 predictions['Druglikeness_Score'] = predictions.apply(189 lambda row: self._calculate_druglikeness(row), axis=1190 )191 192 # Rename columns for clarity193 column_mapping = {194 'BBB_Martins': 'BBB_Permeability',195 'Caco2_Wang': 'Caco2_Permeability',196 'HIA_Hou': 'Intestinal_Absorption',197 'Pgp_Broccatelli': 'Pgp_Substrate',198 'Clearance_Hepatocyte_AZ': 'Clearance',199 'Half_Life_Obach': 'Half_Life',200 'hERG': 'hERG_Blocker',201 'AMES': 'AMES_Toxicity',202 'DILI': 'DILI_Risk',203 }204 205 predictions.rename(columns=column_mapping, inplace=True)206 207 logger.info(f"✓ ADMET-AI predictions completed for {len(predictions)} compounds")208 return predictions209 210 except Exception as e:211 logger.error(f"ADMET-AI prediction failed: {e}")212 import traceback213 traceback.print_exc()214 return self._fallback_predictions(smiles_list)215 216 217 218 def _check_lipinski(self, smiles: str) -> bool:219 """Check Lipinski's Rule of Five"""220 try:221 mol = Chem.MolFromSmiles(smiles)222 if not mol:223 return False224 225 mw = Descriptors.MolWt(mol)226 logp = Descriptors.MolLogP(mol)227 hbd = Descriptors.NumHDonors(mol)228 hba = Descriptors.NumHAcceptors(mol)229 230 violations = sum([mw > 500, logp > 5, hbd > 5, hba > 10])231 return violations <= 1 # Allow 1 violation232 except:233 return False234 235 def _calculate_druglikeness(self, row: pd.Series) -> float:236 """Calculate overall druglikeness score (0-1)"""237 score = 0.0238 239 # Lipinski (20%)240 if row.get('Lipinski_Pass', False):241 score += 0.2242 243 # Absorption (15%)244 if row.get('Intestinal_Absorption', 0) > 0.3:245 score += 0.15246 247 # BBB for CNS drugs (10%)248 if row.get('BBB_Permeability', 0) > 0.3:249 score += 0.1250 251 # Not Pgp substrate (10%)252 if row.get('Pgp_Substrate', 1) < 0.5:253 score += 0.1254 255 # Cardiac safety - not hERG blocker (15%)256 if row.get('hERG_Blocker', 1) < 0.5:257 score += 0.15258 259 # Not AMES toxic (15%)260 if row.get('AMES_Toxicity', 1) < 0.5:261 score += 0.15262 263 # Not hepatotoxic (15%)264 if row.get('DILI_Risk', 1) < 0.5:265 score += 0.15266 267 return min(1.0, score)268 269 def _fallback_predictions(self, smiles_list: List[str]) -> pd.DataFrame:270 """Fallback using RDKit descriptors"""271 logger.warning("Using RDKit fallback (not pretrained ML)")272 273 results = []274 for smiles in smiles_list:275 try:276 mol = Chem.MolFromSmiles(smiles)277 if mol:278 row = {279 'SMILES': smiles,280 'MW': Descriptors.MolWt(mol),281 'LogP': Descriptors.MolLogP(mol),282 'HBD': Descriptors.NumHDonors(mol),283 'HBA': Descriptors.NumHAcceptors(mol),284 'TPSA': Descriptors.TPSA(mol),285 'Lipinski_Pass': self._check_lipinski(smiles),286 'Note': 'RDKit fallback - Install ADMET-AI for ML predictions'287 }288 else:289 row = {'SMILES': smiles, 'Error': 'Invalid SMILES'}290 results.append(row)291 except Exception as e:292 results.append({'SMILES': smiles, 'Error': str(e)})293 294 return pd.DataFrame(results)295 296 297# ============================================================================298# DeepChem ADMET Models (Alternative - also pretrained)299# ============================================================================300 301class DeepChemADMETPredictor:302 """303 ADMET predictions using DeepChem pretrained models304 Already installed if you have deepchem for Tox21305 """306 307 def __init__(self):308 """Initialize DeepChem ADMET predictor"""309 self.available = False310 self.models = {}311 312 try:313 import deepchem as dc314 self.available = True315 logger.info("DeepChem available for ADMET predictions")316 317 # Load pretrained models318 self._load_pretrained_models()319 320 except ImportError:321 logger.warning("DeepChem not available")322 self.available = False323 324 def _load_pretrained_models(self):325 """Load pretrained DeepChem models"""326 import deepchem as dc327 328 # Available pretrained models329 model_names = {330 'bbbp': 'Blood-Brain Barrier',331 'clearance': 'Clearance',332 'sider': 'Side Effects',333 }334 335 for name, description in model_names.items():336 try:337 # Load from DeepChem model zoo338 tasks, datasets, transformers = getattr(dc.molnet, f'load_{name}')()339 logger.info(f"✓ Loaded {description} dataset")340 self.models[name] = {341 'tasks': tasks,342 'datasets': datasets,343 'transformers': transformers344 }345 except Exception as e:346 logger.warning(f"Failed to load {name}: {e}")347 348 def predict(self, smiles_list: List[str]) -> pd.DataFrame:349 """Predict ADMET using DeepChem models"""350 if not self.available or not self.models:351 return pd.DataFrame()352 353 results = []354 for smiles in smiles_list:355 row = {'SMILES': smiles}356 357 # Add predictions from available models358 # (Implementation depends on which models loaded successfully)359 360 results.append(row)361 362 return pd.DataFrame(results)363 364 365# ============================================================================366# Unified ADMET Predictor (Auto-selects best available)367# ============================================================================368 369class RealADMETPredictor:370 """371 Unified ADMET predictor - uses best available pretrained model372 Priority: ADMET-AI > DeepChem > RDKit fallback373 """374 375 def __init__(self):376 """Initialize with best available model"""377 self.predictor = None378 self.method = None379 380 # Try ADMET-AI first (best - 41 endpoints)381 try:382 predictor = ADMETAIPredictor()383 if predictor.available:384 self.predictor = predictor385 self.method = "ADMET-AI (41 endpoints, pretrained)"386 logger.info("✓ Using ADMET-AI pretrained models")387 return388 except Exception as e:389 logger.warning(f"ADMET-AI initialization failed: {e}")390 391 # Try DeepChem second392 try:393 predictor = DeepChemADMETPredictor()394 if predictor.available and predictor.models:395 self.predictor = predictor396 self.method = "DeepChem ADMET (pretrained)"397 logger.info("✓ Using DeepChem pretrained models")398 return399 except Exception as e:400 logger.warning(f"DeepChem ADMET initialization failed: {e}")401 402 # Fallback to RDKit (not ML)403 logger.warning("⚠️ Using RDKit descriptors (not pretrained ML)")404 self.method = "RDKit Descriptors (Not ML - Install ADMET-AI)"405 self.predictor = None406 407 def predict(self, smiles_list: List[str]) -> pd.DataFrame:408 """Predict ADMET properties"""409 if self.predictor:410 return self.predictor.predict(smiles_list)411 else:412 # RDKit fallback413 return self._rdkit_fallback(smiles_list)414 415 def _rdkit_fallback(self, smiles_list: List[str]) -> pd.DataFrame:416 """Basic RDKit predictions (not ML)"""417 from rdkit.Chem import QED418 419 results = []420 for smiles in smiles_list:421 try:422 mol = Chem.MolFromSmiles(smiles)423 if mol:424 results.append({425 'SMILES': smiles,426 'MW': Descriptors.MolWt(mol),427 'LogP': Descriptors.MolLogP(mol),428 'HBD': Descriptors.NumHDonors(mol),429 'HBA': Descriptors.NumHAcceptors(mol),430 'TPSA': Descriptors.TPSA(mol),431 'RotBonds': Descriptors.NumRotatableBonds(mol),432 'Lipinski_Pass': self._check_lipinski(mol),433 'Druglikeness_Score': QED.qed(mol),434 'BBB_Predicted': self._predict_bbb_simple(mol),435 'Note': '⚠️ RDKit fallback - Install ADMET-AI for real ML predictions'436 })437 else:438 results.append({'SMILES': smiles, 'Error': 'Invalid SMILES'})439 except Exception as e:440 results.append({'SMILES': smiles, 'Error': str(e)})441 442 return pd.DataFrame(results)443 444 def _check_lipinski(self, mol) -> bool:445 """Check Lipinski's Rule of Five"""446 mw = Descriptors.MolWt(mol)447 logp = Descriptors.MolLogP(mol)448 hbd = Descriptors.NumHDonors(mol)449 hba = Descriptors.NumHAcceptors(mol)450 violations = sum([mw > 500, logp > 5, hbd > 5, hba > 10])451 return violations <= 1452 453 def _predict_bbb_simple(self, mol) -> str:454 """Simple BBB heuristic (not ML)"""455 tpsa = Descriptors.TPSA(mol)456 logp = Descriptors.MolLogP(mol)457 if tpsa < 90 and logp > 0:458 return "Likely Permeant"459 return "Unlikely Permeant"460 461 def get_method_info(self) -> str:462 """Get information about prediction method"""463 return self.method464 465 466# ============================================================================467# Export for easy import468# ============================================================================469 470__all__ = ['RealADMETPredictor', 'ADMETAIPredictor', 'DeepChemADMETPredictor']