{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":10338,"databundleVersionId":862042,"isSourceIdPinned":false}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"794ed3c9-94de-4829-b8bf-48cd9a3e373c","cell_type":"markdown","source":"# RSNA Pneumonia Pipeline (adapté depuis ton notebook)\n\nOn garde **la même logique générale** :\n1. Charger les données\n2. Créer `train/val/test loaders`\n3. Entraîner le modèle\n4. Sauvegarder `best_model.pth`\n5. Évaluer sur un **test hold-out**\n6. Optionnel : faire l'inférence sur `stage_2_test_images` pour une soumission Kaggle\n","metadata":{}},{"id":"e8702b4f-786c-4044-ba46-155fd23137aa","cell_type":"code","source":"import os, time\nimport pandas as pd\nfrom PIL import Image\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\n\nfrom sklearn.model_selection import train_test_split\n\nprint(\"Torch:\", torch.__version__)\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Device:\", device)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-29T00:22:30.013554Z","iopub.execute_input":"2026-03-29T00:22:30.013990Z","iopub.status.idle":"2026-03-29T00:22:45.540440Z","shell.execute_reply.started":"2026-03-29T00:22:30.013966Z","shell.execute_reply":"2026-03-29T00:22:45.539801Z"}},"outputs":[],"execution_count":null},{"id":"049d45b0-dc80-4d4f-b7e5-b3d6779ddfaa","cell_type":"markdown","source":"## Dataset RSNA\nLe dataset RSNA n'est **pas** structuré en `train/val/test` comme `ImageFolder`.\nIl faut donc :\n- lire les labels depuis `stage_2_train_labels.csv`\n- construire un label image-level (`0/1`) par `patientId`\n- faire un split `train/val/test` à partir du train officiel\n- garder `stage_2_test_images` pour l'inférence uniquement (pas d'accuracy car pas de labels)\n","metadata":{}},{"id":"ad2a3ed8-1106-4100-bc5d-5079e10c45e8","cell_type":"code","source":"# Adapte ce chemin si besoin\nRSNA_DIR = \"/kaggle/input/competitions/rsna-pneumonia-detection-challenge\"\n\nTRAIN_IMG_DIR = os.path.join(RSNA_DIR, \"stage_2_train_images\")\nTEST_IMG_DIR  = os.path.join(RSNA_DIR, \"stage_2_test_images\")\nTRAIN_CSV     = os.path.join(RSNA_DIR, \"stage_2_train_labels.csv\")\nDETAIL_CSV    = os.path.join(RSNA_DIR, \"stage_2_detailed_class_info.csv\")\n\nIMG_SIZE = 224\nBATCH_SIZE = 32\nNUM_EPOCHS = 5\nLR = 1e-3\nSEED = 42\n\ntransform = transforms.Compose([\n    transforms.Resize((IMG_SIZE, IMG_SIZE)),\n    transforms.ToTensor(),\n])\n\nassert os.path.exists(TRAIN_IMG_DIR), f\"Introuvable: {TRAIN_IMG_DIR}\"\nassert os.path.exists(TRAIN_CSV), f\"Introuvable: {TRAIN_CSV}\"\n\nlabels_df = pd.read_csv(TRAIN_CSV)\nlabels_df.head()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-29T00:23:09.324164Z","iopub.execute_input":"2026-03-29T00:23:09.324481Z","iopub.status.idle":"2026-03-29T00:23:09.424551Z","shell.execute_reply.started":"2026-03-29T00:23:09.324451Z","shell.execute_reply":"2026-03-29T00:23:09.423895Z"}},"outputs":[],"execution_count":null},{"id":"0a16a195-e6f8-4a30-bef2-f8ec878669b0","cell_type":"code","source":"# Le CSV contient potentiellement plusieurs lignes par patient (plusieurs boxes).\n# Pour garder ta logique de classification binaire:\n# label = 1 si au moins une ligne a Target=1, sinon 0\n\nimage_labels = (\n    labels_df.groupby(\"patientId\", as_index=False)[\"Target\"]\n    .max()\n    .rename(columns={\"Target\": \"label\"})\n)\n\nimage_labels[\"path\"] = image_labels[\"patientId\"].apply(\n    lambda x: os.path.join(TRAIN_IMG_DIR, f\"{x}.dcm\")\n)\n\nprint(\"Nb images annotées:\", len(image_labels))\nprint(image_labels[\"label\"].value_counts())\n\ntrain_df, temp_df = train_test_split(\n    image_labels,\n    test_size=0.30,\n    random_state=SEED,\n    stratify=image_labels[\"label\"]\n)\n\nval_df, test_df = train_test_split(\n    temp_df,\n    test_size=0.50,\n    random_state=SEED,\n    stratify=temp_df[\"label\"]\n)\n\nprint(\"Train:\", len(train_df), train_df[\"label\"].value_counts().to_dict())\nprint(\"Val  :\", len(val_df), val_df[\"label\"].value_counts().to_dict())\nprint(\"Test :\", len(test_df), test_df[\"label\"].value_counts().to_dict())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-29T00:23:17.108753Z","iopub.execute_input":"2026-03-29T00:23:17.109048Z","iopub.status.idle":"2026-03-29T00:23:17.207692Z","shell.execute_reply.started":"2026-03-29T00:23:17.109022Z","shell.execute_reply":"2026-03-29T00:23:17.207070Z"}},"outputs":[],"execution_count":null},{"id":"3513d4d9-be71-425a-9dcf-b51907fdd4be","cell_type":"code","source":"# Si pydicom n'est pas installé sur Kaggle/Colab, décommente :\n!pip install pydicom -q\n\nimport pydicom\nimport numpy as np\n\nclass RSNAClassificationDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df.reset_index(drop=True)\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def _read_dicom(self, path):\n        dcm = pydicom.dcmread(path)\n        img = dcm.pixel_array.astype(np.float32)\n\n        # normalisation min-max vers [0,255]\n        img = img - img.min()\n        if img.max() > 0:\n            img = img / img.max()\n        img = (img * 255).astype(np.uint8)\n\n        # convertir en RGB pour ResNet\n        img = Image.fromarray(img).convert(\"RGB\")\n        return img\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        image = self._read_dicom(row[\"path\"])\n        label = int(row[\"label\"])\n\n        if self.transform:\n            image = self.transform(image)\n\n        return image, label\n\ntrain_dataset = RSNAClassificationDataset(train_df, transform=transform)\nval_dataset   = RSNAClassificationDataset(val_df, transform=transform)\ntest_dataset  = RSNAClassificationDataset(test_df, transform=transform)\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\nclasses = [\"Normal\", \"Pneumonia\"]\nprint(classes)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-29T00:23:54.451913Z","iopub.execute_input":"2026-03-29T00:23:54.452422Z","iopub.status.idle":"2026-03-29T00:23:59.669782Z","shell.execute_reply.started":"2026-03-29T00:23:54.452394Z","shell.execute_reply":"2026-03-29T00:23:59.668774Z"}},"outputs":[],"execution_count":null},{"id":"9de8f6c0-5869-4824-ba40-bfa278eb79ae","cell_type":"markdown","source":"## MSR Model (on garde la logique)\nJ'ai laissé la même logique que ton notebook :\n- soit tu remplaces par ton import réel `MSR`\n- soit tu testes avec un backbone de secours\n","metadata":{}},{"id":"5e38b54b-63ea-4edd-b366-41fb5ae211a5","cell_type":"code","source":"# IMPORTANT: remplace par ton import réel si ton repo est monté\n# from model.MSR import MSR\n\nclass DummyMSR(nn.Module):\n    def __init__(self, num_classes=2):\n        super().__init__()\n        self.backbone = torch.hub.load('pytorch/vision', 'resnet18', pretrained=True)\n        self.backbone.fc = nn.Linear(self.backbone.fc.in_features, num_classes)\n\n    def forward(self, x):\n        return self.backbone(x)\n\nmodel = DummyMSR(num_classes=2).to(device)\nprint(\"Model chargé\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-29T00:24:12.392542Z","iopub.execute_input":"2026-03-29T00:24:12.392886Z","iopub.status.idle":"2026-03-29T00:24:15.250959Z","shell.execute_reply.started":"2026-03-29T00:24:12.392852Z","shell.execute_reply":"2026-03-29T00:24:15.250277Z"}},"outputs":[],"execution_count":null},{"id":"7062c9c0-f71b-4f1e-81ee-efafd20431b3","cell_type":"markdown","source":"## Training\nMême pipeline que ton notebook, avec ajout du `val_loss`.\n","metadata":{}},{"id":"6b6b6ba2-4e73-42b8-93ab-2f3fe9ac89e5","cell_type":"code","source":"criterion = nn.CrossEntropyLoss()\noptimizer = optim.Adam(model.parameters(), lr=LR)\n\nbest_acc = 0.0\nos.makedirs(\"log\", exist_ok=True)\n\nfor epoch in range(NUM_EPOCHS):\n    print(f\"\\nEpoch {epoch+1}/{NUM_EPOCHS}\")\n\n    model.train()\n    train_loss = 0.0\n    correct, total = 0, 0\n\n    for x, y in train_loader:\n        x, y = x.to(device), y.to(device)\n\n        optimizer.zero_grad()\n        out = model(x)\n        loss = criterion(out, y)\n        loss.backward()\n        optimizer.step()\n\n        train_loss += loss.item() * x.size(0)\n        _, pred = out.max(1)\n        total += y.size(0)\n        correct += pred.eq(y).sum().item()\n\n    train_loss /= total\n    train_acc = 100 * correct / total\n    print(f\"Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.2f}\")\n\n    model.eval()\n    val_loss = 0.0\n    correct, total = 0, 0\n\n    with torch.no_grad():\n        for x, y in val_loader:\n            x, y = x.to(device), y.to(device)\n            out = model(x)\n            loss = criterion(out, y)\n\n            val_loss += loss.item() * x.size(0)\n            _, pred = out.max(1)\n            total += y.size(0)\n            correct += pred.eq(y).sum().item()\n\n    val_loss /= total\n    val_acc = 100 * correct / total\n    print(f\"Val   Loss: {val_loss:.4f} | Val   Acc: {val_acc:.2f}\")\n\n    if val_acc > best_acc:\n        best_acc = val_acc\n        torch.save(model.state_dict(), \"log/best_model.pth\")\n        print(\"Saved best model\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-29T00:24:26.116188Z","iopub.execute_input":"2026-03-29T00:24:26.116468Z","iopub.status.idle":"2026-03-29T00:43:01.928244Z","shell.execute_reply.started":"2026-03-29T00:24:26.116445Z","shell.execute_reply":"2026-03-29T00:43:01.927237Z"}},"outputs":[],"execution_count":null},{"id":"728c0fa1-1ecb-42eb-8662-3481e320ca28","cell_type":"code","source":"from sklearn.metrics import (\n    confusion_matrix,\n    classification_report,\n    precision_score,\n    recall_score,\n    f1_score,\n    roc_auc_score,\n    accuracy_score\n)\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport numpy as np\nimport torch\n\nmodel.eval()\n\nall_labels = []\nall_preds = []\nall_probs = []\n\nwith torch.no_grad():\n    for images, labels in test_loader:\n        images = images.to(device)\n        labels = labels.to(device)\n\n        outputs = model(images)\n\n        # Cas 1 : sortie binaire avec 1 seul logit\n        if outputs.ndim == 1 or outputs.shape[1] == 1:\n            outputs = outputs.squeeze()\n            probs = torch.sigmoid(outputs)\n            preds = (probs >= 0.5).long()\n\n            # labels peut être [B] ou [B,1]\n            if labels.ndim > 1:\n                labels = labels.squeeze()\n\n            labels = labels.long()\n\n            all_labels.extend(labels.cpu().numpy())\n            all_preds.extend(preds.cpu().numpy())\n            all_probs.extend(probs.cpu().numpy())\n\n        # Cas 2 : sortie multi-classes avec 2 logits\n        else:\n            probs = torch.softmax(outputs, dim=1)\n            preds = torch.argmax(probs, dim=1)\n\n            # si labels one-hot -> convertir en indices\n            if labels.ndim > 1 and labels.shape[1] > 1:\n                labels = torch.argmax(labels, dim=1)\n            else:\n                labels = labels.squeeze()\n\n            labels = labels.long()\n\n            all_labels.extend(labels.cpu().numpy())\n            all_preds.extend(preds.cpu().numpy())\n            all_probs.extend(probs[:, 1].cpu().numpy())  # proba classe positive\n\nall_labels = np.array(all_labels).astype(int).reshape(-1)\nall_preds = np.array(all_preds).astype(int).reshape(-1)\nall_probs = np.array(all_probs).reshape(-1)\n\nprint(\"Shape all_labels:\", all_labels.shape)\nprint(\"Shape all_preds :\", all_preds.shape)\nprint(\"Shape all_probs :\", all_probs.shape)\nprint(\"Exemple labels  :\", all_labels[:10])\nprint(\"Exemple preds   :\", all_preds[:10])\n\n# Matrice de confusion\ncm = confusion_matrix(all_labels, all_preds)\n\nplt.figure(figsize=(6, 5))\nsns.heatmap(\n    cm,\n    annot=True,\n    fmt=\"d\",\n    cmap=\"Blues\",\n    xticklabels=[\"Normal\", \"Pneumonia\"],\n    yticklabels=[\"Normal\", \"Pneumonia\"]\n)\nplt.xlabel(\"Predicted\")\nplt.ylabel(\"True\")\nplt.title(\"Confusion Matrix\")\nplt.show()\n\n# Métriques\nacc = accuracy_score(all_labels, all_preds)\nprecision = precision_score(all_labels, all_preds, zero_division=0)\nrecall = recall_score(all_labels, all_preds, zero_division=0)\nf1 = f1_score(all_labels, all_preds, zero_division=0)\n\ntry:\n    auc = roc_auc_score(all_labels, all_probs)\nexcept Exception:\n    auc = None\n\nprint(f\"Accuracy  : {acc:.4f}\")\nprint(f\"Precision : {precision:.4f}\")\nprint(f\"Recall    : {recall:.4f}\")\nprint(f\"F1-score  : {f1:.4f}\")\nif auc is not None:\n    print(f\"AUC ROC   : {auc:.4f}\")\nelse:\n    print(\"AUC ROC   : impossible à calculer\")\n\nprint(\"\\nClassification Report:\\n\")\nprint(classification_report(all_labels, all_preds, target_names=[\"Normal\", \"Pneumonia\"], zero_division=0))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-29T00:52:04.389130Z","iopub.execute_input":"2026-03-29T00:52:04.389478Z","iopub.status.idle":"2026-03-29T00:52:40.038815Z","shell.execute_reply.started":"2026-03-29T00:52:04.389445Z","shell.execute_reply":"2026-03-29T00:52:40.037904Z"}},"outputs":[],"execution_count":null},{"id":"dc137344-bbf6-4e79-bd53-5f0c9a502298","cell_type":"markdown","source":"## Test\nIci on évalue sur le **test split interne** créé depuis `stage_2_train_labels.csv`.\n\n> Important : le dossier `stage_2_test_images` Kaggle n'a pas de labels publics, donc on ne peut pas calculer une vraie accuracy dessus.\n","metadata":{}},{"id":"cf610eca-6a2c-41ca-9255-1c53dc935070","cell_type":"code","source":"model.load_state_dict(torch.load(\"log/best_model.pth\", map_location=device))\nmodel.eval()\n\ntest_loss = 0.0\ncorrect, total = 0, 0\n\nwith torch.no_grad():\n    for x, y in test_loader:\n        x, y = x.to(device), y.to(device)\n        out = model(x)\n        loss = criterion(out, y)\n\n        test_loss += loss.item() * x.size(0)\n        _, pred = out.max(1)\n        total += y.size(0)\n        correct += pred.eq(y).sum().item()\n\ntest_loss /= total\ntest_acc = 100 * correct / total\n\nprint(f\"Test Loss: {test_loss:.4f}\")\nprint(f\"Test Accuracy: {test_acc:.2f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-29T00:43:35.902112Z","iopub.execute_input":"2026-03-29T00:43:35.903033Z","iopub.status.idle":"2026-03-29T00:44:27.685102Z","shell.execute_reply.started":"2026-03-29T00:43:35.902993Z","shell.execute_reply":"2026-03-29T00:44:27.684259Z"}},"outputs":[],"execution_count":null},{"id":"70cbbfa1-b430-4809-b4f8-781e49c3c90e","cell_type":"markdown","source":"## Inference Kaggle sur `stage_2_test_images` (optionnel)\nCette partie sert à générer des prédictions pour la compétition.\nComme il n'y a pas de labels, on exporte juste un CSV.\n","metadata":{}},{"id":"43f05c3d-7be2-44c0-8ff8-234c4ac75f41","cell_type":"code","source":"test_image_ids = []\nif os.path.exists(TEST_IMG_DIR):\n    for fname in os.listdir(TEST_IMG_DIR):\n        if fname.endswith(\".dcm\"):\n            test_image_ids.append(fname.replace(\".dcm\", \"\"))\n\n    test_image_ids = sorted(test_image_ids)\n    print(\"Nb images test compétition:\", len(test_image_ids))\nelse:\n    print(\"Dossier stage_2_test_images introuvable\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-29T00:45:34.292652Z","iopub.execute_input":"2026-03-29T00:45:34.293434Z","iopub.status.idle":"2026-03-29T00:45:34.558450Z","shell.execute_reply.started":"2026-03-29T00:45:34.293398Z","shell.execute_reply":"2026-03-29T00:45:34.557902Z"}},"outputs":[],"execution_count":null},{"id":"fc26793e-761d-4dd8-a652-b60876f0b6f9","cell_type":"code","source":"class RSNATestInferenceDataset(Dataset):\n    def __init__(self, image_ids, image_dir, transform=None):\n        self.image_ids = image_ids\n        self.image_dir = image_dir\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.image_ids)\n\n    def __getitem__(self, idx):\n        patient_id = self.image_ids[idx]\n        path = os.path.join(self.image_dir, f\"{patient_id}.dcm\")\n\n        dcm = pydicom.dcmread(path)\n        img = dcm.pixel_array.astype(np.float32)\n        img = img - img.min()\n        if img.max() > 0:\n            img = img / img.max()\n        img = (img * 255).astype(np.uint8)\n        img = Image.fromarray(img).convert(\"RGB\")\n\n        if self.transform:\n            img = self.transform(img)\n\n        return img, patient_id\n\nif len(test_image_ids) > 0:\n    kaggle_test_dataset = RSNATestInferenceDataset(test_image_ids, TEST_IMG_DIR, transform=transform)\n    kaggle_test_loader = DataLoader(kaggle_test_dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=2)\n\n    preds = []\n    model.eval()\n    with torch.no_grad():\n        for x, patient_ids in kaggle_test_loader:\n            x = x.to(device)\n            out = model(x)\n            prob = torch.softmax(out, dim=1)[:, 1]   # proba classe pneumonia\n            pred = (prob > 0.5).long().cpu().numpy()\n\n            for pid, p in zip(patient_ids, pred):\n                preds.append([pid, int(p)])\n\n    submission_df = pd.DataFrame(preds, columns=[\"patientId\", \"Target\"])\n    submission_df.to_csv(\"stage_2_sample_submission_classification.csv\", index=False)\n    print(\"CSV sauvegardé:\", \"stage_2_sample_submission_classification.csv\")\n    submission_df.head()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-29T00:45:47.893264Z","iopub.execute_input":"2026-03-29T00:45:47.893903Z","iopub.status.idle":"2026-03-29T00:46:31.840511Z","shell.execute_reply.started":"2026-03-29T00:45:47.893873Z","shell.execute_reply":"2026-03-29T00:46:31.839753Z"}},"outputs":[],"execution_count":null}]}