CoolFace
Apppublic

ahmedmiloudi/BioTechLabAI

sourceHugging Faceapache-2.0updated 8mo agoView on Hugging Face
1likes
real_admet_predictor.py470 linesDownload Raw Back to src
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']