Spaces:
Sleeping
Sleeping
| import pandas as pd | |
| import torch | |
| from PIL import Image | |
| from health_multimodal.image.inference_engine import ImageInferenceEngine | |
| from health_multimodal.image.model.pretrained import get_biovil_t_image_encoder | |
| from health_multimodal.image.data.transforms import create_chest_xray_transform_for_inference | |
| from transformers import AutoModel, AutoTokenizer | |
| import streamlit as st | |
| from pathlib import Path | |
| import math | |
| import os | |
| import zipfile | |
| # 1. Configuration de la page | |
| st.set_page_config(page_title="MIROIR", layout="wide") | |
| st.markdown("### 🩺 MIROIR : Modèle d'Intelligence pour le Rapprochement d'Observations, d'Images et de Rapports") | |
| # 2. Mise en cache du chargement du modèle pour éviter de le recharger à chaque interaction | |
| def load_models(): | |
| device = "cuda" if torch.cuda.is_available() else "cpu" | |
| # Configuration des encodeurs d'images basés sur BioViL-T | |
| image_encoder = get_biovil_t_image_encoder() | |
| transform = create_chest_xray_transform_for_inference(resize=512, center_crop_size=448) | |
| image_inference_engine = ImageInferenceEngine(image_encoder, transform) | |
| # Configuration des encodeurs de textes | |
| tokenizer = AutoTokenizer.from_pretrained("microsoft/BiomedVLP-BioViL-T", trust_remote_code=True) | |
| text_model = AutoModel.from_pretrained("microsoft/BiomedVLP-BioViL-T", trust_remote_code=True).to(device) | |
| text_model.eval() | |
| return image_inference_engine, tokenizer, text_model, device | |
| # 3. Chargement modèle | |
| image_engine, tokenizer, text_model, device = load_models() | |
| # 4. Fonction pour calculer le score de similarité cosinus | |
| def compute_similarity(image_path, text_content): | |
| local_device = "cuda" if torch.cuda.is_available() else "cpu" | |
| try: | |
| path_object = Path(image_path) | |
| with torch.no_grad(): | |
| image_embedding = image_engine.get_projected_global_embedding(path_object) | |
| if not isinstance(image_embedding, torch.Tensor): | |
| image_embedding = torch.tensor(image_embedding).to(local_device) | |
| else: | |
| image_embedding = image_embedding.to(local_device) | |
| if image_embedding.ndim == 1: | |
| image_embedding = image_embedding.unsqueeze(0) | |
| image_embedding = image_embedding / image_embedding.norm(dim=-1, keepdim=True) | |
| inputs = tokenizer(text_content, return_tensors="pt", padding="max_length", truncation=True, max_length=512).to(local_device) | |
| with torch.no_grad(): | |
| text_embedding = text_model.get_projected_text_embeddings( | |
| input_ids=inputs["input_ids"], | |
| attention_mask=inputs["attention_mask"] | |
| ) | |
| if text_embedding.ndim == 1: | |
| text_embedding = text_embedding.unsqueeze(0) | |
| text_embedding = text_embedding / text_embedding.norm(dim=-1, keepdim=True) | |
| similarity = torch.mm(image_embedding, text_embedding.t()).item() | |
| similarity_prob = 1 / (1 + math.exp(-similarity * 4)) # Le multiplicateur (ex: 4) ajuste la sensibilité | |
| return similarity_prob | |
| except Exception as e: | |
| st.error(f"Erreur lors du calcul : {e}") | |
| return 0.0 | |
| if "current_pair" not in st.session_state: | |
| st.session_state.current_pair = None | |
| st.session_state.pair_type = None | |
| col_btn1, col_btn2 = st.columns(2) | |
| with col_btn1: | |
| if st.button("✅ Exemple une paire CORRECTE (Match)", use_container_width=True): | |
| df_pos = pd.read_csv("./chexpert_matches_sample.csv", index_col=0, nrows=1000) | |
| st.session_state.current_pair = df_pos.sample(1).iloc[0] | |
| st.session_state.pair_type = "Correcte (Match)" | |
| with col_btn2: | |
| if st.button("❌ Exemple d'une paire INCORRECTE (Mismatch)", use_container_width=True): | |
| df_neg = pd.read_csv("./chexpert_mismatches_swapping_2.csv", index_col=0, nrows=1000) | |
| st.session_state.current_pair = df_neg.sample(1).iloc[0] | |
| st.session_state.pair_type = "Incorrecte (Mismatch)" | |
| if st.session_state.current_pair is not None: | |
| pair = st.session_state.current_pair | |
| st.info(f"Source actuelle : Paire **{st.session_state.pair_type}**") | |
| col1, col2 = st.columns([0.6, 1.4]) | |
| with col1: | |
| st.markdown("#### 🖼️ Radiographie Thoracique :") | |
| img_source = st.radio("Source de l'image :", ["Image originale de la paire", "Uploader une autre image"]) | |
| image_to_process = None | |
| # 1. CAS DE L'IMAGE ORIGINALE DE LA PAIRE | |
| if img_source == "Image originale de la paire": | |
| image_path = pair["path_to_image"] | |
| st.caption(f"Chemin recherché : `{image_path}`") | |
| try: | |
| with zipfile.ZipFile("mini_chexpert.zip", 'r') as z: | |
| liste_fichiers = z.namelist() | |
| chemin_final = None | |
| if image_path in liste_fichiers: | |
| chemin_final = image_path | |
| else: | |
| alternative_path = image_path.replace("CheXpert-v1.0-small/", "") | |
| if alternative_path in liste_fichiers: | |
| chemin_final = alternative_path | |
| if chemin_final is not None: | |
| with z.open(chemin_final) as fichier_image: | |
| img_display = Image.open(fichier_image).convert("RGB") | |
| st.image(img_display, width=250) | |
| image_to_process = "temp_image_paire.jpg" | |
| img_display.save(image_to_process) | |
| else: | |
| st.error(f"⚠️ L'image {image_path} n'est pas dans le ZIP réduit.") | |
| except Exception as e: | |
| st.error("Erreur technique lors de la lecture de l'image.") | |
| print(f"[MIROIR] Erreur chemin ZIP : {e}") | |
| # 2. CAS DU PARCOURS DE L'ARCHIVE ZIP | |
| elif img_source == "Uploader une autre image": | |
| try: | |
| with zipfile.ZipFile("mini_chexpert.zip", 'r') as z: | |
| toutes_les_images = [f for f in z.namelist() if f.lower().endswith(('.jpg', '.jpeg', '.png'))] | |
| if toutes_les_images: | |
| image_choisie = st.selectbox( | |
| "📁 Choisissez une radiographie directement dans l'archive .zip :", | |
| toutes_les_images | |
| ) | |
| with zipfile.ZipFile("mini_chexpert.zip", 'r') as z: | |
| with z.open(image_choisie) as fichier_image: | |
| img_display = Image.open(fichier_image).convert("RGB") | |
| st.image(img_display, width=250) | |
| image_to_process = "temp_image_parcourue.jpg" | |
| img_display.save(image_to_process) | |
| else: | |
| st.warning("Aucune image trouvée à l'intérieur du fichier ZIP.") | |
| except Exception as e: | |
| st.error("Impossible de parcourir le fichier ZIP.") | |
| print(f"[MIROIR] Erreur parcours ZIP : {e}") | |
| with col2: | |
| st.markdown("#### 📝 Compte Rendu (Éditables en direct) :") | |
| default_text = pair["section_impression"] if pd.notna(pair["section_impression"]) else pair["report"] | |
| edited_text = st.text_area( | |
| "Modifiez le texte ci-dessous pour voir le score changer instantanément :", | |
| value=str(default_text), | |
| height=300 | |
| ) | |
| if image_to_process is not None and edited_text: | |
| with st.spinner("Calcul de la similarité..."): | |
| score = compute_similarity(image_to_process, edited_text) | |
| st.markdown("#### 📊 Score en direct :") | |
| percentage = max(0.0, min(1.0, score)) * 100 | |
| custom_html = f""" | |
| <div style="width: 100%; margin-top: 40px; margin-bottom: 10px; position: relative; font-family: sans-serif;"> | |
| <div style="width: 100%; height: 16px; background-color: #e0e4ec; border-radius: 8px; position: relative;"> | |
| <div style="width: {percentage}%; height: 100%; background: linear-gradient(90deg, #ff4b4b, #1f77b4); border-radius: 8px;"></div> | |
| <div style="position: absolute; left: {percentage}%; top: -35px; transform: translateX(-50%); white-space: nowrap;"> | |
| <span style="font-size: 20px; font-weight: bold; color: #1f77b4; background-color: #ffffff; padding: 4px 10px; border-radius: 6px; border: 2px solid #1f77b4; box-shadow: 0px 2px 6px rgba(0,0,0,0.15);"> | |
| {score:.2f} | |
| </span> | |
| </div> | |
| </div> | |
| </div> | |
| """ | |
| st.html(custom_html) | |
| st.write("") # Respiration visuelle | |
| # 3. Affichage immédiat de la conclusion juste en dessous | |
| if score >= 0.75: | |
| st.success(f"🟢 **Excellente cohérence !** La radiographie et le compte rendu médical décrivent parfaitement les mêmes observations.") | |
| elif 0.40 <= score < 0.75: | |
| st.warning(f"⚠️ **Attention : Cohérence modérée.** Le texte a été modifié ou certaines observations du compte rendu ne semblent pas correspondre de manière flagrante aux structures identifiées sur l'image.") | |
| else: | |
| st.error(f"🚨 **ALERTE : Incohérence majeure !** Il y a une contradiction totale ou une absence de lien logique entre cette radiographie et le texte fourni. Risque d'erreur de diagnostic ou d'inversion de dossier patient.") | |
| else: | |
| st.info("Veuillez fournir une image et un texte valides pour calculer le score.") | |
| else: | |
| st.write("👈 Cliquez sur l'un des boutons ci-dessus pour charger une paire aléatoire de données.") |