{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":13451,"datasetId":654585,"databundleVersionId":1188070}],"dockerImageVersionId":31011,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms, models\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import (\n    precision_score, recall_score, f1_score, roc_curve, auc, \n    confusion_matrix, classification_report\n)\nfrom PIL import Image\nimport matplotlib.pyplot as plt\nimport pydicom\nfrom tqdm import tqdm\nimport itertools\n\n\nDATA_DIR    = '/kaggle/input/rsna-intracranial-hemorrhage-detection/rsna-intracranial-hemorrhage-detection'\nCSV_PATH    = os.path.join(DATA_DIR, 'stage_2_train.csv')\nDCM_DIR     = os.path.join(DATA_DIR, 'stage_2_train')\nPNG_ROOT    = '/kaggle/working/png_data'\nOUTPUT_DIR  = '/kaggle/working/output_single_run'\nROC_DIR     = os.path.join(OUTPUT_DIR, 'roc_curves')\nCAM_DIR     = os.path.join(OUTPUT_DIR, 'gradcam')\nMETRICS_DIR = os.path.join(OUTPUT_DIR, 'metrics')\n\nfor d in [PNG_ROOT, OUTPUT_DIR, ROC_DIR, CAM_DIR, METRICS_DIR]:\n    os.makedirs(d, exist_ok=True)\n\nBATCH_SIZE    = 8\nNUM_EPOCHS    = 5\nLEARNING_RATE = 1e-4\nWEIGHT_DECAY  = 1e-6\nPATIENCE      = 10\nDEVICE        = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nCLASS_NAMES   = ['epidural', 'subdural', 'subarachnoid', 'intraparenchymal', 'intraventricular']\nBALANCE_N     = 500\nVAL_SPLIT_SIZE= 0.2\n\ndef display_sample_images_per_class(df, root_dir, classes, num_samples=5):\n    n_cols = num_samples\n    n_rows = len(classes)\n    \n    fig, axes = plt.subplots(n_rows, n_cols, figsize=(4 * n_cols, 4 * n_rows))\n    fig.suptitle('Sample Images Per Hemorrhage Subtype (Positive Examples)', fontsize=16, y=1.02)\n    \n    for i, cls in enumerate(classes):\n        positive_samples = df[df[cls] == 1]\n        samples = positive_samples.head(n_cols)\n        image_ids = samples['image'].tolist()\n\n        for j in range(n_cols):\n            ax = axes[i, j]\n            \n            if j < len(image_ids):\n                img_id = image_ids[j]\n                png_path = os.path.join(root_dir, img_id + '.png')\n                \n                try:\n                    img = Image.open(png_path).convert('RGB')\n                    ax.imshow(img)\n                    ax.set_title(f\"ID: {img_id}\", fontsize=10)\n                except FileNotFoundError:\n                    ax.text(0.5, 0.5, \"PNG Not Found\", ha='center', va='center', transform=ax.transAxes)\n                    ax.set_title(f\"Missing: {img_id}\", fontsize=10)\n            else:\n                ax.axis('off')\n\n            ax.set_xticks([])\n            ax.set_yticks([])\n\n        axes[i, 0].set_ylabel(cls.upper(), fontsize=12, rotation=90, labelpad=20)\n\n    plt.tight_layout()\n    plt.savefig(os.path.join(OUTPUT_DIR, 'sample_images_per_class.png'))\n    plt.close(fig)\n    print(f\"Saved sample image grid to {os.path.join(OUTPUT_DIR, 'sample_images_per_class.png')}\")\n\ndf = pd.read_csv(CSV_PATH)\ndf[['image','subtype']] = df['ID'].str.rsplit('_', n=1, expand=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T13:51:41.319977Z","iopub.execute_input":"2025-12-16T13:51:41.320148Z","iopub.status.idle":"2025-12-16T13:52:05.835485Z","shell.execute_reply.started":"2025-12-16T13:51:41.320132Z","shell.execute_reply":"2025-12-16T13:52:05.834658Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T13:52:12.641228Z","iopub.execute_input":"2025-12-16T13:52:12.641853Z","iopub.status.idle":"2025-12-16T13:52:12.669919Z","shell.execute_reply.started":"2025-12-16T13:52:12.641831Z","shell.execute_reply":"2025-12-16T13:52:12.669242Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport seaborn as sns\n\nplt.figure(figsize=(10, 6))\nsns.countplot(\n    data=df,\n    x=\"subtype\",\n    order=df[\"subtype\"].value_counts().index\n)\n\nplt.title(\"Count Plot of Hemorrhage Subtypes\")\nplt.xlabel(\"Subtype\")\nplt.ylabel(\"Count\")\nplt.xticks(rotation=45)\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T13:52:15.346627Z","iopub.execute_input":"2025-12-16T13:52:15.346907Z","iopub.status.idle":"2025-12-16T13:52:19.089485Z","shell.execute_reply.started":"2025-12-16T13:52:15.346886Z","shell.execute_reply":"2025-12-16T13:52:19.088793Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"subtype_counts = df[\"subtype\"].value_counts()\n\nplt.figure(figsize=(8, 8))\nplt.pie(\n    subtype_counts.values,\n    labels=subtype_counts.index,\n    autopct=\"%1.1f%%\",\n    startangle=90\n)\n\nplt.title(\"Pie Chart of Hemorrhage Subtypes\")\nplt.axis(\"equal\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T13:52:26.410147Z","iopub.execute_input":"2025-12-16T13:52:26.410756Z","iopub.status.idle":"2025-12-16T13:52:26.855279Z","shell.execute_reply.started":"2025-12-16T13:52:26.410731Z","shell.execute_reply":"2025-12-16T13:52:26.854607Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_ml = df.pivot_table(\n    index='image',\n    columns='subtype',\n    values='Label',\n    aggfunc='max',\n    fill_value=0\n).reset_index()\ndf_ml['any'] = df_ml[CLASS_NAMES].max(axis=1)\n\nframes = []\nfor cls in CLASS_NAMES:\n    pos = df_ml[df_ml[cls]==1]\n    neg = df_ml[df_ml[cls]==0]\n    \n    pos = pos.sample(BALANCE_N, random_state=42) if len(pos)>BALANCE_N else pos\n    neg = neg.sample(BALANCE_N, random_state=42) if len(neg)>BALANCE_N else neg\n    \n    frames.append(pos)\n    frames.append(neg)\n\nbalanced_df = pd.concat(frames).drop_duplicates(subset='image').reset_index(drop=True)\nprint(f\"Balanced subset size: {len(balanced_df)} images\")\n\nbalanced_csv = os.path.join(METRICS_DIR, 'subset_multilabel_balanced.csv')\nbalanced_df.to_csv(balanced_csv, index=False)\n\nprint(\"Converting balanced subset DICOM to PNG...\")\nfor img_id in tqdm(balanced_df['image'].unique(), desc=\"DICOM→PNG\"):\n    dcm_path = os.path.join(DCM_DIR, img_id + '.dcm')\n    png_path = os.path.join(PNG_ROOT, img_id + '.png')\n    \n    if os.path.exists(png_path) or not os.path.exists(dcm_path):\n        continue\n        \n    ds = pydicom.dcmread(dcm_path)\n    arr = ds.pixel_array.astype(np.float32)\n    arr = (arr - arr.min()) / (arr.max() - arr.min()) * 255.0\n    arr = arr.astype(np.uint8)\n    Image.fromarray(arr).save(png_path)\n\nprint(\"Done! Balanced subset PNGs are in:\", PNG_ROOT)\n\ndisplay_sample_images_per_class(balanced_df, PNG_ROOT, CLASS_NAMES, num_samples=5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T13:53:28.260716Z","iopub.execute_input":"2025-12-16T13:53:28.261311Z","iopub.status.idle":"2025-12-16T13:57:19.808662Z","shell.execute_reply.started":"2025-12-16T13:53:28.261288Z","shell.execute_reply":"2025-12-16T13:57:19.807977Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class MultiHemoDataset(Dataset):\n    def __init__(self, df, root_dir, transform=None):\n        self.df = df.copy().reset_index(drop=True) \n        self.root = root_dir\n        self.tf = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        rec = self.df.iloc[idx]\n        \n        path = os.path.join(self.root, rec['image'] + '.png')\n        img = Image.open(path).convert('RGB')\n        if self.tf:\n            img = self.tf(img)\n            \n        label_arr = rec[CLASS_NAMES].astype(np.float32).to_numpy()\n        labels = torch.from_numpy(label_arr)\n        \n        return img, labels\n\ntrain_transform = transforms.Compose([\n    transforms.Resize(256),\n    transforms.RandomResizedCrop(224),\n    transforms.RandomHorizontalFlip(),\n    transforms.ToTensor(),\n    transforms.Normalize([0.485,0.456,0.406], [0.229,0.224,0.225])\n])\nval_transform = transforms.Compose([\n    transforms.Resize(256),\n    transforms.CenterCrop(224),\n    transforms.ToTensor(),\n    transforms.Normalize([0.485,0.456,0.406], [0.229,0.224,0.225])\n])\n\ntrain_df, val_df = train_test_split(\n    balanced_df, \n    test_size=VAL_SPLIT_SIZE, \n    random_state=42, \n    stratify=balanced_df['any']\n)\n\ntrain_loader = DataLoader(\n    MultiHemoDataset(train_df, PNG_ROOT, train_transform),\n    batch_size=BATCH_SIZE, shuffle=True,  num_workers=0)\nval_loader   = DataLoader(\n    MultiHemoDataset(val_df,   PNG_ROOT, val_transform),\n    batch_size=BATCH_SIZE, shuffle=False, num_workers=0)\n\nprint(f\"Training images: {len(train_df)}, Validation images: {len(val_df)}\")\n\nmodel     = models.resnet18(weights=models.ResNet18_Weights.IMAGENET1K_V1)\nmodel.fc  = nn.Linear(model.fc.in_features, len(CLASS_NAMES))\nmodel     = model.to(DEVICE)\ncriterion = nn.BCEWithLogitsLoss()\noptimizer = torch.optim.Adam(model.parameters(),\n                             lr=LEARNING_RATE,\n                             weight_decay=WEIGHT_DECAY)\n\nbest_val_loss    = np.inf\nno_improve       = 0\ntrain_losses, train_accs = [], []\nval_losses,   val_accs   = [], []\nbest_model_path  = os.path.join(OUTPUT_DIR, 'best_model.pth')\n\nprint(f\"\\n=== Starting Training (Max {NUM_EPOCHS} Epochs) ===\")\nfor epoch in range(1, NUM_EPOCHS+1):\n    \n    model.train()\n    running_loss = running_corr = running_total = 0\n    for imgs, labels in train_loader:\n        imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n        optimizer.zero_grad()\n        outputs = model(imgs)\n        loss    = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n\n        running_loss  += loss.item() * imgs.size(0)\n        preds         = (torch.sigmoid(outputs) >= 0.5).long()\n        running_corr  += (preds == labels).all(dim=1).sum().item()\n        running_total += imgs.size(0)\n\n    epoch_train_loss = running_loss / running_total\n    epoch_train_acc  = running_corr  / running_total\n    train_losses.append(epoch_train_loss)\n    train_accs.append(epoch_train_acc)\n\n    model.eval()\n    val_running_loss = val_running_corr = val_running_total = 0\n    with torch.no_grad():\n        for imgs, labels in val_loader:\n            imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n            outputs = model(imgs)\n            loss    = criterion(outputs, labels)\n\n            val_running_loss  += loss.item() * imgs.size(0)\n            preds              = (torch.sigmoid(outputs) >= 0.5).long()\n            val_running_corr  += (preds == labels).all(dim=1).sum().item()\n            val_running_total += imgs.size(0)\n\n    epoch_val_loss = val_running_loss / val_running_total\n    epoch_val_acc  = val_running_corr  / val_running_total\n    val_losses.append(epoch_val_loss)\n    val_accs.append(epoch_val_acc)\n\n    print(f\"Epoch {epoch}/{NUM_EPOCHS} — \"\n          f\"Train loss: {epoch_train_loss:.4f}, acc: {epoch_train_acc:.4f} | \"\n          f\" Val loss: {epoch_val_loss:.4f}, acc: {epoch_val_acc:.4f}\")\n\n    if epoch_val_loss < best_val_loss:\n        best_val_loss = epoch_val_loss\n        no_improve    = 0\n        torch.save(model.state_dict(), best_model_path)\n    else:\n        no_improve += 1\n        if no_improve >= PATIENCE and NUM_EPOCHS > PATIENCE:\n            print(f\"Early stopping at epoch {epoch}\")\n            break\n\nprint(\"\\n=== Final Evaluation and Metrics ===\")\nmodel.load_state_dict(torch.load(best_model_path))\nmodel.eval()\n\nval_probs_list, val_targets_list = [], []\nwith torch.no_grad():\n    for imgs, labels in tqdm(val_loader, desc=\"Final Evaluation\"):\n        imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n        outputs = model(imgs)\n        val_probs_list.extend(torch.sigmoid(outputs).cpu().numpy())\n        val_targets_list.extend(labels.cpu().numpy())\n        \nval_probs   = np.array(val_probs_list)\nval_targets = np.array(val_targets_list)\nval_preds   = (val_probs >= 0.5).astype(int)\n\nprint(\"\\n--- Multi-Label Classification Report (Threshold=0.5) ---\")\nreport = classification_report(\n    val_targets, val_preds, \n    target_names=CLASS_NAMES, \n    zero_division=0, \n    output_dict=True\n)\nreport_df = pd.DataFrame(report).transpose().round(4)\nprint(report_df.to_markdown())\nreport_df.to_csv(os.path.join(METRICS_DIR, 'classification_report.csv'))\n\ndef plot_confusion_matrix(cm, classes, title, cmap=plt.cm.Blues):\n    plt.imshow(cm, interpolation='nearest', cmap=cmap)\n    plt.title(title)\n    plt.colorbar()\n    tick_marks = np.arange(len(classes))\n    plt.xticks(tick_marks, classes, rotation=45)\n    plt.yticks(tick_marks, classes)\n\n    fmt = 'd'\n    thresh = cm.max() / 2.\n    for i, j in itertools.product(range(cm.shape[0]), range(cm.shape[1])):\n        plt.text(j, i, format(cm[i, j], fmt),\n                 horizontalalignment=\"center\",\n                 color=\"white\" if cm[i, j] > thresh else \"black\")\n\n    plt.ylabel('True label')\n    plt.xlabel('Predicted label')\n    plt.tight_layout()\n    plt.savefig(os.path.join(METRICS_DIR, f'cm_{title.replace(\" \", \"_\")}.png'))\n    plt.close()\n\nfor i, cls in enumerate(CLASS_NAMES):\n    cm = confusion_matrix(val_targets[:, i], val_preds[:, i])\n    plot_confusion_matrix(\n        cm, \n        classes=['Negative', 'Positive'], \n        title=f'Confusion Matrix: {cls}'\n    )\n    \nprint(f\"\\nSaved {len(CLASS_NAMES)} Confusion Matrix plots to: {METRICS_DIR}\")\n\nplt.figure(figsize=(8,6))\nmetrics_list = []\nfor i, cls in enumerate(CLASS_NAMES):\n    fpr, tpr, _ = roc_curve(val_targets[:, i], val_probs[:, i])\n    roc_auc = auc(fpr, tpr)\n    \n    plt.plot(fpr, tpr, label=f\"{cls} (AUC = {roc_auc:.2f})\")\n    metrics_list.append({\n        'class': cls,\n        'AUC': roc_auc,\n        'Precision': report_df.loc[cls, 'precision'],\n        'Recall': report_df.loc[cls, 'recall'],\n        'F1-Score': report_df.loc[cls, 'f1-score']\n    })\n\nplt.plot([0,1],[0,1],'--', color='gray', label='Random Guess')\nplt.xlabel('False Positive Rate (FPR)'); plt.ylabel('True Positive Rate (TPR)')\nplt.title('Receiver Operating Characteristic (ROC) Curve')\nplt.legend()\nplt.savefig(os.path.join(ROC_DIR, 'roc_all_classes.png'))\nplt.close()\n\npd.DataFrame(metrics_list).to_csv(os.path.join(METRICS_DIR, 'final_metrics_summary.csv'), index=False)\nprint(f\"Saved ROC curve plot to: {ROC_DIR}\")\n\nprint(f\"\\nAll results (metrics, CM, ROC) saved to: {OUTPUT_DIR}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T13:59:12.557535Z","iopub.execute_input":"2025-12-16T13:59:12.557838Z","iopub.status.idle":"2025-12-16T14:03:05.812143Z","shell.execute_reply.started":"2025-12-16T13:59:12.55781Z","shell.execute_reply":"2025-12-16T14:03:05.811532Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class MultiHemoDataset(Dataset):\n    def __init__(self, df, root_dir, transform=None):\n        self.df = df.copy().reset_index(drop=True) \n        self.root = root_dir\n        self.tf = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        rec = self.df.iloc[idx]\n        \n        path = os.path.join(self.root, rec['image'] + '.png')\n        img = Image.open(path).convert('RGB')\n        if self.tf:\n            img = self.tf(img)\n            \n        label_arr = rec[CLASS_NAMES].astype(np.float32).to_numpy()\n        labels = torch.from_numpy(label_arr)\n        \n        return img, labels\n\ntrain_transform = transforms.Compose([\n    transforms.Resize(340), \n    transforms.RandomResizedCrop(299),\n    transforms.RandomHorizontalFlip(),\n    transforms.ToTensor(),\n    transforms.Normalize([0.485,0.456,0.406], [0.229,0.224,0.225])\n])\nval_transform = transforms.Compose([\n    transforms.Resize(299), \n    transforms.CenterCrop(299),\n    transforms.ToTensor(),\n    transforms.Normalize([0.485,0.456,0.406], [0.229,0.224,0.225])\n])\n\ntrain_df, val_df = train_test_split(\n    balanced_df, \n    test_size=VAL_SPLIT_SIZE, \n    random_state=42, \n    stratify=balanced_df['any']\n)\n\ntrain_loader = DataLoader(\n    MultiHemoDataset(train_df, PNG_ROOT, train_transform),\n    batch_size=BATCH_SIZE, shuffle=True,  num_workers=0)\nval_loader   = DataLoader(\n    MultiHemoDataset(val_df,   PNG_ROOT, val_transform),\n    batch_size=BATCH_SIZE, shuffle=False, num_workers=0)\n\nprint(f\"Training images: {len(train_df)}, Validation images: {len(val_df)}\")\n\nmodel = models.inception_v3(weights=models.Inception_V3_Weights.IMAGENET1K_V1, aux_logits=True, transform_input=True)\n\nmodel.AuxLogits = None\nmodel.aux_logits = False\n\nnum_ftrs = model.fc.in_features\nmodel.fc = nn.Linear(num_ftrs, len(CLASS_NAMES))\n\nmodel = model.to(DEVICE)\n\ncriterion = nn.BCEWithLogitsLoss()\noptimizer = torch.optim.Adam(model.parameters(),\n                             lr=LEARNING_RATE,\n                             weight_decay=WEIGHT_DECAY)\n\nbest_val_loss    = np.inf\nno_improve       = 0\ntrain_losses, train_accs = [], []\nval_losses,   val_accs   = [], []\nbest_model_path  = os.path.join(OUTPUT_DIR, 'best_model.pth')\n\nprint(f\"\\n=== Starting Training (Max {NUM_EPOCHS} Epochs) ===\")\nfor epoch in range(1, NUM_EPOCHS+1):\n    \n    model.train()\n    running_loss = running_corr = running_total = 0\n    for imgs, labels in train_loader:\n        imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n        optimizer.zero_grad()\n        \n        outputs = model(imgs) \n        \n        loss = criterion(outputs, labels)\n\n        loss.backward()\n        optimizer.step()\n\n        running_loss  += loss.item() * imgs.size(0)\n        preds         = (torch.sigmoid(outputs) >= 0.5).long()\n        running_corr  += (preds == labels).all(dim=1).sum().item()\n        running_total += imgs.size(0)\n\n    epoch_train_loss = running_loss / running_total\n    epoch_train_acc  = running_corr  / running_total\n    train_losses.append(epoch_train_loss)\n    train_accs.append(epoch_train_acc)\n\n    model.eval()\n    val_running_loss = val_running_corr = val_running_total = 0\n    with torch.no_grad():\n        for imgs, labels in val_loader:\n            imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n            outputs = model(imgs)\n\n            val_running_loss  += criterion(outputs, labels).item() * imgs.size(0)\n            preds              = (torch.sigmoid(outputs) >= 0.5).long()\n            val_running_corr  += (preds == labels).all(dim=1).sum().item()\n            val_running_total += imgs.size(0)\n\n    epoch_val_loss = val_running_loss / val_running_total\n    epoch_val_acc  = val_running_corr  / val_running_total\n    val_losses.append(epoch_val_loss)\n    val_accs.append(epoch_val_acc)\n\n    print(f\"Epoch {epoch}/{NUM_EPOCHS} — \"\n          f\"Train loss: {epoch_train_loss:.4f}, acc: {epoch_train_acc:.4f} | \"\n          f\" Val loss: {epoch_val_loss:.4f}, acc: {epoch_val_acc:.4f}\")\n\n    if epoch_val_loss < best_val_loss:\n        best_val_loss = epoch_val_loss\n        no_improve    = 0\n        torch.save(model.state_dict(), best_model_path)\n    else:\n        no_improve += 1\n        if no_improve >= PATIENCE and NUM_EPOCHS > PATIENCE:\n            print(f\"Early stopping at epoch {epoch}\")\n            break\n\nprint(\"\\n=== Final Evaluation and Metrics ===\")\nmodel.load_state_dict(torch.load(best_model_path))\nmodel.eval()\n\nval_probs_list, val_targets_list = [], []\nwith torch.no_grad():\n    for imgs, labels in tqdm(val_loader, desc=\"Final Evaluation\"):\n        imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n        outputs = model(imgs)\n        val_probs_list.extend(torch.sigmoid(outputs).cpu().numpy())\n        val_targets_list.extend(labels.cpu().numpy())\n        \nval_probs   = np.array(val_probs_list)\nval_targets = np.array(val_targets_list)\nval_preds   = (val_probs >= 0.5).astype(int)\n\nprint(\"\\n--- Multi-Label Classification Report (Threshold=0.5) ---\")\nreport = classification_report(\n    val_targets, val_preds, \n    target_names=CLASS_NAMES, \n    zero_division=0, \n    output_dict=True\n)\nreport_df = pd.DataFrame(report).transpose().round(4)\nprint(report_df.to_markdown())\nreport_df.to_csv(os.path.join(METRICS_DIR, 'classification_report.csv'))\n\ndef plot_confusion_matrix(cm, classes, title, cmap=plt.cm.Blues):\n    plt.imshow(cm, interpolation='nearest', cmap=cmap)\n    plt.title(title)\n    plt.colorbar()\n    tick_marks = np.arange(len(classes))\n    plt.xticks(tick_marks, classes, rotation=45)\n    plt.yticks(tick_marks, classes)\n\n    fmt = 'd'\n    thresh = cm.max() / 2.\n    for i, j in itertools.product(range(cm.shape[0]), range(cm.shape[1])):\n        plt.text(j, i, format(cm[i, j], fmt),\n                 horizontalalignment=\"center\",\n                 color=\"white\" if cm[i, j] > thresh else \"black\")\n\n    plt.ylabel('True label')\n    plt.xlabel('Predicted label')\n    plt.tight_layout()\n    plt.savefig(os.path.join(METRICS_DIR, f'cm_{title.replace(\" \", \"_\")}.png'))\n    plt.close()\n\nfor i, cls in enumerate(CLASS_NAMES):\n    cm = confusion_matrix(val_targets[:, i], val_preds[:, i])\n    plot_confusion_matrix(\n        cm, \n        classes=['Negative', 'Positive'], \n        title=f'Confusion Matrix: {cls}'\n    )\n    \nprint(f\"\\nSaved {len(CLASS_NAMES)} Confusion Matrix plots to: {METRICS_DIR}\")\n\nplt.figure(figsize=(8,6))\nmetrics_list = []\nfor i, cls in enumerate(CLASS_NAMES):\n    fpr, tpr, _ = roc_curve(val_targets[:, i], val_probs[:, i])\n    roc_auc = auc(fpr, tpr)\n    \n    plt.plot(fpr, tpr, label=f\"{cls} (AUC = {roc_auc:.2f})\")\n    metrics_list.append({\n        'class': cls,\n        'AUC': roc_auc,\n        'Precision': report_df.loc[cls, 'precision'],\n        'Recall': report_df.loc[cls, 'recall'],\n        'F1-Score': report_df.loc[cls, 'f1-score']\n    })\n\nplt.plot([0,1],[0,1],'--', color='gray', label='Random Guess')\nplt.xlabel('False Positive Rate (FPR)'); plt.ylabel('True Positive Rate (TPR)')\nplt.title('Receiver Operating Characteristic (ROC) Curve')\nplt.legend()\nplt.savefig(os.path.join(ROC_DIR, 'roc_all_classes.png'))\nplt.close()\n\npd.DataFrame(metrics_list).to_csv(os.path.join(METRICS_DIR, 'final_metrics_summary.csv'), index=False)\nprint(f\"Saved ROC curve plot to: {ROC_DIR}\")\n\nprint(f\"\\nAll results (metrics, CM, ROC) saved to: {OUTPUT_DIR}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T14:11:16.318462Z","iopub.execute_input":"2025-12-16T14:11:16.319023Z","iopub.status.idle":"2025-12-16T14:18:58.916762Z","shell.execute_reply.started":"2025-12-16T14:11:16.319001Z","shell.execute_reply":"2025-12-16T14:18:58.916131Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class MultiHemoDataset(Dataset):\n    def __init__(self, df, root_dir, transform=None):\n        self.df = df.copy().reset_index(drop=True) \n        self.root = root_dir\n        self.tf = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        rec = self.df.iloc[idx]\n        \n        path = os.path.join(self.root, rec['image'] + '.png')\n        img = Image.open(path).convert('RGB')\n        if self.tf:\n            img = self.tf(img)\n            \n        label_arr = rec[CLASS_NAMES].astype(np.float32).to_numpy()\n        labels = torch.from_numpy(label_arr)\n        \n        return img, labels\n\ntrain_transform = transforms.Compose([\n    transforms.Resize(340), \n    transforms.RandomResizedCrop(299),\n    transforms.RandomHorizontalFlip(),\n    transforms.ToTensor(),\n    transforms.Normalize([0.485,0.456,0.406], [0.229,0.224,0.225])\n])\nval_transform = transforms.Compose([\n    transforms.Resize(299), \n    transforms.CenterCrop(299),\n    transforms.ToTensor(),\n    transforms.Normalize([0.485,0.456,0.406], [0.229,0.224,0.225])\n])\n\ntrain_df, val_df = train_test_split(\n    balanced_df, \n    test_size=VAL_SPLIT_SIZE, \n    random_state=42, \n    stratify=balanced_df['any']\n)\n\nclass_counts = train_df[CLASS_NAMES].sum()\ntotal_samples = len(train_df)\n\nweights = (total_samples - class_counts) / (class_counts + 1e-6)\nweights = weights / weights.mean() \nclass_weights = torch.tensor(weights.values, dtype=torch.float).to(DEVICE)\n\ntrain_loader = DataLoader(\n    MultiHemoDataset(train_df, PNG_ROOT, train_transform),\n    batch_size=BATCH_SIZE, shuffle=True,  num_workers=0)\nval_loader   = DataLoader(\n    MultiHemoDataset(val_df,   PNG_ROOT, val_transform),\n    batch_size=BATCH_SIZE, shuffle=False, num_workers=0)\n\nprint(f\"Training images: {len(train_df)}, Validation images: {len(val_df)}\")\n\nmodel = models.inception_v3(weights=models.Inception_V3_Weights.IMAGENET1K_V1, aux_logits=True, transform_input=True)\n\nmodel.AuxLogits = None\nmodel.aux_logits = False\n\nnum_ftrs = model.fc.in_features\n\nmodel.fc = nn.Linear(num_ftrs, len(CLASS_NAMES))\n\nmodel = model.to(DEVICE)\n\ncriterion = nn.BCEWithLogitsLoss(weight=class_weights)\noptimizer = torch.optim.Adam(model.parameters(),\n                             lr=LEARNING_RATE,\n                             weight_decay=WEIGHT_DECAY)\n\nbest_val_loss    = np.inf\nno_improve       = 0\ntrain_losses, train_accs = [], []\nval_losses,   val_accs   = [], []\nbest_model_path  = os.path.join(OUTPUT_DIR, 'best_model.pth')\n\nprint(f\"\\n=== Starting Training (Max {NUM_EPOCHS} Epochs) ===\")\nfor epoch in range(1, NUM_EPOCHS+1):\n    \n    model.train()\n    running_loss = running_corr = running_total = 0\n    for imgs, labels in train_loader:\n        imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n        optimizer.zero_grad()\n        \n        outputs = model(imgs) \n        \n        loss = criterion(outputs, labels)\n\n        loss.backward()\n        optimizer.step()\n\n        running_loss  += loss.item() * imgs.size(0)\n       \n        preds         = (torch.sigmoid(outputs) >= 0.5).long()\n        running_corr  += (preds == labels).all(dim=1).sum().item()\n        running_total += imgs.size(0)\n\n    epoch_train_loss = running_loss / running_total\n    epoch_train_acc  = running_corr  / running_total\n    train_losses.append(epoch_train_loss)\n    train_accs.append(epoch_train_acc)\n\n    model.eval()\n    val_running_loss = val_running_corr = val_running_total = 0\n    with torch.no_grad():\n        for imgs, labels in val_loader:\n            imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n            outputs = model(imgs)\n\n            unweighted_criterion = nn.BCEWithLogitsLoss() \n            val_running_loss  += unweighted_criterion(outputs, labels).item() * imgs.size(0)\n            preds              = (torch.sigmoid(outputs) >= 0.5).long()\n            val_running_corr  += (preds == labels).all(dim=1).sum().item()\n            val_running_total += imgs.size(0)\n\n    epoch_val_loss = val_running_loss / val_running_total\n    epoch_val_acc  = val_running_corr  / val_running_total\n    val_losses.append(epoch_val_loss)\n    val_accs.append(epoch_val_acc)\n\n    print(f\"Epoch {epoch}/{NUM_EPOCHS} — \"\n          f\"Train loss: {epoch_train_loss:.4f}, acc: {epoch_train_acc:.4f} | \"\n          f\" Val loss: {epoch_val_loss:.4f}, acc: {epoch_val_acc:.4f}\")\n\n    if epoch_val_loss < best_val_loss:\n        best_val_loss = epoch_val_loss\n        no_improve    = 0\n        torch.save(model.state_dict(), best_model_path)\n    else:\n        no_improve += 1\n        if no_improve >= PATIENCE and NUM_EPOCHS > PATIENCE:\n            print(f\"Early stopping at epoch {epoch}\")\n            break\n\nprint(\"\\n=== Final Evaluation and Metrics ===\")\n\nmodel.load_state_dict(torch.load(best_model_path))\nmodel.eval()\n\nval_probs_list, val_targets_list = [], []\nwith torch.no_grad():\n    for imgs, labels in tqdm(val_loader, desc=\"Final Evaluation\"):\n        imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n        outputs = model(imgs)\n        val_probs_list.extend(torch.sigmoid(outputs).cpu().numpy())\n        val_targets_list.extend(labels.cpu().numpy())\n        \nval_probs   = np.array(val_probs_list)\nval_targets = np.array(val_targets_list)\nval_preds   = (val_probs >= 0.5).astype(int)\n\nprint(\"\\n--- Multi-Label Classification Report (Threshold=0.5) ---\")\n\nreport = classification_report(\n    val_targets, val_preds, \n    target_names=CLASS_NAMES, \n    zero_division=0, \n    output_dict=True\n)\nreport_df = pd.DataFrame(report).transpose().round(4)\nprint(report_df.to_markdown())\nreport_df.to_csv(os.path.join(METRICS_DIR, 'classification_report.csv'))\n\ndef plot_confusion_matrix(cm, classes, title, cmap=plt.cm.Blues):\n    plt.imshow(cm, interpolation='nearest', cmap=cmap)\n    plt.title(title)\n    plt.colorbar()\n    tick_marks = np.arange(len(classes))\n    plt.xticks(tick_marks, classes, rotation=45)\n    plt.yticks(tick_marks, classes)\n\n    fmt = 'd'\n    thresh = cm.max() / 2.\n    for i, j in itertools.product(range(cm.shape[0]), range(cm.shape[1])):\n        plt.text(j, i, format(cm[i, j], fmt),\n                 horizontalalignment=\"center\",\n                 color=\"white\" if cm[i, j] > thresh else \"black\")\n\n    plt.ylabel('True label')\n    plt.xlabel('Predicted label')\n    plt.tight_layout()\n    plt.savefig(os.path.join(METRICS_DIR, f'cm_{title.replace(\" \", \"_\")}.png'))\n    plt.close()\n\nfor i, cls in enumerate(CLASS_NAMES):\n    cm = confusion_matrix(val_targets[:, i], val_preds[:, i])\n    plot_confusion_matrix(\n        cm, \n        classes=['Negative', 'Positive'], \n        title=f'Confusion Matrix: {cls}'\n    )\n    \nprint(f\"\\nSaved {len(CLASS_NAMES)} Confusion Matrix plots to: {METRICS_DIR}\")\n\nplt.figure(figsize=(8,6))\nmetrics_list = []\nfor i, cls in enumerate(CLASS_NAMES):\n    fpr, tpr, _ = roc_curve(val_targets[:, i], val_probs[:, i])\n    roc_auc = auc(fpr, tpr)\n    \n    plt.plot(fpr, tpr, label=f\"{cls} (AUC = {roc_auc:.2f})\")\n    metrics_list.append({\n        'class': cls,\n        'AUC': roc_auc,\n        'Precision': report_df.loc[cls, 'precision'],\n        'Recall': report_df.loc[cls, 'recall'],\n        'F1-Score': report_df.loc[cls, 'f1-score']\n    })\n\nplt.plot([0,1],[0,1],'--', color='gray', label='Random Guess')\nplt.xlabel('False Positive Rate (FPR)'); plt.ylabel('True Positive Rate (TPR)')\nplt.title('Receiver Operating Characteristic (ROC) Curve')\nplt.legend()\nplt.savefig(os.path.join(ROC_DIR, 'roc_all_classes.png'))\nplt.close()\n\npd.DataFrame(metrics_list).to_csv(os.path.join(METRICS_DIR, 'final_metrics_summary.csv'), index=False)\nprint(f\"Saved ROC curve plot to: {ROC_DIR}\")\n\nprint(f\"\\nAll results (metrics, CM, ROC) saved to: {OUTPUT_DIR}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T14:24:25.870488Z","iopub.execute_input":"2025-12-16T14:24:25.870758Z","iopub.status.idle":"2025-12-16T14:32:02.730036Z","shell.execute_reply.started":"2025-12-16T14:24:25.870741Z","shell.execute_reply":"2025-12-16T14:32:02.729357Z"}},"outputs":[],"execution_count":null}]}