import os import re import torch import torch.nn as nn import numpy as np import gradio as gr from PIL import Image # RDKit imports from rdkit import Chem, RDLogger from rdkit.Chem import rdMolDescriptors, AllChem, Draw # DRFP and ProtT5 imports from drfp import DrfpEncoder from transformers import T5Tokenizer, T5EncoderModel # KAN network from kan import KAN RDLogger.DisableLog("rdApp.*") # ========================================== # 1. Model Architecture # ========================================== class HybridPairKAN(nn.Module): def __init__( self, esm_dim: int, drfp_dim: int, react_dim: int, hidden: int, dropout: float = 0.0 ): super().__init__() self.esm_ln = nn.LayerNorm(esm_dim) self.drfp_ln = nn.LayerNorm(drfp_dim) self.react_ln = nn.LayerNorm(react_dim) in_dim = esm_dim + drfp_dim + react_dim self.bottleneck_dim = hidden self.compressor = nn.Sequential( nn.Linear(in_dim, self.bottleneck_dim), nn.LayerNorm(self.bottleneck_dim), nn.SiLU(), nn.Dropout(dropout) ) self.net = KAN([self.bottleneck_dim, hidden, 1]) def forward( self, esm: torch.Tensor, drfp: torch.Tensor, react: torch.Tensor, update_grid: bool = False, return_latent: bool = False ): e_norm = self.esm_ln(esm) d_norm = self.drfp_ln(drfp) r_norm = self.react_ln(react) x = torch.cat([e_norm, d_norm, r_norm], dim=-1) latent_feat = self.compressor(x) logit = self.net(latent_feat, update_grid=update_grid).squeeze(-1) if return_latent: return logit, latent_feat return logit # ========================================== # 2. Environment Setup # ========================================== device = torch.device("cuda" if torch.cuda.is_available() else "cpu") print(f"🖥️ Current computing device: {device}") PROTT5_DIM = 1024 FP_DIM = 2048 REACT_DIM = 2048 HIDDEN_DIM = 512 MAHA_STAT_PATH = "train_distribution_stat.pt" # ========================================== # 3. Load Mahalanobis Training Statistics # ========================================== def load_maha_stats(path: str): if not os.path.exists(path): print(f"⚠️ Mahalanobis stat file not found: {path}") print(" The app will still run probability prediction, but Mahalanobis distance will be unavailable.") return None stat = torch.load(path, map_location=device) if "mean" not in stat or "inv_cov" not in stat: raise KeyError( "train_distribution_stat.pt must contain keys: 'mean' and 'inv_cov'." ) mu = stat["mean"].float().to(device) inv_cov = stat["inv_cov"].float().to(device) # Allow either [D] or [1, D] mu = mu.view(-1) if inv_cov.ndim != 2 or inv_cov.shape[0] != inv_cov.shape[1]: raise ValueError( f"inv_cov should be a square matrix, but got shape {tuple(inv_cov.shape)}." ) if mu.shape[0] != inv_cov.shape[0]: raise ValueError( f"Mahalanobis mean and inv_cov dimension mismatch: " f"mean={tuple(mu.shape)}, inv_cov={tuple(inv_cov.shape)}." ) if mu.shape[0] != HIDDEN_DIM: print( f"⚠️ Warning: Mahalanobis stat dimension is {mu.shape[0]}, " f"but HIDDEN_DIM is {HIDDEN_DIM}. Please make sure the stat file " f"was generated from the same model architecture." ) print(f"✅ Loaded Mahalanobis stats from {path}") print(f" mean shape: {tuple(mu.shape)} | inv_cov shape: {tuple(inv_cov.shape)}") return { "mean": mu, "inv_cov": inv_cov } maha_stats = load_maha_stats(MAHA_STAT_PATH) def compute_mahalanobis_distance(latent_feat: torch.Tensor): """ latent_feat: [B, hidden_dim] return: [B] """ if maha_stats is None: return None mu = maha_stats["mean"] inv_cov = maha_stats["inv_cov"] if latent_feat.shape[-1] != mu.shape[0]: raise ValueError( f"Latent feature dimension ({latent_feat.shape[-1]}) does not match " f"Mahalanobis stat dimension ({mu.shape[0]})." ) delta = latent_feat.float() - mu.view(1, -1) maha_sq = torch.sum(delta * torch.matmul(delta, inv_cov), dim=-1) maha_dist = torch.sqrt(torch.clamp(maha_sq, min=1e-9)) return maha_dist # ========================================== # 4. Load Multiple Model Ensembles # ========================================== model_categories = { "🧬 General (Pan-Enzyme)": ("binarycls_best_val_seed", ".pt"), "🩸 Cytochrome P450": ("ft_p450_best_seed", ".pt"), "🧪 Phosphatase": ("ft_phosphatase_best_seed", ".pt"), "🌿 Terpene Synthase": ("ft_terpene_best_seed", ".pt") } ensembles = {cat: [] for cat in model_categories.keys()} print("⏳ Loading ensemble model weights for all categories...") for cat, (prefix, ext) in model_categories.items(): for seed in range(40, 45): path = f"{prefix}{seed}{ext}" model = HybridPairKAN( esm_dim=PROTT5_DIM, drfp_dim=FP_DIM, react_dim=REACT_DIM, hidden=HIDDEN_DIM, dropout=0.0 ) if os.path.exists(path): ckpt = torch.load(path, map_location="cpu") model.load_state_dict(ckpt.get("model", ckpt)) model.to(device).eval() ensembles[cat].append(model) else: print(f" ⚠️ File not found: {path}, skipping.") print(f"✅ Loaded {len(ensembles[cat])} models for {cat}") # ========================================== # 5. Feature Extractors # ========================================== prott5_name = "Rostlab/prot_t5_xl_half_uniref50-enc" tokenizer = T5Tokenizer.from_pretrained(prott5_name, do_lower_case=False) if torch.cuda.is_available(): prott5_extractor = T5EncoderModel.from_pretrained( prott5_name, torch_dtype=torch.float16 ).to(device) else: prott5_extractor = T5EncoderModel.from_pretrained(prott5_name).to(device) prott5_extractor.eval() def preprocess_seq_for_prott5(seq: str, max_len: int = 1022) -> str: seq = seq[:max_len] seq = re.sub(r"[UZOBuzob]", "X", seq.upper()) return " ".join(list(seq)) def mean_pool_reps(hidden_states: torch.Tensor, attention_mask: torch.Tensor) -> torch.Tensor: mask_expanded = attention_mask.unsqueeze(-1).float() summed = (hidden_states * mask_expanded).sum(dim=1) denom = mask_expanded.sum(dim=1).clamp(min=1e-9) return summed / denom @torch.no_grad() def get_prott5(sequence: str, max_len: int = 1022) -> torch.Tensor: processed_seq = preprocess_seq_for_prott5(sequence, max_len) inputs = tokenizer( [processed_seq], return_tensors="pt", padding=True, truncation=True, max_length=max_len + 1 ) inputs = {k: v.to(device) for k, v in inputs.items()} if torch.cuda.is_available(): with torch.autocast("cuda", dtype=torch.float16): outputs = prott5_extractor(**inputs) else: outputs = prott5_extractor(**inputs) reps = outputs.last_hidden_state pooled = mean_pool_reps(reps, inputs["attention_mask"]) return pooled.squeeze(0).cpu().float() def get_reactant_morgan_fp( rxn_smiles: str, radius: int = 2, nBits: int = 2048 ) -> np.ndarray: if not isinstance(rxn_smiles, str) or ">>" not in rxn_smiles: return np.zeros((nBits,), dtype=np.float32) reactants_smi = rxn_smiles.split(">>")[0] mol = Chem.MolFromSmiles(reactants_smi) if mol is None: return np.zeros((nBits,), dtype=np.float32) fp = rdMolDescriptors.GetMorganFingerprintAsBitVect( mol, radius, nBits=nBits ) arr = np.zeros((nBits,), dtype=np.float32) Chem.DataStructs.ConvertToNumpyArray(fp, arr) n = np.linalg.norm(arr, keepdims=True) return arr / np.maximum(n, 1e-12) def get_drfp(rxn_smiles: str, nBits: int = 2048) -> np.ndarray: fps = DrfpEncoder.encode([rxn_smiles], n_folded_length=nBits)[0] arr = np.asarray(fps, dtype=np.float32) n = np.linalg.norm(arr, keepdims=True) return arr / np.maximum(n, 1e-12) def smiles_to_reaction_image(rxn_smiles: str): try: rxn = AllChem.ReactionFromSmarts(rxn_smiles, useSmiles=True) if rxn is not None: return Draw.ReactionToImage(rxn) return None except Exception: return None # ========================================== # 6. Inference & UI # ========================================== @torch.no_grad() def predict_interaction( model_choice: str, protein_seq: str, rxn_smiles: str, prob_threshold: float, maha_threshold: float ): protein_seq = protein_seq.strip() rxn_smiles = rxn_smiles.strip() if not protein_seq or not rxn_smiles: return "⚠️ **Please provide both the protein sequence and the reaction SMILES.**", None selected_ensemble = ensembles.get(model_choice, []) if not selected_ensemble: return ( f"🚨 **Error**: No models loaded for category '{model_choice}'. " f"Please check server logs.", None ) rxn_image = smiles_to_reaction_image(rxn_smiles) try: # Extract features prott5_tensor = get_prott5(protein_seq).unsqueeze(0).to(device) drfp_tensor = torch.tensor( get_drfp(rxn_smiles, FP_DIM), dtype=torch.float32 ).unsqueeze(0).to(device) react_tensor = torch.tensor( get_reactant_morgan_fp(rxn_smiles, nBits=REACT_DIM), dtype=torch.float32 ).unsqueeze(0).to(device) # Ensemble inference all_probs = [] all_latents = [] for model in selected_ensemble: logit, latent_feat = model( prott5_tensor, drfp_tensor, react_tensor, return_latent=True ) prob = torch.sigmoid(logit).item() all_probs.append(prob) all_latents.append(latent_feat.detach()) all_probs_np = np.array(all_probs, dtype=np.float32) ensemble_mean = float(np.mean(all_probs_np)) ensemble_latent = torch.stack(all_latents, dim=0).mean(dim=0) # Mahalanobis distance maha_dist_tensor = compute_mahalanobis_distance(ensemble_latent) if maha_dist_tensor is not None: maha_dist = float(maha_dist_tensor.item()) else: maha_dist = None # Epistemic uncertainty by ensemble mutual information eps = 1e-10 entropy_of_mean = -( ensemble_mean * np.log(ensemble_mean + eps) + (1 - ensemble_mean) * np.log(1 - ensemble_mean + eps) ) entropy_of_preds = -( all_probs_np * np.log(all_probs_np + eps) + (1 - all_probs_np) * np.log(1 - all_probs_np + eps) ) mean_of_entropy = float(np.mean(entropy_of_preds)) ensemble_mi = float(entropy_of_mean - mean_of_entropy) percentage = ensemble_mean * 100 prob_pass = ensemble_mean >= prob_threshold maha_available = maha_dist is not None maha_pass = maha_available and (maha_dist <= maha_threshold) # Format output prob_color = "#2e7d32" if prob_pass else "#c62828" result_md = f"### 📊 Ensemble Prediction ({model_choice} | N={len(selected_ensemble)})\n\n" result_md += ( f"**Mean Match Probability**: " f"" f"**{percentage:.2f}%**\n\n" ) result_md += f"**Epistemic Uncertainty (MI)**: **{ensemble_mi:.5f}**\n\n" if maha_available: maha_color = "#2e7d32" if maha_pass else "#c62828" result_md += ( f"**Mahalanobis Distance**: " f"" f"**{maha_dist:.4f}**\n\n" ) result_md += f"**Mahalanobis Threshold**: **{maha_threshold:.2f}**\n\n" else: result_md += ( "⚠️ **Mahalanobis Distance**: unavailable. " "`train_distribution_stat.pt` was not found or not loaded.\n\n" ) result_md += "---\n#### 💡 Diagnostic Report\n\n" if maha_available: if prob_pass and maha_pass: result_md += ( "✅ **Conclusion**: The model predicts this enzyme–reaction pair as " "**compatible**, and its latent representation is close to the learned " "training distribution.\n\n" ) elif prob_pass and not maha_pass: result_md += ( "⚠️ **Conclusion**: The predicted probability is high, but the pair is " "**far from the learned training distribution**. This candidate should be " "treated as lower reliability or possible OOD.\n\n" ) elif not prob_pass and maha_pass: result_md += ( "❌ **Conclusion**: The pair appears to be within the learned distribution, " "but the predicted compatibility is low.\n\n" ) else: result_md += ( "🔴 **Conclusion**: The predicted compatibility is low, and the pair is " "also far from the learned training distribution.\n\n" ) else: if prob_pass: result_md += ( "✅ **Conclusion**: The model ensemble leans toward this pair being " "**compatible**, but no latent-space distribution check is available.\n\n" ) else: result_md += ( "❌ **Conclusion**: The model ensemble leans toward this pair being " "**incompatible**, but no latent-space distribution check is available.\n\n" ) if ensemble_mi < 0.05: result_md += "🟢 **Ensemble Agreement**: High. The models are in strong agreement." elif ensemble_mi < 0.15: result_md += "🟡 **Ensemble Agreement**: Moderate. The models show some disagreement." else: result_md += ( "🔴 **Ensemble Agreement**: Low. The models disagree substantially, " "so the prediction should be interpreted cautiously." ) return result_md, rxn_image except Exception as e: return f"🚨 **Inference Error**: {str(e)}", rxn_image # ========================================== # 7. Gradio UI # ========================================== demo = gr.Interface( fn=predict_interaction, inputs=[ gr.Dropdown( choices=list(model_categories.keys()), value="🧬 General (Pan-Enzyme)", label="1. Select Prediction Model", info=( "Use the general model for broad screening, or choose a specialized " "model when the enzyme family is known." ) ), gr.Textbox( lines=4, placeholder="e.g. MTEYKLVVVG...", label="2. Enzyme Protein Sequence" ), gr.Textbox( lines=2, placeholder="e.g. C(C)=O>>C(C)O", label="3. Reaction SMILES" ), gr.Slider( minimum=0.01, maximum=0.99, value=0.30, step=0.01, label="4. Probability Decision Threshold" ), gr.Slider( minimum=1.0, maximum=80.0, value=20.0, step=0.5, label="5. Mahalanobis Distance Threshold" ) ], outputs=[ gr.Markdown(label="Prediction Analysis Panel"), gr.Image(type="pil", label="Reaction Visualization") ], title="🧬 EZHit: Enzyme–Reaction Catalytic Potential Predictor", description=( "EZHit estimates enzyme–reaction compatibility using an ensemble model and " "adds a latent-space distribution check based on Mahalanobis distance. " "A high probability suggests potential catalytic compatibility, while a lower " "Mahalanobis distance suggests that the input pair is closer to the learned " "training distribution." ), examples=[ [ "🧬 General (Pan-Enzyme)", "MTEYKPTVRLATSQERENPTINLADMLKNRGIGLGIAFSSMGGAWGKGGIGGLGLAIAGWGLGGLAIGYLGGAWGKGGIGGLGLAIAGWGLGGLAIGYL", "C=C(C(C)C)CC[C@@H](C)C1CCC2C3=CC=C4C[C@@H](O)CC[C@]4(C)C3CC[C@@]21C>>C=C(C(C)C)CC[C@@H](C)C1CCC2C3CC=C4C[C@@H](O)CC[C@]4(C)C3CC[C@@]21C", 0.30, 20.0 ], [ "🩸 Cytochrome P450", "MTEYKPTVRLATSQERENPTINLADMLKNRGIGLGIAFSSMGGAWGKGGIGGLGLAIAGWGLGGLAIGYLGGAWGKGGIGGLGLAIAGWGLGGLAIGYL", "C=C(C(C)C)CC[C@@H](C)C1CCC2C3=CC=C4C[C@@H](O)CC[C@]4(C)C3CC[C@@]21C>>C=C(C(C)C)CC[C@@H](C)C1CCC2C3CC=C4C[C@@H](O)CC[C@]4(C)C3CC[C@@]21C", 0.30, 20.0 ] ], flagging_mode="never" ) demo.launch()