{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":39272,"databundleVersionId":4629629,"isSourceIdPinned":false,"sourceType":"competition"}],"dockerImageVersionId":31193,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# =============================================================================\n# Section 1 - Imports, configuration, chargement du CSV et sous-ensemble équilibré\n# =============================================================================\nimport os\nfrom pathlib import Path\nimport warnings\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\nimport pydicom\nfrom PIL import Image\nfrom tqdm.auto import tqdm\n\nimport torch\nfrom torch import nn\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as T\nimport torchvision.models as models\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import (\n    accuracy_score,\n    precision_recall_fscore_support,\n    roc_auc_score,\n    roc_curve,\n    confusion_matrix,\n    classification_report,\n)\n\nwarnings.filterwarnings(\"ignore\", category=RuntimeWarning)\nwarnings.filterwarnings(\"ignore\", category=FutureWarning)\n\nsns.set(style=\"whitegrid\")\nplt.rcParams[\"figure.figsize\"] = (8, 5)\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Device utilisé :\", device)\n\n# Dossier RSNA brut\nDATA_DIR = Path(\"/kaggle/input/rsna-breast-cancer-detection\")\nTRAIN_CSV = DATA_DIR / \"train.csv\"\nTRAIN_IMG_DIR = DATA_DIR / \"train_images\"\n\nprint(\"TRAIN_CSV existe ?\", TRAIN_CSV.exists())\nprint(\"TRAIN_IMG_DIR existe ?\", TRAIN_IMG_DIR.exists())\n\n# Chargement du CSV\ntrain = pd.read_csv(TRAIN_CSV)\nprint(\"Taille du DataFrame train :\", train.shape)\ndisplay(train.head())\n\nprint(\"Répartition globale du label 'cancer' :\")\nprint(train[\"cancer\"].value_counts())\n\n# Sous-ensemble équilibré : toutes les positives + même nb de négatives\npos_df = train[train[\"cancer\"] == 1].copy()\nneg_df = train[train[\"cancer\"] == 0].sample(len(pos_df), random_state=42)\n\nsubset_df = pd.concat([pos_df, neg_df], ignore_index=True)\nsubset_df = subset_df.sample(frac=1.0, random_state=42).reset_index(drop=True)\n\nprint(\"\\nTaille subset_df :\", subset_df.shape)\nprint(\"Répartition des labels dans subset_df :\")\nprint(subset_df[\"cancer\"].value_counts())\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-12-08T20:02:22.19136Z","iopub.execute_input":"2025-12-08T20:02:22.191594Z","iopub.status.idle":"2025-12-08T20:02:32.730817Z","shell.execute_reply.started":"2025-12-08T20:02:22.191569Z","shell.execute_reply":"2025-12-08T20:02:32.730208Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# Section 2 - Fonctions de pré-traitement (DICOM -> array -> 512×512)\n# =============================================================================\n\ndef load_dicom_raw(patient_id, image_id, verbose=False):\n    \"\"\"\n    Lit un DICOM à partir de patient_id et image_id.\n    Retourne (image numpy float32, objet dcm) ou (None, None) en cas d'échec.\n    Gère les problèmes de compression via try/except.\n    \"\"\"\n    patient_id = int(patient_id)\n    image_id = int(image_id)\n    dcm_path = TRAIN_IMG_DIR / str(patient_id) / f\"{image_id}.dcm\"\n\n    try:\n        dcm = pydicom.dcmread(dcm_path)\n        img = dcm.pixel_array.astype(np.float32)\n\n        # Application éventuelle de RescaleSlope / RescaleIntercept\n        intercept = getattr(dcm, \"RescaleIntercept\", 0.0)\n        slope = getattr(dcm, \"RescaleSlope\", 1.0)\n        img = img * slope + intercept\n\n        return img, dcm\n    except Exception as e:\n        if verbose:\n            print(f\"[load_dicom_raw] Erreur lecture {dcm_path} : {e}\")\n        return None, None\n\n\ndef window_image(img, lower_percentile=5.0, upper_percentile=99.5):\n    \"\"\"\n    Améliore le contraste via un windowing basé sur des percentiles.\n    Normalise ensuite en [0,1].\n    \"\"\"\n    if img is None:\n        return None\n\n    lo = np.percentile(img, lower_percentile)\n    hi = np.percentile(img, upper_percentile)\n    if lo >= hi:\n        return None\n\n    img = np.clip(img, lo, hi)\n    img = img - img.min()\n    if img.max() > 0:\n        img = img / img.max()\n    return img\n\n\ndef crop_to_breast(img, threshold_ratio=0.02):\n    \"\"\"\n    Recadre l'image autour de la zone du sein en détectant les pixels \"non noirs\".\n    threshold_ratio : fraction du max pour définir le seuil.\n    \"\"\"\n    if img is None:\n        return None\n\n    thr = threshold_ratio * img.max()\n    mask = img > thr\n    if not mask.any():\n        # rien à croper, on retourne l'image telle quelle\n        return img\n\n    ys, xs = np.where(mask)\n    y_min, y_max = ys.min(), ys.max()\n    x_min, x_max = xs.min(), xs.max()\n    return img[y_min:y_max + 1, x_min:x_max + 1]\n\n\ndef make_square(img):\n    \"\"\"\n    Rend l'image carrée en recadrant au centre avec la plus petite dimension.\n    \"\"\"\n    if img is None:\n        return None\n    h, w = img.shape\n    side = min(h, w)\n    y_start = (h - side) // 2\n    x_start = (w - side) // 2\n    return img[y_start:y_start + side, x_start:x_start + side]\n\n\ndef to_pil_512(img, size=512):\n    \"\"\"\n    Convertit une image numpy [0,1] en image PIL 8 bits et la redimensionne en size x size.\n    \"\"\"\n    if img is None:\n        return None\n    img = (img * 255.0).clip(0, 255).astype(np.uint8)\n    pil_img = Image.fromarray(img)\n    pil_img = pil_img.resize((size, size), resample=Image.BILINEAR)\n    return pil_img\n\n\ndef process_row_to_pil(row, verbose=False):\n    \"\"\"\n    Pipeline complet :\n    DICOM -> windowing -> crop -> carré -> resize 512 -> PIL.Image\n    Retourne (pil_img, dcm) ou (None, None).\n    \"\"\"\n    img, dcm = load_dicom_raw(row[\"patient_id\"], row[\"image_id\"], verbose=verbose)\n    if img is None:\n        return None, None\n\n    img = window_image(img)\n    if img is None:\n        return None, dcm\n\n    img = crop_to_breast(img)\n    img = make_square(img)\n    pil_img = to_pil_512(img)\n\n    if pil_img is None:\n        return None, dcm\n\n    return pil_img, dcm\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-08T20:02:32.732585Z","iopub.execute_input":"2025-12-08T20:02:32.732848Z","iopub.status.idle":"2025-12-08T20:02:32.743869Z","shell.execute_reply.started":"2025-12-08T20:02:32.732832Z","shell.execute_reply":"2025-12-08T20:02:32.743222Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# Section 3 - Visualisation AVANT / APRÈS (un cas positif, un cas négatif)\n# =============================================================================\n\ndef get_decodable_example(df, label, max_tries=80):\n    \"\"\"\n    Cherche une ligne dont l'image est décodable et pré-traitable pour le label donné (0 ou 1).\n    \"\"\"\n    subset = df[df[\"cancer\"] == label]\n    if len(subset) == 0:\n        raise RuntimeError(f\"Aucune image pour label={label} dans le DataFrame.\")\n\n    subset = subset.sample(min(max_tries, len(subset)), random_state=0)\n\n    for _, row in subset.iterrows():\n        img, dcm = load_dicom_raw(row[\"patient_id\"], row[\"image_id\"], verbose=False)\n        if img is None:\n            continue\n        processed, _ = process_row_to_pil(row)\n        if processed is not None:\n            return row\n\n    raise RuntimeError(f\"Aucune image décodable trouvée pour label={label} après {max_tries} essais.\")\n\n\ndef show_before_after(row, title_prefix=\"\"):\n    \"\"\"\n    Affiche :\n    - l'image DICOM windowée,\n    - l'image finale 512×512.\n    \"\"\"\n    raw_img, dcm = load_dicom_raw(row[\"patient_id\"], row[\"image_id\"], verbose=False)\n    if raw_img is None:\n        print(\"Impossible de lire l'image brute.\")\n        return\n\n    windowed = window_image(raw_img)\n    processed, _ = process_row_to_pil(row)\n\n    if windowed is None or processed is None:\n        print(\"Pré-traitement impossible pour cet exemple.\")\n        return\n\n    fig, axes = plt.subplots(1, 2, figsize=(10, 5))\n    axes[0].imshow(windowed, cmap=\"gray\")\n    axes[0].set_title(f\"{title_prefix}DICOM après windowing\")\n    axes[0].axis(\"off\")\n\n    axes[1].imshow(processed, cmap=\"gray\")\n    axes[1].set_title(\"Image finale 512×512\")\n    axes[1].axis(\"off\")\n\n    plt.tight_layout()\n    plt.show()\n\n    print(\"Taille brute :\", raw_img.shape)\n    print(\"Taille finale (512×512) :\", processed.size)\n\n\nprint(\"Exemple POSITIF (cancer=1)\")\npos_example = get_decodable_example(subset_df, label=1)\nshow_before_after(pos_example, title_prefix=\"[POS] \")\n\nprint(\"Exemple NÉGATIF (cancer=0)\")\nneg_example = get_decodable_example(subset_df, label=0)\nshow_before_after(neg_example, title_prefix=\"[NEG] \")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-08T20:02:32.744516Z","iopub.execute_input":"2025-12-08T20:02:32.744733Z","iopub.status.idle":"2025-12-08T20:02:47.988572Z","shell.execute_reply.started":"2025-12-08T20:02:32.744708Z","shell.execute_reply":"2025-12-08T20:02:47.987886Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# Section 4 - Génération des PNG 512×512 + CSV de métadonnées\n# =============================================================================\n\nOUTPUT_DIR = Path(\"/kaggle/working/rsna-processed-512\")\nOUTPUT_IMG_DIR = OUTPUT_DIR / \"images\"\nOUTPUT_IMG_DIR.mkdir(parents=True, exist_ok=True)\n\nmetadata = []\nn_ok = 0\nn_fail = 0\n\nfor idx, row in tqdm(subset_df.iterrows(), total=len(subset_df)):\n    pil_img, dcm = process_row_to_pil(row)\n    if pil_img is None:\n        n_fail += 1\n        continue\n\n    label = int(row[\"cancer\"])\n    patient_id = int(row[\"patient_id\"])\n    image_id = int(row[\"image_id\"])\n\n    # Dossier : images/train/{label}/\n    out_folder = OUTPUT_IMG_DIR / \"train\" / str(label)\n    out_folder.mkdir(parents=True, exist_ok=True)\n\n    out_name = f\"{patient_id}_{image_id}.png\"\n    out_path = out_folder / out_name\n\n    pil_img.save(out_path)\n\n    metadata.append(\n        {\n            \"patient_id\": patient_id,\n            \"image_id\": image_id,\n            \"cancer\": label,\n            \"filepath\": str(out_path.relative_to(OUTPUT_DIR)),\n        }\n    )\n    n_ok += 1\n\nprint(\"Images traitées avec succès :\", n_ok)\nprint(\"Images ignorées (erreurs DICOM ou pré-traitement) :\", n_fail)\n\nmeta_df = pd.DataFrame(metadata)\nmeta_csv_path = OUTPUT_DIR / \"processed_metadata.csv\"\nmeta_df.to_csv(meta_csv_path, index=False)\n\nprint(\"CSV de métadonnées sauvegardé sous :\", meta_csv_path)\nprint(\"Fichier existe ? \", meta_csv_path.exists())\n\nprint(\"\\nContenu de /kaggle/working :\")\nprint(os.listdir(\"/kaggle/working\"))\n\nprint(\"\\nContenu de /kaggle/working/rsna-processed-512 :\")\nprint(os.listdir(OUTPUT_DIR))\n\ndisplay(meta_df.head())\nprint(\"Répartition des labels dans meta_df :\")\nprint(meta_df[\"cancer\"].value_counts())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-08T20:02:47.98929Z","iopub.execute_input":"2025-12-08T20:02:47.989549Z","iopub.status.idle":"2025-12-08T20:53:36.798636Z","shell.execute_reply.started":"2025-12-08T20:02:47.989525Z","shell.execute_reply":"2025-12-08T20:53:36.798023Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# Section 5 - Split train / val / test (stratifié par patiente)\n# =============================================================================\n\n# Table \"par patiente\" : si au moins une image est positive, la patiente est notée 1\npatient_df = (\n    meta_df.groupby(\"patient_id\")[\"cancer\"]\n    .max()\n    .reset_index()\n)\n\nprint(\"Nombre de patientes :\", len(patient_df))\nprint(\"Répartition par patiente :\")\nprint(patient_df[\"cancer\"].value_counts())\n\n# 1) Split train vs (val+test) -> 70% / 30%\ntrain_patients, temp_patients = train_test_split(\n    patient_df[\"patient_id\"],\n    test_size=0.3,\n    random_state=42,\n    stratify=patient_df[\"cancer\"],\n)\n\n# 2) Split (val+test) en val et test (50/50) -> 15% / 15%\ntemp_df = patient_df[patient_df[\"patient_id\"].isin(temp_patients)]\n\nval_patients, test_patients = train_test_split(\n    temp_df[\"patient_id\"],\n    test_size=0.5,\n    random_state=42,\n    stratify=temp_df[\"cancer\"],\n)\n\ntrain_patients = set(train_patients)\nval_patients   = set(val_patients)\ntest_patients  = set(test_patients)\n\nprint(\"Patientes en train :\", len(train_patients))\nprint(\"Patientes en val   :\", len(val_patients))\nprint(\"Patientes en test  :\", len(test_patients))\n\n# Application au niveau des images\ntrain_df = meta_df[meta_df[\"patient_id\"].isin(train_patients)].reset_index(drop=True)\nval_df   = meta_df[meta_df[\"patient_id\"].isin(val_patients)].reset_index(drop=True)\ntest_df  = meta_df[meta_df[\"patient_id\"].isin(test_patients)].reset_index(drop=True)\n\nprint(\"Taille train_df :\", train_df.shape)\nprint(\"Taille val_df   :\", val_df.shape)\nprint(\"Taille test_df  :\", test_df.shape)\n\nprint(\"\\nRépartition des labels dans train_df :\")\nprint(train_df[\"cancer\"].value_counts())\n\nprint(\"\\nRépartition des labels dans val_df :\")\nprint(val_df[\"cancer\"].value_counts())\n\nprint(\"\\nRépartition des labels dans test_df :\")\nprint(test_df[\"cancer\"].value_counts())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-08T20:53:36.799492Z","iopub.execute_input":"2025-12-08T20:53:36.799738Z","iopub.status.idle":"2025-12-08T20:53:36.826203Z","shell.execute_reply.started":"2025-12-08T20:53:36.799719Z","shell.execute_reply":"2025-12-08T20:53:36.825417Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# Section 6 - Dataset PyTorch & DataLoader (train / val / test)\n# =============================================================================\n\nIMG_ROOT = OUTPUT_DIR  # racine pour les chemins relatifs \"filepath\"\n\nclass MammographyDataset(Dataset):\n    \"\"\"\n    Dataset PyTorch pour lire les PNG pré-traitées.\n    \"\"\"\n    def __init__(self, df, img_root, transform=None):\n        self.df = df.reset_index(drop=True)\n        self.img_root = Path(img_root)\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img_path = self.img_root / row[\"filepath\"]\n        label = float(row[\"cancer\"])\n\n        img = Image.open(img_path).convert(\"L\")\n        img = img.convert(\"RGB\")\n\n        if self.transform is not None:\n            img = self.transform(img)\n\n        return img, torch.tensor(label, dtype=torch.float32)\n\n\ntrain_transforms = T.Compose([\n    T.RandomHorizontalFlip(p=0.5),\n    T.RandomRotation(degrees=5),\n    T.ToTensor(),\n    T.Normalize(mean=[0.485, 0.456, 0.406],\n                std=[0.229, 0.224, 0.225]),\n])\n\nval_test_transforms = T.Compose([\n    T.ToTensor(),\n    T.Normalize(mean=[0.485, 0.456, 0.406],\n                std=[0.229, 0.224, 0.225]),\n])\n\ntrain_dataset = MammographyDataset(train_df, IMG_ROOT, transform=train_transforms)\nval_dataset   = MammographyDataset(val_df,   IMG_ROOT, transform=val_test_transforms)\ntest_dataset  = MammographyDataset(test_df,  IMG_ROOT, transform=val_test_transforms)\n\nBATCH_SIZE = 16\n\ntrain_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True,  num_workers=2)\nval_loader   = DataLoader(val_dataset,   batch_size=BATCH_SIZE, shuffle=False, num_workers=2)\ntest_loader  = DataLoader(test_dataset,  batch_size=BATCH_SIZE, shuffle=False, num_workers=2)\n\nprint(\"Batches train :\", len(train_loader))\nprint(\"Batches val   :\", len(val_loader))\nprint(\"Batches test  :\", len(test_loader))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-08T20:53:36.827131Z","iopub.execute_input":"2025-12-08T20:53:36.827409Z","iopub.status.idle":"2025-12-08T20:53:36.843316Z","shell.execute_reply.started":"2025-12-08T20:53:36.827378Z","shell.execute_reply":"2025-12-08T20:53:36.842569Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# Section 7 - Modèle CNN & entraînement (train + val)\n# =============================================================================\n\n# Ici on NE TÉLÉCHARGE PAS de poids pré-entraînés (pas d'accès internet),\n# on crée un ResNet18 \"from scratch\" (weights=None).\n# Si tu préfères un modèle plus petit, tu peux aussi utiliser resnet10 ou un CNN custom.\n\n# Version sans pré-entraînement (pas de téléchargement)\nbase_model = models.resnet18(weights=None)  # ou pretrained=False pour les anciennes versions\n\nnum_features = base_model.fc.in_features\nbase_model.fc = nn.Linear(num_features, 1)  # sortie binaire (1 logit)\n\nmodel = base_model.to(device)\n\n# Perte & optimiseur\ncriterion = nn.BCEWithLogitsLoss()\noptimizer = torch.optim.Adam(model.parameters(), lr=1e-3)  # un peu plus haut que 1e-4 sans pré-training\n\n\ndef train_one_epoch(model, loader, optimizer, criterion, device):\n    model.train()\n    running_loss = 0.0\n    for images, labels in loader:\n        images = images.to(device)\n        labels = labels.to(device).unsqueeze(1)\n\n        optimizer.zero_grad()\n        outputs = model(images)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n\n        running_loss += loss.item() * images.size(0)\n\n    return running_loss / len(loader.dataset)\n\n\ndef evaluate(model, loader, criterion, device):\n    model.eval()\n    running_loss = 0.0\n    all_labels = []\n    all_logits = []\n\n    with torch.no_grad():\n        for images, labels in loader:\n            images = images.to(device)\n            labels = labels.to(device).unsqueeze(1)\n\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n\n            running_loss += loss.item() * images.size(0)\n            all_labels.append(labels.cpu().numpy())\n            all_logits.append(outputs.cpu().numpy())\n\n    epoch_loss = running_loss / len(loader.dataset)\n    all_labels = np.vstack(all_labels).ravel()\n    all_logits = np.vstack(all_logits).ravel()\n    return epoch_loss, all_labels, all_logits\n\n\n# Nombre d'epochs (tu peux monter à 8–10 si tu as le temps de calcul)\nEPOCHS = 5\ntrain_losses, val_losses = [], []\nbest_val_auc = 0.0\nbest_state_dict = None\n\nfor epoch in range(1, EPOCHS + 1):\n    train_loss = train_one_epoch(model, train_loader, optimizer, criterion, device)\n    val_loss, val_labels, val_logits = evaluate(model, val_loader, criterion, device)\n\n    # Passage logits -> probabilités\n    val_probs = 1 / (1 + np.exp(-val_logits))\n    try:\n        val_auc = roc_auc_score(val_labels, val_probs)\n    except ValueError:\n        val_auc = np.nan\n\n    train_losses.append(train_loss)\n    val_losses.append(val_loss)\n\n    print(\n        f\"Epoch {epoch}/{EPOCHS} - \"\n        f\"Train loss: {train_loss:.4f} | \"\n        f\"Val loss: {val_loss:.4f} | \"\n        f\"Val AUC: {val_auc:.4f}\"\n    )\n\n    # Sauvegarde du meilleur modèle selon l'AUC validation\n    if not np.isnan(val_auc) and val_auc > best_val_auc:\n        best_val_auc = val_auc\n        best_state_dict = model.state_dict().copy()\n\n# Reload du meilleur modèle\nif best_state_dict is not None:\n    model.load_state_dict(best_state_dict)\n    print(f\"Meilleur modèle rechargé (Val AUC = {best_val_auc:.4f})\")\n\n# Visualisation des courbes de loss\nplt.figure(figsize=(6, 4))\nplt.plot(range(1, EPOCHS + 1), train_losses, label=\"Train loss\")\nplt.plot(range(1, EPOCHS + 1), val_losses,   label=\"Val loss\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Loss\")\nplt.title(\"Évolution des pertes (train/val)\")\nplt.legend()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-08T20:56:10.38427Z","iopub.execute_input":"2025-12-08T20:56:10.385021Z","iopub.status.idle":"2025-12-08T20:57:08.785113Z","shell.execute_reply.started":"2025-12-08T20:56:10.384992Z","shell.execute_reply":"2025-12-08T20:57:08.784433Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# Section 8 - Évaluation détaillée sur le set de TEST\n# =============================================================================\n\ntest_loss, test_labels, test_logits = evaluate(model, test_loader, criterion, device)\ntest_probs = 1 / (1 + np.exp(-test_logits))\ntest_preds = (test_probs >= 0.5).astype(int)\n\nprint(f\"Test loss : {test_loss:.4f}\")\n\nacc = accuracy_score(test_labels, test_preds)\nprec, rec, f1, _ = precision_recall_fscore_support(\n    test_labels, test_preds, average=\"binary\", zero_division=0\n)\n\nprint(f\"Accuracy : {acc:.3f}\")\nprint(f\"Précision : {prec:.3f}\")\nprint(f\"Rappel (sensibilité) : {rec:.3f}\")\nprint(f\"F1-score : {f1:.3f}\")\n\ntry:\n    auc = roc_auc_score(test_labels, test_probs)\n    print(f\"AUC : {auc:.3f}\")\nexcept ValueError:\n    auc = np.nan\n    print(\"AUC non calculable.\")\n\nprint(\"\\nRapport de classification :\")\nprint(classification_report(test_labels, test_preds, digits=3))\n\ncm = confusion_matrix(test_labels, test_preds)\nfig, ax = plt.subplots()\nsns.heatmap(cm, annot=True, fmt=\"d\", cmap=\"Blues\", ax=ax)\nax.set_xlabel(\"Prédiction\")\nax.set_ylabel(\"Vérité terrain\")\nax.set_title(\"Matrice de confusion (test)\")\nplt.show()\n\nif not np.isnan(auc):\n    fpr, tpr, thresholds = roc_curve(test_labels, test_probs)\n    fig, ax = plt.subplots()\n    ax.plot(fpr, tpr, label=f\"ROC (AUC = {auc:.3f})\")\n    ax.plot([0, 1], [0, 1], \"k--\", label=\"Aléatoire\")\n    ax.set_xlabel(\"Taux de faux positifs\")\n    ax.set_ylabel(\"Taux de vrais positifs\")\n    ax.set_title(\"Courbe ROC (test)\")\n    ax.legend()\n    plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-08T20:57:42.807285Z","iopub.execute_input":"2025-12-08T20:57:42.807562Z","iopub.status.idle":"2025-12-08T20:57:44.42835Z","shell.execute_reply.started":"2025-12-08T20:57:42.807541Z","shell.execute_reply":"2025-12-08T20:57:44.427506Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# Section 9 - Visualisation de quelques prédictions sur le set de TEST\n# =============================================================================\n\ndef show_predictions(model, dataset, n=6):\n    model.eval()\n    indices = np.random.choice(len(dataset), size=min(n, len(dataset)), replace=False)\n\n    fig, axes = plt.subplots(2, max(1, n // 2), figsize=(12, 6))\n    axes = np.array(axes).ravel()\n\n    with torch.no_grad():\n        for ax, idx in zip(axes, indices):\n            img, label = dataset[idx]\n            img_input = img.unsqueeze(0).to(device)\n\n            logits = model(img_input)\n            prob = torch.sigmoid(logits)[0].item()\n            pred = int(prob >= 0.5)\n\n            # Dénormalisation approximative pour affichage\n            img_disp = img.cpu().numpy().transpose(1, 2, 0)\n            img_disp = img_disp * np.array([0.229, 0.224, 0.225]) + np.array([0.485, 0.456, 0.406])\n            img_disp = np.clip(img_disp, 0, 1)\n\n            ax.imshow(img_disp, cmap=\"gray\")\n            ax.axis(\"off\")\n            ax.set_title(f\"GT={int(label.item())} | Pred={pred}\\nProb={prob:.2f}\")\n\n    plt.tight_layout()\n    plt.show()\n\nprint(\"Quelques prédictions sur le set de test :\")\nshow_predictions(model, test_dataset, n=6)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-08T20:57:44.429931Z","iopub.execute_input":"2025-12-08T20:57:44.430231Z","iopub.status.idle":"2025-12-08T20:57:45.827731Z","shell.execute_reply.started":"2025-12-08T20:57:44.4302Z","shell.execute_reply":"2025-12-08T20:57:45.827094Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}