ClearDDI_AlertSystem / backend /inference.py
malikachaaban
app
22c0953
Raw
History Blame Contribute Delete
18.7 kB
"""
ClearDDI — moteur d'inférence
Charge les artefacts réels du pipeline PolypharmaSafe (XGBoost binaire + multiclasse,
SHAP, mappings nom -> DrugBank ID -> SMILES) et expose les fonctions utilisées par l'API.
Prototype académique d'aide à la décision — non destiné à un usage clinique réel.
"""
from __future__ import annotations
import json
import pickle
import re
import difflib
import warnings
from pathlib import Path
from functools import lru_cache
import numpy as np
import joblib
import shap
from rdkit import Chem
from rdkit.Chem import rdFingerprintGenerator, AllChem, Draw
warnings.filterwarnings("ignore")
ARTIFACTS_DIR = Path(__file__).resolve().parent.parent / "artifacts"
# ──────────────────────────────────────────────────────────────────────────
# Chargement des artefacts (une seule fois, au démarrage du serveur)
# ──────────────────────────────────────────────────────────────────────────
def _load_artifacts():
with open(ARTIFACTS_DIR / "name_to_id.pkl", "rb") as f:
name_to_id = pickle.load(f)
with open(ARTIFACTS_DIR / "id_to_smiles.pkl", "rb") as f:
id_to_smiles = pickle.load(f)
with open(ARTIFACTS_DIR / "inference_config.json", encoding="utf-8") as f:
config = json.load(f)
with open(ARTIFACTS_DIR / "metrics.json", encoding="utf-8") as f:
metrics = json.load(f)
binary_model = joblib.load(ARTIFACTS_DIR / "ddi_binary_model.joblib")
multi_model = joblib.load(ARTIFACTS_DIR / "ddi_multiclass_model.joblib")
label_encoder = joblib.load(ARTIFACTS_DIR / "ddi_multiclass_label_encoder.joblib")
return {
"name_to_id": name_to_id,
"id_to_smiles": id_to_smiles,
"config": config,
"metrics": metrics,
"binary_model": binary_model,
"multi_model": multi_model,
"label_encoder": label_encoder,
}
_A = _load_artifacts()
NAME_TO_ID: dict[str, str] = _A["name_to_id"]
ID_TO_SMILES: dict[str, str] = _A["id_to_smiles"]
CONFIG: dict = _A["config"]
METRICS: dict = _A["metrics"]
BINARY_MODEL = _A["binary_model"]
MULTI_MODEL = _A["multi_model"]
LABEL_ENCODER = _A["label_encoder"]
FP_RADIUS = CONFIG["fp_radius"]
FP_BITS = CONFIG["fp_bits"]
DECISION_THRESHOLD = CONFIG.get("decision_threshold", 0.5)
_MFPGEN = rdFingerprintGenerator.GetMorganGenerator(radius=FP_RADIUS, fpSize=FP_BITS)
_SHAP_EXPLAINER = shap.TreeExplainer(BINARY_MODEL)
# Liste triée des noms connus, pour la recherche / les suggestions
_ALL_NAMES_SORTED = sorted(NAME_TO_ID.keys())
# Alias courants (français / noms commerciaux fréquents) vers le nom DrugBank
# officiel utilisé comme clé dans name_to_id. Complète la recherche exacte,
# car la base utilise la nomenclature DrugBank (souvent le nom anglais de la
# substance active).
_COMMON_ALIASES = {
"aspirine": "acetylsalicylic acid",
"aspirin": "acetylsalicylic acid",
"acide acetylsalicylique": "acetylsalicylic acid",
"warfarine": "warfarin",
"paracetamol": "acetaminophen",
"paracétamol": "acetaminophen",
"doliprane": "acetaminophen",
"tylenol": "acetaminophen",
}
# ──────────────────────────────────────────────────────────────────────────
# Connaissance métier additionnelle : libellés et effets secondaires typiques
# par type clinique. Construite manuellement à partir des 11 classes du
# modèle multiclasse (le texte source DrugBank n'est pas inclus dans le
# bundle de déploiement) — niveau prototype académique, à valider par un
# professionnel de santé avant tout usage réel.
# ──────────────────────────────────────────────────────────────────────────
CLINICAL_TYPE_INFO = {
"Hémorragique": {
"label": "Risque hémorragique",
"description": "Effet additif ou synergique sur l'hémostase, augmentant le risque de saignement.",
"side_effects": [
"Risque hémorragique accru",
"Ecchymoses, saignements de nez",
"Hémorragie digestive (cas sévères)",
"Allongement du temps de coagulation",
],
},
"Cardiaque": {
"label": "Risque cardiaque (rythme)",
"description": "Effet potentiel sur la conduction cardiaque ou le rythme (allongement du QTc, arythmies).",
"side_effects": [
"Allongement de l'intervalle QTc",
"Risque de troubles du rythme (tachycardie, bradycardie)",
"Palpitations",
],
},
"Cardiovasculaire": {
"label": "Risque cardiovasculaire (tension)",
"description": "Interaction affectant la pression artérielle ou la fonction cardiovasculaire globale.",
"side_effects": [
"Hypotension ou hypertension",
"Étourdissements posturaux",
"Modification de la fréquence cardiaque",
],
},
"Rénal": {
"label": "Risque rénal",
"description": "Effet potentiellement néphrotoxique ou altérant la clairance rénale.",
"side_effects": [
"Altération de la fonction rénale",
"Réduction de l'élimination des toxines",
"Risque accru chez l'insuffisant rénal",
],
},
"Hépatique": {
"label": "Risque hépatique",
"description": "Effet potentiel sur le métabolisme ou la fonction hépatique.",
"side_effects": [
"Élévation des enzymes hépatiques",
"Risque d'hépatotoxicité",
],
},
"Métabolique (CYP)": {
"label": "Interaction métabolique (CYP450)",
"description": "Inhibition ou induction probable d'une enzyme du cytochrome P450, modifiant les concentrations plasmatiques.",
"side_effects": [
"Augmentation ou diminution de la concentration plasmatique d'un des deux médicaments",
"Risque de surdosage relatif ou de sous-dosage",
"Effets toxiques liés à l'accumulation",
],
},
"Excrétion": {
"label": "Interaction sur l'excrétion",
"description": "Modification probable de la clairance ou de l'élimination d'un des deux médicaments.",
"side_effects": [
"Variation du taux sérique du médicament",
"Risque d'accumulation ou d'élimination accélérée",
],
},
"Hématologique": {
"label": "Risque hématologique",
"description": "Effet potentiel sur les cellules sanguines ou le transport de l'oxygène.",
"side_effects": [
"Risque de méthémoglobinémie",
"Anomalies de la numération sanguine",
],
},
"Sédation/SNC": {
"label": "Sédation / dépression du SNC",
"description": "Effet dépresseur additif sur le système nerveux central.",
"side_effects": [
"Somnolence accrue",
"Risque de dépression respiratoire (cas sévères)",
"Altération de la vigilance",
],
},
"Efficacité": {
"label": "Modification d'efficacité thérapeutique",
"description": "Risque de réduction ou de potentialisation de l'effet thérapeutique d'un des deux médicaments.",
"side_effects": [
"Perte d'efficacité thérapeutique",
"Effet thérapeutique exagéré",
],
},
"Pharmacovigilance (TWOSIDES)": {
"label": "Signal de pharmacovigilance",
"description": "Association identifiée par signal statistique de pharmacovigilance (base TWOSIDES), sans mécanisme unique établi.",
"side_effects": [
"Profil d'effets indésirables non spécifique",
"Surveillance clinique recommandée",
],
},
"Autre": {
"label": "Mécanisme non catégorisé",
"description": "Interaction détectée sans correspondance claire avec les catégories cliniques principales.",
"side_effects": [
"Mécanisme à investiguer au cas par cas",
],
},
}
# ──────────────────────────────────────────────────────────────────────────
# Recherche de médicament par nom
# ──────────────────────────────────────────────────────────────────────────
def normalize_name(name: str) -> str:
return re.sub(r"\s+", " ", name.strip().lower())
def find_drug(name: str) -> dict | None:
"""Cherche un médicament par nom (insensible à la casse). Retourne None si absent."""
key = normalize_name(name)
drug_id = NAME_TO_ID.get(key)
if drug_id is None:
alias = _COMMON_ALIASES.get(key)
if alias:
drug_id = NAME_TO_ID.get(alias)
key = alias
if drug_id is None:
return None
return {
"query": name,
"matched_name": key,
"drugbank_id": drug_id,
"smiles": ID_TO_SMILES[drug_id],
}
def suggest_names(name: str, n: int = 5) -> list[str]:
"""Suggestions de noms proches (utilisées uniquement dans le message d'erreur)."""
key = normalize_name(name)
matches = difflib.get_close_matches(key, _ALL_NAMES_SORTED, n=n, cutoff=0.6)
if not matches:
# repli : recherche par sous-chaîne
matches = [n2 for n2 in _ALL_NAMES_SORTED if key in n2][:n]
return matches
# ──────────────────────────────────────────────────────────────────────────
# Featurisation
# ──────────────────────────────────────────────────────────────────────────
@lru_cache(maxsize=4096)
def _get_fp(smiles: str) -> np.ndarray:
mol = Chem.MolFromSmiles(smiles)
if mol is None:
raise ValueError(f"SMILES invalide : {smiles}")
return _MFPGEN.GetFingerprintAsNumPy(mol).astype(np.int16)
def pair_features(smi_a: str, smi_b: str) -> np.ndarray:
fa, fb = _get_fp(smi_a), _get_fp(smi_b)
return np.concatenate([fa + fb, np.abs(fa - fb)]).reshape(1, -1)
def bit_to_fragment(smiles: str, bit_idx: int) -> str | None:
mol = Chem.MolFromSmiles(str(smiles))
if mol is None:
return None
bit_info: dict = {}
AllChem.GetMorganFingerprintAsBitVect(mol, FP_RADIUS, nBits=FP_BITS, bitInfo=bit_info)
if bit_idx not in bit_info:
return None
atom_idx, rad = bit_info[bit_idx][0]
env = Chem.FindAtomEnvironmentOfRadiusN(mol, rad, atom_idx)
submol = Chem.PathToSubmol(mol, env, atomMap={})
frag_smiles = Chem.MolToSmiles(submol)
return frag_smiles if frag_smiles else None
def molecule_svg(smiles: str, width: int = 280, height: int = 220) -> str:
"""Génère un SVG 2D de la molécule (pour affichage dans l'interface)."""
mol = Chem.MolFromSmiles(smiles)
if mol is None:
return ""
from rdkit.Chem.Draw import rdMolDraw2D
drawer = rdMolDraw2D.MolDraw2DSVG(width, height)
opts = drawer.drawOptions()
opts.clearBackground = False
rdMolDraw2D.PrepareAndDrawMolecule(drawer, mol)
drawer.FinishDrawing()
svg = drawer.GetDrawingText()
return svg
# ──────────────────────────────────────────────────────────────────────────
# Risque & confiance
# ──────────────────────────────────────────────────────────────────────────
def risk_level(proba: float) -> str:
if proba < 0.35:
return "faible"
elif proba < 0.70:
return "modéré"
return "élevé"
def confidence_score(proba: float) -> float:
"""Score de confiance du modèle : distance à la frontière de décision (0.5),
normalisée sur [0,1]. Une proba proche de 0 ou 1 -> confiance élevée ;
une proba proche de 0.5 -> confiance faible (zone d'incertitude)."""
return float(abs(proba - 0.5) * 2)
# ──────────────────────────────────────────────────────────────────────────
# Pipeline complet : prédiction + explicabilité
# ──────────────────────────────────────────────────────────────────────────
def analyze_pair(name_a: str, name_b: str, top_k_fragments: int = 6) -> dict:
drug_a = find_drug(name_a)
if drug_a is None:
return {"error": "not_found", "which": "A", "query": name_a,
"suggestions": suggest_names(name_a)}
drug_b = find_drug(name_b)
if drug_b is None:
return {"error": "not_found", "which": "B", "query": name_b,
"suggestions": suggest_names(name_b)}
smi_a, smi_b = drug_a["smiles"], drug_b["smiles"]
mol_a, mol_b = Chem.MolFromSmiles(smi_a), Chem.MolFromSmiles(smi_b)
if mol_a is None or mol_b is None:
return {"error": "invalid_smiles"}
feat = pair_features(smi_a, smi_b)
# --- Prédiction binaire ---
proba = float(BINARY_MODEL.predict_proba(feat)[0, 1])
level = risk_level(proba)
confidence = confidence_score(proba)
interaction_predicted = proba >= DECISION_THRESHOLD
# --- Prédiction multiclasse (type clinique) ---
multi_probs = MULTI_MODEL.predict_proba(feat)[0]
order = np.argsort(multi_probs)[::-1]
clinical_ranking = [
{"type": LABEL_ENCODER.classes_[i], "probability": float(multi_probs[i])}
for i in order
]
top_clinical_type = clinical_ranking[0]["type"]
clinical_info = CLINICAL_TYPE_INFO.get(top_clinical_type, {
"label": top_clinical_type, "description": "", "side_effects": []
})
# --- SHAP : fragments responsables ---
shap_values = _SHAP_EXPLAINER.shap_values(feat)[0]
top_bits = np.argsort(np.abs(shap_values))[::-1][:top_k_fragments]
fragments = []
for bit in top_bits:
is_sum = bit < FP_BITS
local_bit = int(bit % FP_BITS)
block = "Somme (A+B)" if is_sum else "Différence absolue |A−B|"
frag_a = bit_to_fragment(smi_a, local_bit)
frag_b = bit_to_fragment(smi_b, local_bit)
frag = frag_a or frag_b or None
source_drug = "A" if frag_a else ("B" if frag_b else None)
fragments.append({
"bit": local_bit,
"block": block,
"shap_value": float(shap_values[bit]),
"direction": "augmente le risque" if shap_values[bit] > 0 else "réduit le risque",
"fragment_smiles": frag,
"source_drug": source_drug,
})
# --- Raisons lisibles (synthèse pour le pharmacien) ---
reasons = build_reasons(fragments, clinical_info, proba)
return {
"error": None,
"drug_a": {"name": name_a.strip(), "drugbank_id": drug_a["drugbank_id"], "smiles": smi_a},
"drug_b": {"name": name_b.strip(), "drugbank_id": drug_b["drugbank_id"], "smiles": smi_b},
"prediction": {
"probability": proba,
"interaction_predicted": bool(interaction_predicted),
"risk_level": level,
"confidence": confidence,
"decision_threshold": DECISION_THRESHOLD,
},
"clinical_type": {
"top": top_clinical_type,
"label": clinical_info["label"],
"description": clinical_info["description"],
"ranking": clinical_ranking[:5],
},
"side_effects": clinical_info["side_effects"],
"fragments": fragments,
"reasons": reasons,
"model_metrics": {
"cold_start_auc": METRICS.get("cold_start", {}).get("auc"),
"warm_start_auc": METRICS.get("warm_start", {}).get("auc"),
},
}
def build_reasons(fragments: list[dict], clinical_info: dict, proba: float) -> list[dict]:
"""Construit 2-4 raisons lisibles à partir des fragments SHAP et du type clinique."""
reasons = []
positive_frags = [f for f in fragments if f["shap_value"] > 0 and f["fragment_smiles"]]
if positive_frags:
reasons.append({
"title": "Sous-structures moléculaires partagées",
"detail": (
f"{len(positive_frags)} fragment(s) commun(s) ou similaires identifiés par SHAP "
f"contribuent positivement au score de risque, suggérant une parenté structurelle "
f"avec des paires d'interaction connues."
),
})
if clinical_info.get("description"):
reasons.append({
"title": clinical_info["label"],
"detail": clinical_info["description"],
})
if proba >= 0.70:
reasons.append({
"title": "Score de probabilité élevé",
"detail": "La probabilité prédite dépasse largement le seuil de décision, indiquant une similarité forte avec des interactions documentées dans les données d'entraînement.",
})
elif proba < 0.35:
reasons.append({
"title": "Score de probabilité faible",
"detail": "Le modèle ne détecte pas de similarité significative avec des paires d'interaction connues.",
})
return reasons
def list_sample_names(n: int = 12) -> list[str]:
"""Quelques noms réels présents dans la base, pour peupler des exemples dans l'UI."""
preferred = [
"acetylsalicylic acid", "warfarin", "simvastatin", "clarithromycin",
"furosemide", "digoxin", "fluoxetine", "omeprazole",
]
found = [p for p in preferred if p in NAME_TO_ID]
return found[:n]