{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# ARVI-RX : Assistant Radiologique Virtuel Intelligent\n\nCe notebook implémente le pipeline principal du projet sur Kaggle.\nIl contient les étapes suivantes :\n1. **Préparation et normalisation** des images DICOM.\n2. **Définition du prompt structuré** et du schéma JSON.\n3. **Chargement d'un VLM** (PaliGemma/MedGemma) et exécution.\n4. **Évaluation** sur un sous-ensemble.","metadata":{}},{"cell_type":"code","source":"!pip install -q pydicom\n!pip install -q transformers accelerate bitsandbytes Pillow opencv-python pandas scikit-learn","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-25T14:02:37.812243Z","iopub.execute_input":"2026-06-25T14:02:37.813249Z","iopub.status.idle":"2026-06-25T14:02:44.637709Z","shell.execute_reply.started":"2026-06-25T14:02:37.813212Z","shell.execute_reply":"2026-06-25T14:02:44.636882Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport pydicom\nimport cv2\nimport json\nimport time\nimport torch\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nfrom transformers import AutoProcessor, AutoModelForCausalLM\nfrom sklearn.metrics import accuracy_score, f1_score, confusion_matrix\nimport warnings\nwarnings.filterwarnings('ignore')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-25T14:02:44.639336Z","iopub.execute_input":"2026-06-25T14:02:44.639567Z","iopub.status.idle":"2026-06-25T14:02:44.645244Z","shell.execute_reply.started":"2026-06-25T14:02:44.639543Z","shell.execute_reply":"2026-06-25T14:02:44.644609Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 1. Définition des chemins et chargement des métadonnées","metadata":{}},{"cell_type":"code","source":"# Configuration des chemins. Sur Kaggle, les datasets liés sont sous /kaggle/input/\n# Si vous utilisez ce notebook en local, assurez-vous que les dossiers stage_2_* sont dans le dossier courant.\nif os.path.exists('/kaggle/input/competitions/rsna-pneumonia-detection-challenge'):\n    BASE_DIR = '/kaggle/input/competitions/rsna-pneumonia-detection-challenge/'\nelse:\n    BASE_DIR = '.'  # Chemin local\n\nTRAIN_IMAGES_DIR = os.path.join(BASE_DIR, 'stage_2_train_images')\nTRAIN_LABELS_CSV = os.path.join(BASE_DIR, 'stage_2_train_labels.csv')\nCLASS_INFO_CSV = os.path.join(BASE_DIR, 'stage_2_detailed_class_info.csv')\n\ndf_labels = pd.read_csv(TRAIN_LABELS_CSV)\ndf_class_info = pd.read_csv(CLASS_INFO_CSV)\n\n# Jointure des deux fichiers CSV\ndf_merged = pd.merge(df_labels, df_class_info, on='patientId', how='left')\n\n# Un patient peut avoir plusieurs boîtes englobantes, on déduplique pour la classification image entière\ndf_unique = df_merged.drop_duplicates(subset=['patientId']).copy()\nprint(f\"Nombre total de patients uniques : {len(df_unique)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-25T14:02:44.646076Z","iopub.execute_input":"2026-06-25T14:02:44.646411Z","iopub.status.idle":"2026-06-25T14:02:44.738114Z","shell.execute_reply.started":"2026-06-25T14:02:44.646392Z","shell.execute_reply":"2026-06-25T14:02:44.737473Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 2. Sélection d'un sous-ensemble (150 cas)\nAfin de conserver des temps d'inférence raisonnables dans le prototype, nous tirons au hasard 150 cas.","metadata":{}},{"cell_type":"code","source":"# On sélectionne 75 cas 'Normal' et 75 cas 'Lung Opacity'\ndf_normal = df_unique[df_unique['class'] == 'Normal'].sample(75, random_state=42)\ndf_opacity = df_unique[df_unique['class'] == 'Lung Opacity'].sample(75, random_state=42)\n\ndf_subset = pd.concat([df_normal, df_opacity]).sample(frac=1, random_state=42).reset_index(drop=True)\n\n# Mapping vers les classes de notre projet : normal | suspected_opacity | uncertain\ndef map_class(cls_name):\n    if cls_name == 'Normal':\n        return 'normal'\n    return 'suspected_opacity'\n\ndf_subset['project_class'] = df_subset['class'].apply(map_class)\nprint(\"Répartition dans le sous-ensemble :\")\nprint(df_subset['project_class'].value_counts())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-25T14:02:44.740050Z","iopub.execute_input":"2026-06-25T14:02:44.740355Z","iopub.status.idle":"2026-06-25T14:02:44.758718Z","shell.execute_reply.started":"2026-06-25T14:02:44.740335Z","shell.execute_reply":"2026-06-25T14:02:44.758067Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3. Prétraitement des images DICOM\nLes VLMs attendent généralement une image encodée en 8-bits, en 3 canaux RGB.","metadata":{}},{"cell_type":"code","source":"def process_dicom(patient_id):\n    \"\"\"\n    Lit un fichier DICOM, normalise les pixels, applique un CLAHE pour améliorer\n    les contrastes et retourne une image PIL RGB.\n    \"\"\"\n    path = os.path.join(TRAIN_IMAGES_DIR, f\"{patient_id}.dcm\")\n    if not os.path.exists(path):\n        return None\n    \n    dcm = pydicom.dcmread(path)\n    img = dcm.pixel_array\n    \n    # Normalisation Min-Max\n    img = (img - np.min(img)) / (np.max(img) - np.min(img)) * 255.0\n    img = img.astype(np.uint8)\n    \n    # Amélioration du contraste (CLAHE)\n    clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8,8))\n    img = clahe.apply(img)\n    \n    # Conversion en RGB car les VLM attendent souvent 3 canaux\n    img_rgb = cv2.cvtColor(img, cv2.COLOR_GRAY2RGB)\n    return Image.fromarray(img_rgb)\n\n# Test visuel sur une image du sous-ensemble\nsample_patient = df_subset.iloc[0]['patientId']\nsample_img = process_dicom(sample_patient)\n\nif sample_img:\n    plt.figure(figsize=(5,5))\n    plt.imshow(sample_img)\n    plt.title(f\"Classe: {df_subset.iloc[0]['project_class']}\")\n    plt.axis('off')\n    plt.show()\nelse:\n    print(f\"Image DICOM introuvable pour le patient {sample_patient}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-25T14:02:44.759637Z","iopub.execute_input":"2026-06-25T14:02:44.759918Z","iopub.status.idle":"2026-06-25T14:02:44.934037Z","shell.execute_reply.started":"2026-06-25T14:02:44.759899Z","shell.execute_reply":"2026-06-25T14:02:44.933239Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 4. Définition du Prompt Structuré (Contrat JSON)\nCe prompt force le modèle à adopter un schéma strict en JSON, tout en lui donnant les règles métier.","metadata":{}},{"cell_type":"code","source":"BASELINE_PROMPT = \"\"\"You are an AI assistant designed for educational purposes only. You must analyze this frontal chest X-ray.\nYou will be given x-ray and you will produce a JSON output strictly matching this schema:\n{\n  \"image_quality\": \"good | limited | poor\",\n  \"predicted_class\": \"normal | suspected_opacity | uncertain\",\n  \"confidence\": <float between 0.0 and 1.0>,\n  \"visual_evidence\": [\"observation 1\", \"observation 2\"],\n  \"justification\": \"short justification\",\n  \"limitations\": [\"limitation 1\"],\n  \"warning\": \"Prototype pédagogique. Non destiné au diagnostic médical.\"\n}\n\nRules:\n1. \"normal\" means no suspected opacity.\n2. \"suspected_opacity\" means you see a potential lung opacity.\n3. If you really cannot choose between \"normal\" or \"suspected_opacity\" classify as \"uncertain\"\n3. Output ONLY valid JSON, nothing else.\n\"\"\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-25T14:02:44.935068Z","iopub.execute_input":"2026-06-25T14:02:44.935366Z","iopub.status.idle":"2026-06-25T14:02:44.939589Z","shell.execute_reply.started":"2026-06-25T14:02:44.935336Z","shell.execute_reply":"2026-06-25T14:02:44.938928Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 5. Chargement du Modèle VLM\nNous utilisons `google/medgemma-4b-it`, assurez-vous d'avoir accepté les conditions sur Hugging Face et d'avoir inséré votre token hugging face dans les secret kaggle du notebook sous le nom HF_TOKEN.","metadata":{}},{"cell_type":"code","source":"from kaggle_secrets import UserSecretsClient\nuser_secrets = UserSecretsClient()\nHF_TOKEN = user_secrets.get_secret(\"HF_TOKEN\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-25T14:02:44.940543Z","iopub.execute_input":"2026-06-25T14:02:44.940726Z","iopub.status.idle":"2026-06-25T14:02:45.188499Z","shell.execute_reply.started":"2026-06-25T14:02:44.940706Z","shell.execute_reply":"2026-06-25T14:02:45.187560Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from huggingface_hub import login\nlogin(HF_TOKEN)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-25T14:02:45.189759Z","iopub.execute_input":"2026-06-25T14:02:45.190079Z","iopub.status.idle":"2026-06-25T14:02:45.329558Z","shell.execute_reply.started":"2026-06-25T14:02:45.190050Z","shell.execute_reply":"2026-06-25T14:02:45.328976Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from transformers import BitsAndBytesConfig\n\nMODEL_NAME = \"google/medgemma-4b-it\"\n\nprint(f\"Chargement du processeur et du modèle {MODEL_NAME}...\")\nprocessor = AutoProcessor.from_pretrained(MODEL_NAME)\n\n# L'astuce est ici : on force le calcul en bfloat16 pour éviter les NaNs de Gemma\nquantization_config = BitsAndBytesConfig(\n    load_in_4bit=True,\n    bnb_4bit_compute_dtype=torch.bfloat16 \n)\n\nmodel = AutoModelForCausalLM.from_pretrained(\n    MODEL_NAME, \n    device_map=\"auto\",\n    torch_dtype=torch.bfloat16,\n    quantization_config=quantization_config\n)\nprint(\"Modèle chargé avec succès.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-25T14:02:52.048902Z","iopub.execute_input":"2026-06-25T14:02:52.049515Z","iopub.status.idle":"2026-06-25T14:03:10.012414Z","shell.execute_reply.started":"2026-06-25T14:02:52.049487Z","shell.execute_reply":"2026-06-25T14:03:10.011460Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 6. Exécution de l'Inférence (Baseline) et Extraction JSON","metadata":{}},{"cell_type":"code","source":"results = []\nsubset_eval = df_subset.head(20)\n\nprint(\"Début de l'inférence...\")\nfor idx, row in subset_eval.iterrows():\n    patient_id = row['patientId']\n    img = process_dicom(patient_id)\n    if img is None:\n        continue\n        \n    messages = [\n        {\n            \"role\": \"user\",\n            \"content\": [\n                {\"type\": \"image\"},\n                {\"type\": \"text\", \"text\": f\"{BASELINE_PROMPT}\\n\\nAnalyze the following image.\"}\n            ]\n        }\n    ]\n    \n    prompt_formate = processor.apply_chat_template(messages, add_generation_prompt=True, tokenize=False)\n    \n    # Un simple .to(model.device) suffit maintenant !\n    inputs = processor(text=prompt_formate, images=img, return_tensors=\"pt\").to(model.device)\n    \n    start_time = time.time()\n    with torch.no_grad():\n        outputs = model.generate(\n            **inputs, \n            max_new_tokens=512\n        )\n    latency = time.time() - start_time\n    \n    input_len = inputs[\"input_ids\"].shape[1]\n    generated_tokens = outputs[0][input_len:] \n    response_text = processor.decode(generated_tokens, skip_special_tokens=True).strip()\n    \n    print(f\"\\n--- REPONSE DU MODELE (Image {idx+1}) ---\")\n    print(response_text)\n    print(\"---------------------------------------\\n\")\n    \n\n    # Parsing JSON\n    try:\n        start_idx = response_text.find(\"{\")\n        end_idx = response_text.rfind(\"}\") + 1\n        if start_idx != -1 and end_idx != 0:\n            json_str = response_text[start_idx:end_idx]\n            pred_json = json.loads(json_str)\n        else:\n            raise ValueError(\"Aucun JSON détecté\")\n    except Exception as e:\n        pred_json = {\n            \"predicted_class\": \"uncertain\",\n            \"confidence\": 0.0,\n            \"error\": f\"Erreur parsing: {str(e)}\"\n        }\n    \n    if pred_json.get(\"confidence\", 0.0) < 0.6:\n        pred_json[\"predicted_class\"] = \"uncertain\"\n        \n    results.append({\n        \"patientId\": patient_id,\n        \"true_class\": row['project_class'],\n        \"predicted_class\": pred_json.get(\"predicted_class\", \"uncertain\"),\n        \"confidence\": pred_json.get(\"confidence\", 0.0),\n        \"latency\": latency,\n        \"full_json\": json.dumps(pred_json)\n    })\n    \n    print(f\"Image {idx+1}/{len(subset_eval)} | Vrai: {row['project_class']:<17} | Prédit: {pred_json.get('predicted_class')}\")\n\ndf_results = pd.DataFrame(results)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-25T14:05:42.105820Z","iopub.execute_input":"2026-06-25T14:05:42.106647Z","iopub.status.idle":"2026-06-25T14:13:57.354877Z","shell.execute_reply.started":"2026-06-25T14:05:42.106616Z","shell.execute_reply":"2026-06-25T14:13:57.354102Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 7. Évaluation et Métriques\nCalcul des métriques (Accuracy, Macro-F1) en tenant compte de la classe \"uncertain\".","metadata":{}},{"cell_type":"code","source":"print(\"\\n--- RÉSULTATS GLOBAUX ---\")\n\nvalid_preds = df_results[df_results['predicted_class'].isin(['normal', 'suspected_opacity'])]\nuncertain_rate = (len(df_results) - len(valid_preds)) / len(df_results)\n\nprint(f\"Taux d'incertitude (rejets) : {uncertain_rate:.1%}\")\n\nif len(valid_preds) > 0:\n    acc = accuracy_score(valid_preds['true_class'], valid_preds['predicted_class'])\n    f1 = f1_score(valid_preds['true_class'], valid_preds['predicted_class'], average='macro')\n    print(f\"Accuracy (sur prédictions fermes) : {acc:.3f}\")\n    print(f\"Macro F1 (sur prédictions fermes) : {f1:.3f}\")\nelse:\n    print(\"Toutes les prédictions ont été classées comme incertaines.\")\n\nprint(\"\\nMatrice de confusion complète :\")\nlabels = ['normal', 'suspected_opacity', 'uncertain']\ncm = confusion_matrix(df_results['true_class'], df_results['predicted_class'], labels=labels)\ncm_df = pd.DataFrame(cm, index=[f\"True_{l}\" for l in labels], columns=[f\"Pred_{l}\" for l in labels])\ndisplay(cm_df)\n\n# Sauvegarde pour analyse d'erreurs ultérieure\ndf_results.to_csv(\"baseline_results.csv\", index=False)\nprint(\"\\nRésultats sauvegardés dans baseline_results.csv\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-25T14:21:44.365740Z","iopub.execute_input":"2026-06-25T14:21:44.366415Z","iopub.status.idle":"2026-06-25T14:21:44.407301Z","shell.execute_reply.started":"2026-06-25T14:21:44.366387Z","shell.execute_reply":"2026-06-25T14:21:44.406567Z"}},"outputs":[],"execution_count":null}]}