| import os |
| import sys |
| import json |
| import torch |
| import numpy as np |
| from fastapi import FastAPI, HTTPException |
| from fastapi.middleware.cors import CORSMiddleware |
| from pydantic import BaseModel |
| from typing import List, Optional |
| from huggingface_hub import hf_hub_download |
| from rdkit import Chem |
|
|
| |
| sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), '..', 'workspace', 'src'))) |
| sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), 'workspace', 'src'))) |
|
|
| from model import MultiTaskEvidentialGNN |
| from dataloader import smiles_to_graph |
| from explain import explain_molecule, render_molecular_explanation |
| from scaffold import generate_synthetic_analogs |
|
|
| app = FastAPI(title="Explainable & Uncertainty-Aware Multi-Task GNN API") |
|
|
| |
| app.add_middleware( |
| CORSMiddleware, |
| allow_origins=["*"], |
| allow_credentials=True, |
| allow_methods=["*"], |
| allow_headers=["*"], |
| ) |
|
|
| |
| model = None |
| device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') |
|
|
| def load_evidential_model(): |
| global model |
| try: |
| model = MultiTaskEvidentialGNN( |
| node_in_dim=31, |
| edge_in_dim=6, |
| hidden_dim=64, |
| num_reg_tasks=2, |
| num_clf_tasks=2 |
| ) |
| |
| |
| checkpoint_path = "workspace/models/best_checkpoint.pt" |
| if not os.path.exists(checkpoint_path): |
| checkpoint_path = "../workspace/models/best_checkpoint.pt" |
| |
| if not os.path.exists(checkpoint_path): |
| print("[Backend] Local checkpoint not found. Downloading from Hugging Face Hub...") |
| token_val = None |
| token_path = "/home/anamitra/Downloads/API_Keys_and_Secrets/hf_token" |
| if os.path.exists(token_path): |
| with open(token_path, 'r') as f: |
| token_val = f.read().strip() |
| |
| try: |
| |
| checkpoint_path = hf_hub_download( |
| repo_id="Arko007/multi-task-evidential-gnn", |
| filename="best_checkpoint.pt", |
| token=token_val |
| ) |
| print(f"[Backend] Checkpoint downloaded successfully to: {checkpoint_path}") |
| except Exception as hf_err: |
| print(f"[Backend] HF download failed: {hf_err}. Attempting to run uninitialized (fallback only).") |
| return |
| |
| |
| model.load_state_dict(torch.load(checkpoint_path, map_location=device)) |
| model.to(device) |
| model.eval() |
| print("[Backend] Evidential multi-task model loaded successfully.") |
| except Exception as e: |
| print(f"[Backend] ERROR loading GNN model: {e}") |
|
|
| @app.on_event("startup") |
| def startup_event(): |
| load_evidential_model() |
| |
| if not os.environ.get("GROQ_API_KEY"): |
| key_path = "/home/anamitra/Downloads/API_Keys_and_Secrets/groq_api.txt" |
| if os.path.exists(key_path): |
| with open(key_path, 'r') as f: |
| os.environ["GROQ_API_KEY"] = f.read().strip() |
|
|
| |
| class SmilesRequest(BaseModel): |
| smiles: str |
|
|
| class AnalogRequest(BaseModel): |
| smiles: str |
| num_variants: Optional[int] = 30 |
| reaction_smarts: Optional[str] = None |
|
|
| class GroqAuditRequest(BaseModel): |
| smiles: str |
| predictions: dict |
| important_atoms: List[str] |
|
|
| |
| @app.post("/api/predict") |
| def predict_smiles(req: SmilesRequest): |
| if model is None: |
| raise HTTPException(status_code=500, detail="GNN model is not loaded. Please train it first.") |
| |
| mol = Chem.MolFromSmiles(req.smiles) |
| if not mol: |
| raise HTTPException(status_code=400, detail="Invalid SMILES string provided.") |
| |
| data = smiles_to_graph(req.smiles) |
| if data is None: |
| raise HTTPException(status_code=400, detail="Graph extraction failed for SMILES.") |
| |
| data = data.to(device) |
| |
| with torch.no_grad(): |
| |
| batch = torch.zeros(data.x.size(0), dtype=torch.long, device=device) |
| pred_reg, pred_clf = model(data.x, data.edge_index, data.edge_attr, batch) |
| |
| |
| |
| caco2_gamma, caco2_nu, caco2_alpha, caco2_beta = pred_reg[0, 0].cpu().numpy() |
| sol_gamma, sol_nu, sol_alpha, sol_beta = pred_reg[0, 1].cpu().numpy() |
| |
| caco2_aleatoric = float(np.sqrt(caco2_beta / (caco2_alpha - 1.0))) |
| caco2_epistemic = float(1.0 / np.sqrt(caco2_nu)) |
| |
| sol_aleatoric = float(np.sqrt(sol_beta / (sol_alpha - 1.0))) |
| sol_epistemic = float(1.0 / np.sqrt(sol_nu)) |
| |
| |
| |
| herg_alpha = pred_clf[0, 0].cpu().numpy() |
| ames_alpha = pred_clf[0, 1].cpu().numpy() |
| |
| herg_S = float(np.sum(herg_alpha)) |
| herg_prob = float(herg_alpha[1] / herg_S) |
| herg_epistemic = float(2.0 / herg_S) |
| |
| ames_S = float(np.sum(ames_alpha)) |
| ames_prob = float(ames_alpha[1] / ames_S) |
| ames_epistemic = float(2.0 / ames_S) |
| |
| return { |
| "smiles": req.smiles, |
| "caco2": float(caco2_gamma), |
| "caco2_aleatoric": caco2_aleatoric, |
| "caco2_epistemic": caco2_epistemic, |
| "solubility": float(sol_gamma), |
| "solubility_aleatoric": sol_aleatoric, |
| "solubility_epistemic": sol_epistemic, |
| "herg": herg_prob, |
| "herg_epistemic": herg_epistemic, |
| "ames": ames_prob, |
| "ames_epistemic": ames_epistemic |
| } |
|
|
| |
| @app.post("/api/explain") |
| def explain_smiles(req: SmilesRequest): |
| if model is None: |
| raise HTTPException(status_code=500, detail="GNN model is not loaded.") |
|
|
| mol = Chem.MolFromSmiles(req.smiles) |
| if not mol: |
| raise HTTPException(status_code=400, detail="Invalid SMILES.") |
|
|
| try: |
| data = smiles_to_graph(req.smiles) |
| data = data.to(device) |
|
|
| |
| |
| atom_w, edge_w = explain_molecule(model, data, task_type='clf', task_idx=0, epochs=80) |
|
|
| svg_dir = "workspace/explanations" |
| os.makedirs(svg_dir, exist_ok=True) |
| temp_svg_path = os.path.join(svg_dir, "temp_explain.svg") |
|
|
| |
| render_molecular_explanation(req.smiles, atom_w, edge_w, data, temp_svg_path, is_toxic_or_bad=True) |
|
|
| |
| with open(temp_svg_path, 'r') as f: |
| svg_content = f.read() |
|
|
| |
| atom_symbols = [] |
| for idx, w in enumerate(atom_w): |
| atom_symbol = mol.GetAtomWithIdx(idx).GetSymbol() |
| atom_symbols.append({"index": idx, "symbol": atom_symbol, "weight": float(w)}) |
|
|
| |
| atom_symbols = sorted(atom_symbols, key=lambda x: x["weight"], reverse=True) |
|
|
| return { |
| "svg": svg_content, |
| "attributions": atom_symbols[:5] |
| } |
| except Exception as e: |
| import traceback |
| print(f"[Backend] /api/explain failed: {e}\n{traceback.format_exc()}") |
| raise HTTPException(status_code=500, detail=f"Explanation generation failed: {e}") |
|
|
| |
| @app.post("/api/generate") |
| def generate_analogs_route(req: AnalogRequest): |
| if model is None: |
| raise HTTPException(status_code=500, detail="GNN model is not loaded.") |
| |
| analogs = generate_synthetic_analogs(req.smiles, req.reaction_smarts, req.num_variants) |
| if not analogs: |
| return {"analogs": []} |
| |
| results = [] |
| |
| for smiles in analogs: |
| data = smiles_to_graph(smiles) |
| if data is None: |
| continue |
| data = data.to(device) |
| |
| with torch.no_grad(): |
| batch = torch.zeros(data.x.size(0), dtype=torch.long, device=device) |
| pred_reg, pred_clf = model(data.x, data.edge_index, data.edge_attr, batch) |
| |
| caco2_gamma = float(pred_reg[0, 0, 0].cpu().item()) |
| sol_gamma = float(pred_reg[0, 1, 0].cpu().item()) |
| |
| herg_alpha = pred_clf[0, 0].cpu().numpy() |
| herg_prob = float(herg_alpha[1] / np.sum(herg_alpha)) |
| |
| ames_alpha = pred_clf[0, 1].cpu().numpy() |
| ames_prob = float(ames_alpha[1] / np.sum(ames_alpha)) |
| |
| results.append({ |
| "smiles": smiles, |
| "caco2": caco2_gamma, |
| "solubility": sol_gamma, |
| "herg": herg_prob, |
| "ames": ames_prob, |
| |
| "safety_score": 1.0 - max(herg_prob, ames_prob) |
| }) |
| |
| |
| results = sorted(results, key=lambda x: (x["safety_score"], x["solubility"]), reverse=True) |
| |
| return {"analogs": results} |
|
|
| |
| @app.post("/api/explain_text") |
| def explain_predictions_groq_route(req: GroqAuditRequest): |
| groq_key = os.environ.get("GROQ_API_KEY") |
| if not groq_key: |
| raise HTTPException(status_code=500, detail="Groq API Key is not set in the server environment.") |
| |
| try: |
| from groq import Groq |
| client = Groq(api_key=groq_key) |
| |
| prompt = f""" |
| You are an expert computational toxicologist and medicinal chemist. |
| A multi-task Graph Neural Network (GNN) trained on Tox21 and ADME datasets has evaluated the following molecule: |
| SMILES: {req.smiles} |
| |
| Predictions: |
| - Cardiac Toxicity (hERG Blockade): {"Toxic (hERG blocker)" if req.predictions['herg'] > 0.5 else "Safe (Non-blocker)"} (Confidence probability: {req.predictions['herg']:.2f}, Epistemic Uncertainty: {req.predictions['herg_epistemic']:.3f}) |
| - Mutagenicity (AMES test): {"Mutagenic" if req.predictions['ames'] > 0.5 else "Non-mutagenic"} (Confidence probability: {req.predictions['ames']:.2f}, Epistemic Uncertainty: {req.predictions['ames_epistemic']:.3f}) |
| - Intestinal Permeability (Caco-2 LogPapp): {req.predictions['caco2']:.3f} cm/s (Epistemic Uncertainty: {req.predictions['caco2_epistemic']:.3f}) |
| - Aqueous Solubility (LogS): {req.predictions['solubility']:.3f} (Epistemic Uncertainty: {req.predictions['solubility_epistemic']:.3f}) |
| |
| GNNExplainer highlighted the following atoms/regions as highly influential for predictions: {req.important_atoms} |
| |
| Based on these predictions, provide a brief, professional chemical audit report explaining: |
| 1. The biological implications of these ADMET predictions. |
| 2. How the highlighted structural features (e.g. aromatic rings, hydrogen bond donors/acceptors, polar groups) correlate with the predicted toxicity and ADME profiles. |
| 3. Recommendations for modifying the chemical structure to improve solubility or reduce toxicity (e.g. adding specific polar groups, reducing lipophilicity, or using bioisosteres). |
| |
| Be concise (maximum 3 paragraphs). Use professional medicinal chemistry terms. Do not repeat the prompt. |
| """ |
| response = client.chat.completions.create( |
| model="llama-3.3-70b-versatile", |
| messages=[{"role": "user", "content": prompt}], |
| temperature=0.2, |
| max_tokens=800 |
| ) |
| report = response.choices[0].message.content |
| return {"report": report} |
| except Exception as e: |
| raise HTTPException(status_code=500, detail=f"Groq generation failed: {e}") |
|
|
| |
| @app.get("/api/diagnostics") |
| def get_diagnostics(): |
| metrics_path = "workspace/models/test_metrics.json" |
| if not os.path.exists(metrics_path): |
| metrics_path = "../workspace/models/test_metrics.json" |
| |
| if os.path.exists(metrics_path): |
| with open(metrics_path, 'r') as f: |
| metrics = json.load(f) |
| return metrics |
| else: |
| |
| return { |
| "reg_Caco-2_mae": 0.354, |
| "reg_Caco-2_spearman": 0.684, |
| "reg_Caco-2_mca": 0.082, |
| "reg_Solubility_mae": 0.412, |
| "reg_Solubility_spearman": 0.725, |
| "reg_Solubility_mca": 0.076, |
| "clf_hERG_accuracy": 0.845, |
| "clf_hERG_mca": 0.048, |
| "clf_AMES_accuracy": 0.812, |
| "clf_AMES_mca": 0.061, |
| "info": "Placeholder values (Model has not been trained locally yet. Run training to generate local metrics)." |
| } |
|
|
| from fastapi.staticfiles import StaticFiles |
| dist_path = os.path.abspath(os.path.join(os.path.dirname(__file__), 'dist')) |
| if os.path.exists(dist_path): |
| app.mount("/", StaticFiles(directory=dist_path, html=True), name="static") |
|
|
|
|