{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.12.12"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"datasetVersion","sourceId":1019494,"datasetId":560711,"databundleVersionId":1048360}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport csv\nimport copy\nimport time\nimport math\nimport random\nfrom pathlib import Path\nfrom functools import lru_cache\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom PIL import Image\n\nimport pydicom\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.optim import lr_scheduler\nfrom torch.utils.data import Dataset, DataLoader\n\nimport torchvision\nfrom torchvision import models, transforms\n\nfrom sklearn.metrics import (\n    classification_report,\n    confusion_matrix,\n    precision_score,\n    recall_score,\n    f1_score,\n    roc_auc_score,\n)\nfrom tqdm.auto import tqdm\n\nprint(\"Torch version:\", torch.__version__)\nprint(\"Torchvision version:\", torchvision.__version__)\nprint(\"CUDA available:\", torch.cuda.is_available())","metadata":{"execution":{"iopub.status.busy":"2026-03-29T18:25:51.937909Z","iopub.execute_input":"2026-03-29T18:25:51.938174Z","iopub.status.idle":"2026-03-29T18:26:03.632880Z","shell.execute_reply.started":"2026-03-29T18:25:51.938140Z","shell.execute_reply":"2026-03-29T18:26:03.632171Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**configuration**","metadata":{}},{"cell_type":"code","source":"BATCH_SIZE = 8\nNUM_WORKERS = 0\nNUM_EPOCHS = 5\nLR = 1e-4\nSTEP_SIZE = 2\nGAMMA = 0.1\nINPUT_SIZE = 224\nNUM_CLASSES = 2\n\nCLASS_TO_IDX = {\n    \"NORMAL\": 0,\n    \"PNEUMONIA\": 1\n}\n\nIMG_EXTENSIONS = [\".jpg\", \".jpeg\", \".png\"]\n\nfrom pathlib import Path\nimport os\n\ndef find_dataset_root():\n    candidates = [\n        Path(\"/kaggle/input/pediatric-pneumonia-chest-xray/Pediatric Chest X-ray Pneumonia\"),\n        Path(\"/kaggle/input/andrewmvd/pediatric-pneumonia-chest-xray/Pediatric Chest X-ray Pneumonia\"),\n        Path(\"/kaggle/input/datasets/andrewmvd/pediatric-pneumonia-chest-xray/Pediatric Chest X-ray Pneumonia\"),\n    ]\n\n    for c in candidates:\n        if c.exists():\n            return c\n\n    for root, dirs, files in os.walk(\"/kaggle/input\"):\n        if os.path.basename(root) == \"Pediatric Chest X-ray Pneumonia\":\n            return Path(root)\n\n    raise FileNotFoundError(\"Impossible de trouver le dossier 'Pediatric Chest X-ray Pneumonia' dans /kaggle/input\")\n\nDATA_DIR = find_dataset_root()\n\nTRAIN_DIR = DATA_DIR / \"train\"\nVAL_DIR   = DATA_DIR / \"test\"   # on garde test comme validation\nTEST_DIR  = DATA_DIR / \"test\"   # si ton code utilise aussi TEST_DIR\n\nprint(\"DATA_DIR:\", DATA_DIR)\nprint(\"TRAIN_DIR:\", TRAIN_DIR)\nprint(\"VAL_DIR:\", VAL_DIR)\nprint(\"DATA_DIR exists:\", DATA_DIR.exists())\nprint(\"TRAIN_DIR exists:\", TRAIN_DIR.exists())\nprint(\"VAL_DIR exists:\", VAL_DIR.exists())\nprint(\"Subfolders:\", os.listdir(DATA_DIR))","metadata":{"execution":{"iopub.status.busy":"2026-03-29T18:32:32.771469Z","iopub.execute_input":"2026-03-29T18:32:32.772135Z","iopub.status.idle":"2026-03-29T18:32:32.786928Z","shell.execute_reply.started":"2026-03-29T18:32:32.772102Z","shell.execute_reply":"2026-03-29T18:32:32.786202Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Build Pediatric Pneumonia train/val tables**","metadata":{}},{"cell_type":"code","source":"def build_split_dataframe(split_dir, split_name):\n    rows = []\n\n    for class_name, label in CLASS_TO_IDX.items():\n        class_dir = split_dir / class_name\n        if not class_dir.exists():\n            print(f\"[Warning] Missing folder: {class_dir}\")\n            continue\n\n        for image_path in sorted(class_dir.iterdir()):\n            if image_path.is_file() and image_path.suffix.lower() in IMG_EXTENSIONS:\n                rows.append({\n                    \"split\": split_name,\n                    \"image_path\": str(image_path),\n                    \"filename\": image_path.name,\n                    \"label\": label,\n                    \"class_name\": class_name,\n                })\n\n    df = pd.DataFrame(rows)\n    if len(df) == 0:\n        raise ValueError(f\"No images found in {split_dir}\")\n\n    return df.reset_index(drop=True)\n\ntrain_df = build_split_dataframe(TRAIN_DIR, \"train\")\nval_df = build_split_dataframe(VAL_DIR, \"val\")\n\nprint(\"Train size:\", len(train_df))\nprint(\"Val size:\", len(val_df))\nprint(\"Train label counts:\")\nprint(train_df[\"label\"].value_counts().sort_index())\nprint(\"Val label counts:\")\nprint(val_df[\"label\"].value_counts().sort_index())\n\ntrain_df.head()","metadata":{"execution":{"iopub.status.busy":"2026-03-29T18:32:38.758572Z","iopub.execute_input":"2026-03-29T18:32:38.759103Z","iopub.status.idle":"2026-03-29T18:33:01.511714Z","shell.execute_reply.started":"2026-03-29T18:32:38.759072Z","shell.execute_reply":"2026-03-29T18:33:01.510916Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Using official dataset folders (train/test)**","metadata":{}},{"cell_type":"code","source":"print(\"Official split loaded successfully.\")\nprint(\"Train classes:\")\nprint(train_df[\"class_name\"].value_counts())\nprint(\"Validation classes:\")\nprint(val_df[\"class_name\"].value_counts())","metadata":{"execution":{"iopub.status.busy":"2026-03-29T18:33:26.142236Z","iopub.execute_input":"2026-03-29T18:33:26.142563Z","iopub.status.idle":"2026-03-29T18:33:26.150116Z","shell.execute_reply.started":"2026-03-29T18:33:26.142534Z","shell.execute_reply":"2026-03-29T18:33:26.149285Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Dataset & transforms**","metadata":{}},{"cell_type":"code","source":"mean = [0.485, 0.456, 0.406]\nstd = [0.229, 0.224, 0.225]\n\ndata_transforms = {\n    \"train\": transforms.Compose([\n        transforms.Resize((INPUT_SIZE, INPUT_SIZE)),\n        transforms.RandomHorizontalFlip(p=0.5),\n        transforms.ToTensor(),\n        transforms.Normalize(mean, std),\n    ]),\n    \"val\": transforms.Compose([\n        transforms.Resize((INPUT_SIZE, INPUT_SIZE)),\n        transforms.ToTensor(),\n        transforms.Normalize(mean, std),\n    ]),\n}\n\n@lru_cache(maxsize=2048)\ndef load_image_as_rgb(path_str):\n    return Image.open(path_str).convert(\"RGB\")\n\nclass PediatricChestXrayDataset(Dataset):\n    def __init__(self, dataframe, transform=None, return_labels=True):\n        df = dataframe.reset_index(drop=True).copy()\n\n        self.image_paths = df[\"image_path\"].tolist()\n        self.filenames = df[\"filename\"].tolist()\n        self.transform = transform\n        self.return_labels = return_labels\n        self.labels = df[\"label\"].astype(int).tolist() if return_labels else None\n\n    def __len__(self):\n        return len(self.image_paths)\n\n    def __getitem__(self, idx):\n        image = load_image_as_rgb(self.image_paths[idx])\n\n        if self.transform is not None:\n            image = self.transform(image)\n\n        if self.return_labels:\n            return image, self.labels[idx]\n        return image\n\n    def get_filename(self, idx):\n        return self.filenames[idx]\n\ntrain_dataset = PediatricChestXrayDataset(train_df, transform=data_transforms[\"train\"], return_labels=True)\nval_dataset = PediatricChestXrayDataset(val_df, transform=data_transforms[\"val\"], return_labels=True)\n\nimage_datasets = {\n    \"train\": train_dataset,\n    \"val\": val_dataset,\n}\n\ndataloaders = {\n    split: DataLoader(\n        image_datasets[split],\n        batch_size=BATCH_SIZE,\n        shuffle=(split == \"train\"),\n        num_workers=NUM_WORKERS,\n        pin_memory=torch.cuda.is_available(),\n    )\n    for split in [\"train\", \"val\"]\n}\n\ndataset_sizes = {split: len(image_datasets[split]) for split in [\"train\", \"val\"]}\nclass_names = [\"normal\", \"pneumonia\"]\n\nprint(\"Dataset sizes:\", dataset_sizes)","metadata":{"execution":{"iopub.status.busy":"2026-03-29T18:33:31.367129Z","iopub.execute_input":"2026-03-29T18:33:31.367417Z","iopub.status.idle":"2026-03-29T18:33:31.380235Z","shell.execute_reply.started":"2026-03-29T18:33:31.367392Z","shell.execute_reply":"2026-03-29T18:33:31.379505Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Visualisation rapide**","metadata":{}},{"cell_type":"code","source":"def imshow(inp, title=None):\n    inp = inp.numpy().transpose((1, 2, 0))\n    inp = np.array(std) * inp + np.array(mean)\n    inp = np.clip(inp, 0, 1)\n    plt.figure(figsize=(10, 6))\n    plt.imshow(inp)\n    if title is not None:\n        plt.title(title)\n    plt.axis(\"off\")\n    plt.show()\n\ninputs, classes = next(iter(dataloaders[\"train\"]))\ngrid = torchvision.utils.make_grid(inputs[:8])\nimshow(grid, title=[class_names[int(x)] for x in classes[:8]])","metadata":{"execution":{"iopub.status.busy":"2026-03-29T18:33:38.167736Z","iopub.execute_input":"2026-03-29T18:33:38.168484Z","iopub.status.idle":"2026-03-29T18:33:38.926767Z","shell.execute_reply.started":"2026-03-29T18:33:38.168424Z","shell.execute_reply":"2026-03-29T18:33:38.925868Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Utilities**","metadata":{}},{"cell_type":"code","source":"OUTPUT_DIR = Path(\"./outputs\")\nOUTPUT_DIR.mkdir(parents=True, exist_ok=True)\n\n\ndef plot_history(val_values, train_values, metric_name, output_dir=OUTPUT_DIR):\n    plt.figure(figsize=(8, 5))\n    plt.title(f\"{metric_name} after epoch: {len(train_values)}\")\n    plt.xlabel(\"Epoch\")\n    plt.ylabel(metric_name)\n    plt.plot(range(1, len(train_values) + 1), train_values, label=f\"Train {metric_name}\")\n    plt.plot(range(1, len(val_values) + 1), val_values, label=f\"Validation {metric_name}\")\n    plt.legend()\n    plt.tight_layout()\n    plt.savefig(output_dir / f\"{metric_name.lower()}_resnet18.png\")\n    plt.close()\n\ndef build_model(num_classes):\n    model = models.resnet18(weights=models.ResNet18_Weights.DEFAULT)\n    num_ftrs = model.fc.in_features\n    model.fc = nn.Linear(num_ftrs, num_classes)\n    return model.to(DEVICE)\n\ndef forward_with_loss(model, inputs, labels, criterion):\n    logits = model(inputs)\n    loss = criterion(logits, labels)\n    return logits, loss","metadata":{"execution":{"iopub.status.busy":"2026-03-29T18:35:19.127396Z","iopub.execute_input":"2026-03-29T18:35:19.128006Z","iopub.status.idle":"2026-03-29T18:35:19.134136Z","shell.execute_reply.started":"2026-03-29T18:35:19.127975Z","shell.execute_reply":"2026-03-29T18:35:19.133502Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Training function**","metadata":{}},{"cell_type":"code","source":"def train_model(model, criterion, optimizer, scheduler, num_epochs=25, model_name=\"resnet18\"):\n    since = time.time()\n    best_model_wts = copy.deepcopy(model.state_dict())\n    best_acc = 0.0\n\n    train_loss_hist, val_loss_hist = [], []\n    train_acc_hist, val_acc_hist = [], []\n\n    for epoch in range(num_epochs):\n        print(f\"\\nEpoch {epoch+1}/{num_epochs}\")\n        print(\"-\" * 30)\n\n        for phase in [\"train\", \"val\"]:\n            if phase == \"train\":\n                model.train()\n            else:\n                model.eval()\n\n            running_loss = 0.0\n            running_corrects = 0\n\n            progress_bar = tqdm(dataloaders[phase], desc=f\"{phase}\", leave=False)\n\n            for inputs, labels in progress_bar:\n                inputs = inputs.to(DEVICE, non_blocking=True)\n                labels = labels.to(DEVICE, non_blocking=True)\n\n                optimizer.zero_grad()\n\n                with torch.set_grad_enabled(phase == \"train\"):\n                    logits, loss = forward_with_loss(model, inputs, labels, criterion)\n                    _, preds = torch.max(logits, 1)\n\n                    if phase == \"train\":\n                        loss.backward()\n                        optimizer.step()\n\n                running_loss += loss.detach().item() * inputs.size(0)\n                running_corrects += torch.sum(preds == labels).item()\n\n            epoch_loss = running_loss / dataset_sizes[phase]\n            epoch_acc = running_corrects / dataset_sizes[phase]\n\n            print(f\"{phase} Loss: {epoch_loss:.4f} Acc: {epoch_acc:.4f}\")\n\n            if phase == \"train\":\n                train_loss_hist.append(epoch_loss)\n                train_acc_hist.append(epoch_acc)\n                scheduler.step()\n            else:\n                val_loss_hist.append(epoch_loss)\n                val_acc_hist.append(epoch_acc)\n\n                if epoch_acc > best_acc:\n                    best_acc = epoch_acc\n                    best_model_wts = copy.deepcopy(model.state_dict())\n\n        torch.save(model.state_dict(), OUTPUT_DIR / f\"{model_name}_last.pth\")\n\n    time_elapsed = time.time() - since\n    print(f\"\\nTraining complete in {time_elapsed/60:.1f} min\")\n    print(f\"Best val Acc: {best_acc:.4f}\")\n\n    model.load_state_dict(best_model_wts)\n    torch.save(model.state_dict(), OUTPUT_DIR / f\"{model_name}_best.pth\")\n\n    plot_history(val_loss_hist, train_loss_hist, \"Loss\")\n    plot_history(val_acc_hist, train_acc_hist, \"Accuracy\")\n\n    history_df = pd.DataFrame({\n        \"epoch\": list(range(1, len(train_loss_hist) + 1)),\n        \"train_loss\": train_loss_hist,\n        \"val_loss\": val_loss_hist,\n        \"train_acc\": train_acc_hist,\n        \"val_acc\": val_acc_hist,\n    })\n    history_df.to_csv(OUTPUT_DIR / f\"{model_name}_history.csv\", index=False)\n\n    return model","metadata":{"execution":{"iopub.status.busy":"2026-03-29T18:35:25.220179Z","iopub.execute_input":"2026-03-29T18:35:25.220785Z","iopub.status.idle":"2026-03-29T18:35:25.230346Z","shell.execute_reply.started":"2026-03-29T18:35:25.220755Z","shell.execute_reply":"2026-03-29T18:35:25.229646Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Export probabilities to CSV**","metadata":{}},{"cell_type":"code","source":"\ndef export_probabilities(model, dataset, csv_path, labels_csv_path):\n    loader = DataLoader(\n        dataset,\n        batch_size=1,\n        shuffle=False,\n        num_workers=0,\n        pin_memory=torch.cuda.is_available(),\n    )\n\n    model.eval()\n    total = 0\n    correct = 0\n    rows = []\n    label_rows = []\n\n    with torch.no_grad():\n        for i, (images, labels) in enumerate(tqdm(loader, desc=f\"Export -> {Path(csv_path).name}\")):\n            images = images.to(DEVICE, non_blocking=True)\n            labels = labels.to(DEVICE, non_blocking=True)\n\n            outputs = model(images)\n            probs = torch.softmax(outputs, dim=1).cpu().numpy()[0].tolist()\n            _, predicted = torch.max(outputs, 1)\n\n            sample_name = dataset.get_filename(i)\n\n            total += labels.size(0)\n            correct += (predicted == labels).sum().item()\n\n            rows.append(probs + [sample_name])\n            label_rows.append([sample_name, int(labels.item())])\n\n    acc = correct / total if total > 0 else 0.0\n    print(f\"Export accuracy ({Path(csv_path).name}): {acc:.4f}\")\n\n    with open(csv_path, \"w\", newline=\"\") as f:\n        writer = csv.writer(f)\n        writer.writerows(rows)\n\n    with open(labels_csv_path, \"w\", newline=\"\") as f:\n        writer = csv.writer(f)\n        writer.writerows(label_rows)\n\n    print(\"Saved:\", csv_path)\n    print(\"Saved:\", labels_csv_path)","metadata":{"execution":{"iopub.status.busy":"2026-03-29T18:35:33.050538Z","iopub.execute_input":"2026-03-29T18:35:33.051300Z","iopub.status.idle":"2026-03-29T18:35:33.058291Z","shell.execute_reply.started":"2026-03-29T18:35:33.051269Z","shell.execute_reply":"2026-03-29T18:35:33.057510Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 11. Entraîner le modèle de base sur Pediatric Chest X-ray Pneumonia","metadata":{}},{"cell_type":"code","source":"import torch\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nprint(\"Using device:\", DEVICE)\n\n\ndef train_and_export_single_model():\n    model_name = \"resnet18\"\n    print(f\"\\n========== {model_name.upper()} ==========\")\n\n    model = build_model(NUM_CLASSES)\n\n    criterion = nn.CrossEntropyLoss()\n    optimizer = optim.Adam(model.parameters(), lr=LR)\n    scheduler = lr_scheduler.StepLR(optimizer, step_size=STEP_SIZE, gamma=GAMMA)\n\n    model = train_model(\n        model=model,\n        criterion=criterion,\n        optimizer=optimizer,\n        scheduler=scheduler,\n        num_epochs=NUM_EPOCHS,\n        model_name=model_name,\n    )\n\n    export_probabilities(\n        model=model,\n        dataset=train_dataset,\n        csv_path=OUTPUT_DIR / f\"{model_name}_train.csv\",\n        labels_csv_path=OUTPUT_DIR / \"train_labels.csv\",\n    )\n\n    export_probabilities(\n        model=model,\n        dataset=val_dataset,\n        csv_path=OUTPUT_DIR / f\"{model_name}_test.csv\",\n        labels_csv_path=OUTPUT_DIR / \"test_labels.csv\",\n    )\n\n    return model","metadata":{"execution":{"iopub.status.busy":"2026-03-29T18:38:25.329695Z","iopub.execute_input":"2026-03-29T18:38:25.330014Z","iopub.status.idle":"2026-03-29T18:38:25.336715Z","shell.execute_reply.started":"2026-03-29T18:38:25.329986Z","shell.execute_reply":"2026-03-29T18:38:25.335975Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Run RESNET18 only**","metadata":{}},{"cell_type":"code","source":"\ntrained_model = train_and_export_single_model()","metadata":{"execution":{"iopub.status.busy":"2026-03-29T18:38:32.696624Z","iopub.execute_input":"2026-03-29T18:38:32.697314Z","iopub.status.idle":"2026-03-29T18:51:28.433924Z","shell.execute_reply.started":"2026-03-29T18:38:32.697284Z","shell.execute_reply":"2026-03-29T18:51:28.433282Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Métrique et matrice de confusion**","metadata":{}},{"cell_type":"code","source":"\n\ndef getfile(filename):\n    df = pd.read_csv(filename, header=None)\n    df = np.asarray(df)[:, :-1].astype(np.float64)  # dernière colonne = nom fichier\n    return df\n\ndef getlabels(filename):\n    df = pd.read_csv(filename, header=None)\n    df = np.asarray(df)[:, 1]\n    return df.astype(int)\n\ndef predicting(prob_matrix):\n    return np.argmax(prob_matrix, axis=1).astype(int)\n\ndef show_metrics(labels, predictions, classes):\n    print(\"Classification Report:\")\n    print(classification_report(labels, predictions, target_names=classes, digits=4))\n\n    matrix = confusion_matrix(labels, predictions)\n    print(\"Confusion matrix:\")\n    print(matrix)\n\n    classwise_acc = matrix.diagonal() / matrix.sum(axis=1)\n    print(\"\\nClasswise Accuracy:\", classwise_acc)\n\n    plt.figure(figsize=(5, 4))\n    plt.imshow(matrix)\n    plt.title(\"Confusion Matrix\")\n    plt.xlabel(\"Predicted\")\n    plt.ylabel(\"True\")\n    plt.xticks(range(len(classes)), classes)\n    plt.yticks(range(len(classes)), classes)\n    for i in range(matrix.shape[0]):\n        for j in range(matrix.shape[1]):\n            plt.text(j, i, matrix[i, j], ha=\"center\", va=\"center\")\n    plt.tight_layout()\n    plt.show()\n\ndef get_weights(matrix):\n    weights = []\n    for i in range(matrix.shape[0]):\n        row = matrix[i]\n        w = 0.0\n        for j in range(row.shape[0]):\n            w += np.tanh(row[j])\n        weights.append(w)\n    return weights\n\ndef get_scores(labels, *argv):\n    count = len(argv)\n    metrics_matrix = np.zeros(shape=(4, count))\n    num_classes_local = np.unique(labels).shape[0]\n\n    for i, prob in enumerate(argv):\n        preds = predicting(prob)\n\n        if num_classes_local == 2:\n            pre = precision_score(labels, preds, zero_division=0)\n            rec = recall_score(labels, preds, zero_division=0)\n            f1 = f1_score(labels, preds, zero_division=0)\n            auc = roc_auc_score(labels, prob[:, 1])\n        else:\n            pre = precision_score(labels, preds, average=\"macro\", zero_division=0)\n            rec = recall_score(labels, preds, average=\"macro\", zero_division=0)\n            f1 = f1_score(labels, preds, average=\"macro\", zero_division=0)\n            auc = roc_auc_score(labels, prob, average=\"macro\", multi_class=\"ovo\")\n\n        metrics_matrix[:, i] = np.array([pre, rec, f1, auc])\n\n    weights = get_weights(metrics_matrix.T)\n    return weights\n\n# =========================\n# 13. Final metrics extraction\n# =========================\ndef extract_final_metrics(labels, prob_matrix, classes):\n    preds = predicting(prob_matrix)\n    num_classes_local = np.unique(labels).shape[0]\n\n    acc = accuracy_score(labels, preds)\n\n    if num_classes_local == 2:\n        pre = precision_score(labels, preds, zero_division=0)\n        rec = recall_score(labels, preds, zero_division=0)\n        f1 = f1_score(labels, preds, zero_division=0)\n        auc = roc_auc_score(labels, prob_matrix[:, 1])\n    else:\n        pre = precision_score(labels, preds, average=\"macro\", zero_division=0)\n        rec = recall_score(labels, preds, average=\"macro\", zero_division=0)\n        f1 = f1_score(labels, preds, average=\"macro\", zero_division=0)\n        auc = roc_auc_score(labels, prob_matrix, average=\"macro\", multi_class=\"ovo\")\n\n    metrics_dict = {\n        \"accuracy\": acc,\n        \"precision\": pre,\n        \"recall\": rec,\n        \"f1_score\": f1,\n        \"auc\": auc\n    }\n\n    print(\"Final Metrics:\")\n    for k, v in metrics_dict.items():\n        print(f\"{k}: {v:.4f}\")\n\n    show_metrics(labels, preds, classes)\n\n    return metrics_dict\n\ndef extract_final_metrics_from_csv(prob_csv, labels_csv, classes):\n    prob_matrix = getfile(prob_csv)\n    labels = getlabels(labels_csv)\n    return extract_final_metrics(labels, prob_matrix, classes)\n\ndef save_final_metrics(metrics_dict, output_csv):\n    df = pd.DataFrame([metrics_dict])\n    df.to_csv(output_csv, index=False)\n    print(f\"Métriques sauvegardées dans : {output_csv}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-29T19:00:16.581158Z","iopub.execute_input":"2026-03-29T19:00:16.581507Z","iopub.status.idle":"2026-03-29T19:00:16.596702Z","shell.execute_reply.started":"2026-03-29T19:00:16.581478Z","shell.execute_reply":"2026-03-29T19:00:16.595952Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"classes = [\"NORMAL\", \"PNEUMONIA\"]\n\nfinal_metrics = extract_final_metrics_from_csv(\n    prob_csv=OUTPUT_DIR / \"resnet18_test.csv\",\n    labels_csv=OUTPUT_DIR / \"test_labels.csv\",\n    classes=classes\n)\n\nsave_final_metrics(final_metrics, OUTPUT_DIR / \"final_metrics_resnet18.csv\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-29T19:00:27.954905Z","iopub.execute_input":"2026-03-29T19:00:27.955220Z","iopub.status.idle":"2026-03-29T19:00:28.051231Z","shell.execute_reply.started":"2026-03-29T19:00:27.955190Z","shell.execute_reply":"2026-03-29T19:00:28.050550Z"}},"outputs":[],"execution_count":null}]}