{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","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":99552,"databundleVersionId":13851420},{"sourceType":"datasetVersion","sourceId":15802569,"datasetId":10117562,"databundleVersionId":16749854},{"sourceType":"datasetVersion","sourceId":15879874,"datasetId":10181565,"databundleVersionId":16833317}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport random\nimport warnings\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\nfrom sklearn.model_selection import GroupShuffleSplit\nfrom sklearn.metrics import accuracy_score, f1_score, roc_auc_score\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import models\n\nwarnings.filterwarnings(\"ignore\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T07:50:09.828621Z","iopub.execute_input":"2026-04-24T07:50:09.829237Z","iopub.status.idle":"2026-04-24T07:50:09.83454Z","shell.execute_reply.started":"2026-04-24T07:50:09.829162Z","shell.execute_reply":"2026-04-24T07:50:09.833819Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"SEED = 42\n\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\ntorch.cuda.manual_seed_all(SEED)\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"DEVICE:\", DEVICE)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T07:50:09.836056Z","iopub.execute_input":"2026-04-24T07:50:09.83644Z","iopub.status.idle":"2026-04-24T07:50:09.852725Z","shell.execute_reply.started":"2026-04-24T07:50:09.836419Z","shell.execute_reply":"2026-04-24T07:50:09.852042Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"PROCESSED_PATH = \"/kaggle/input/datasets/sonamaggarwal/additional-dataset-of-rsna-aneurysm\"\nNPY_ROOT = \"/kaggle/input/datasets/sonamaggarwal/preprocess-of-rsna-aneurysm/precomputed_slices\"\n\nSAMPLES_CSV = os.path.join(PROCESSED_PATH, \"samples_df.csv\")\nSERIES_LENGTHS_CSV = os.path.join(PROCESSED_PATH, \"series_lengths.csv\")\n\nprint(\"SAMPLES_CSV exists:\", os.path.exists(SAMPLES_CSV))\nprint(\"SERIES_LENGTHS_CSV exists:\", os.path.exists(SERIES_LENGTHS_CSV))\nprint(\"NPY_ROOT exists:\", os.path.exists(NPY_ROOT))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T07:50:09.854336Z","iopub.execute_input":"2026-04-24T07:50:09.854707Z","iopub.status.idle":"2026-04-24T07:50:09.873615Z","shell.execute_reply.started":"2026-04-24T07:50:09.854683Z","shell.execute_reply":"2026-04-24T07:50:09.872981Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"samples_df = pd.read_csv(SAMPLES_CSV)\nseries_lengths_df = pd.read_csv(SERIES_LENGTHS_CSV)\n\nsamples_df[\"SeriesInstanceUID\"] = samples_df[\"SeriesInstanceUID\"].astype(str)\nseries_lengths_df[\"SeriesInstanceUID\"] = series_lengths_df[\"SeriesInstanceUID\"].astype(str)\n\nseries_lengths = dict(\n    zip(series_lengths_df[\"SeriesInstanceUID\"], series_lengths_df[\"num_slices\"])\n)\n\nprint(\"samples_df shape:\", samples_df.shape)\nprint(\"series_lengths_df shape:\", series_lengths_df.shape)\n\ndisplay(samples_df.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T07:50:09.87453Z","iopub.execute_input":"2026-04-24T07:50:09.874812Z","iopub.status.idle":"2026-04-24T07:50:09.941481Z","shell.execute_reply.started":"2026-04-24T07:50:09.874783Z","shell.execute_reply":"2026-04-24T07:50:09.940757Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"TARGET_COL = \"label_binary\"\n\nLOCATION_COLS = [\n    \"Left Infraclinoid Internal Carotid Artery\",\n    \"Right Infraclinoid Internal Carotid Artery\",\n    \"Left Supraclinoid Internal Carotid Artery\",\n    \"Right Supraclinoid Internal Carotid Artery\",\n    \"Left Middle Cerebral Artery\",\n    \"Right Middle Cerebral Artery\",\n    \"Anterior Communicating Artery\",\n    \"Left Anterior Cerebral Artery\",\n    \"Right Anterior Cerebral Artery\",\n    \"Left Posterior Communicating Artery\",\n    \"Right Posterior Communicating Artery\",\n    \"Basilar Tip\",\n    \"Other Posterior Circulation\",\n]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T07:50:09.943443Z","iopub.execute_input":"2026-04-24T07:50:09.943738Z","iopub.status.idle":"2026-04-24T07:50:09.947855Z","shell.execute_reply.started":"2026-04-24T07:50:09.943716Z","shell.execute_reply":"2026-04-24T07:50:09.947234Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"samples_df[\"slice_idx\"] = pd.to_numeric(samples_df[\"slice_idx\"], errors=\"coerce\")\nsamples_df[\"label_binary\"] = pd.to_numeric(samples_df[\"label_binary\"], errors=\"coerce\")\n\nsamples_df = samples_df.dropna(subset=[\"SeriesInstanceUID\", \"slice_idx\", \"label_binary\"]).copy()\nsamples_df[\"slice_idx\"] = samples_df[\"slice_idx\"].astype(int)\nsamples_df[\"label_binary\"] = samples_df[\"label_binary\"].astype(int)\n\nprint(samples_df.shape)\nprint(samples_df[\"label_binary\"].value_counts())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T07:50:09.948744Z","iopub.execute_input":"2026-04-24T07:50:09.949258Z","iopub.status.idle":"2026-04-24T07:50:09.972123Z","shell.execute_reply.started":"2026-04-24T07:50:09.949218Z","shell.execute_reply":"2026-04-24T07:50:09.971443Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.model_selection import GroupShuffleSplit\n\ngss = GroupShuffleSplit(n_splits=1, test_size=0.15, random_state=42)\n\ngroups = samples_df[\"SeriesInstanceUID\"].values\ny = samples_df[\"label_binary\"].values\n\ntrain_idx, temp_idx = next(gss.split(samples_df, y=y, groups=groups))\n\ntrain_df = samples_df.iloc[train_idx].reset_index(drop=True)\ntemp_df = samples_df.iloc[temp_idx].reset_index(drop=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T07:50:09.973354Z","iopub.execute_input":"2026-04-24T07:50:09.973661Z","iopub.status.idle":"2026-04-24T07:50:09.989435Z","shell.execute_reply.started":"2026-04-24T07:50:09.973612Z","shell.execute_reply":"2026-04-24T07:50:09.988692Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"gss2 = GroupShuffleSplit(n_splits=1, test_size=0.5, random_state=42)\n\ngroups_temp = temp_df[\"SeriesInstanceUID\"].values\ny_temp = temp_df[\"label_binary\"].values\n\nval_idx, test_idx = next(gss2.split(temp_df, y=y_temp, groups=groups_temp))\n\nvalid_df = temp_df.iloc[val_idx].reset_index(drop=True)\ntest_df = temp_df.iloc[test_idx].reset_index(drop=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T07:50:09.990434Z","iopub.execute_input":"2026-04-24T07:50:09.990755Z","iopub.status.idle":"2026-04-24T07:50:10.004952Z","shell.execute_reply.started":"2026-04-24T07:50:09.99073Z","shell.execute_reply":"2026-04-24T07:50:10.004181Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_series = set(train_df[\"SeriesInstanceUID\"])\nval_series = set(valid_df[\"SeriesInstanceUID\"])\ntest_series = set(test_df[\"SeriesInstanceUID\"])\n\nprint(\"Train ∩ Val:\", len(train_series & val_series))\nprint(\"Train ∩ Test:\", len(train_series & test_series))\nprint(\"Val ∩ Test:\", len(val_series & test_series))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T07:50:10.005925Z","iopub.execute_input":"2026-04-24T07:50:10.00623Z","iopub.status.idle":"2026-04-24T07:50:10.017729Z","shell.execute_reply.started":"2026-04-24T07:50:10.006197Z","shell.execute_reply":"2026-04-24T07:50:10.017101Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df.to_csv(\"/kaggle/working/train_split.csv\", index=False)\nvalid_df.to_csv(\"/kaggle/working/valid_split.csv\", index=False)\ntest_df.to_csv(\"/kaggle/working/test_split.csv\", index=False)\n\nprint(\"Saved train/valid/test split CSVs\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T07:50:10.018952Z","iopub.execute_input":"2026-04-24T07:50:10.019307Z","iopub.status.idle":"2026-04-24T07:50:10.112772Z","shell.execute_reply.started":"2026-04-24T07:50:10.019279Z","shell.execute_reply":"2026-04-24T07:50:10.112068Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_neighbor_indices(center_idx, n_slices, num_slices=5):\n    assert num_slices % 2 == 1\n    half = num_slices // 2\n\n    out = []\n    for offset in range(-half, half + 1):\n        idx = center_idx + offset\n        idx = max(0, min(idx, n_slices - 1))\n        out.append(idx)\n    return out","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T07:50:10.114894Z","iopub.execute_input":"2026-04-24T07:50:10.11526Z","iopub.status.idle":"2026-04-24T07:50:10.120681Z","shell.execute_reply.started":"2026-04-24T07:50:10.115228Z","shell.execute_reply":"2026-04-24T07:50:10.119988Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def npy_path_for_slice(series_uid, slice_idx):\n    return os.path.join(NPY_ROOT, str(series_uid), f\"{slice_idx}.npy\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T07:50:10.121516Z","iopub.execute_input":"2026-04-24T07:50:10.121746Z","iopub.status.idle":"2026-04-24T07:50:10.136531Z","shell.execute_reply.started":"2026-04-24T07:50:10.121725Z","shell.execute_reply":"2026-04-24T07:50:10.135854Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class AneurysmNPYDataset(Dataset):\n    def __init__(self, df, series_lengths, num_slices=5, include_multilabel=True):\n        self.df = df.reset_index(drop=True)\n        self.series_lengths = series_lengths\n        self.num_slices = num_slices\n        self.include_multilabel = include_multilabel\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n\n        series_uid = str(row[\"SeriesInstanceUID\"])\n        center_idx = int(row[\"slice_idx\"])\n        n_slices = int(self.series_lengths[series_uid])\n\n        neighbor_indices = get_neighbor_indices(center_idx, n_slices, self.num_slices)\n\n        imgs = []\n        for sidx in neighbor_indices:\n            path = npy_path_for_slice(series_uid, sidx)\n\n            if os.path.exists(path):\n                img = np.load(path).astype(np.float32)\n            else:\n                # safety fallback\n                img = np.zeros((256, 256), dtype=np.float32)\n\n            imgs.append(img)\n\n        stack = np.stack(imgs, axis=0).astype(np.float32)   # (C,H,W)\n\n        y_binary = np.float32(row[\"label_binary\"])\n        sample = {\n            \"image\": torch.tensor(stack, dtype=torch.float32),\n            \"binary\": torch.tensor(y_binary, dtype=torch.float32),\n            \"series_uid\": series_uid,\n            \"slice_idx\": center_idx,\n        }\n\n        if self.include_multilabel:\n            y_locations = row[LOCATION_COLS].values.astype(np.float32)\n            sample[\"locations\"] = torch.tensor(y_locations, dtype=torch.float32)\n\n        return sample","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T07:50:10.137371Z","iopub.execute_input":"2026-04-24T07:50:10.137605Z","iopub.status.idle":"2026-04-24T07:50:10.149992Z","shell.execute_reply.started":"2026-04-24T07:50:10.137586Z","shell.execute_reply":"2026-04-24T07:50:10.149467Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"BATCH_SIZE = 16\nNUM_SLICES = 5","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T07:50:10.150868Z","iopub.execute_input":"2026-04-24T07:50:10.151229Z","iopub.status.idle":"2026-04-24T07:50:10.162941Z","shell.execute_reply.started":"2026-04-24T07:50:10.151159Z","shell.execute_reply":"2026-04-24T07:50:10.162135Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_dataset = AneurysmNPYDataset(\n    train_df,\n    series_lengths=series_lengths,\n    num_slices=NUM_SLICES,\n    include_multilabel=True\n)\n\nvalid_dataset = AneurysmNPYDataset(\n    valid_df,\n    series_lengths=series_lengths,\n    num_slices=NUM_SLICES,\n    include_multilabel=True\n)\n\ntest_dataset = AneurysmNPYDataset(\n    test_df,\n    series_lengths=series_lengths,\n    num_slices=NUM_SLICES,\n    include_multilabel=True\n)\n\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=True,\n    num_workers=4,\n    pin_memory=torch.cuda.is_available()\n)\n\nvalid_loader = DataLoader(\n    valid_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=False,\n    num_workers=4,\n    pin_memory=torch.cuda.is_available()\n)\n\ntest_loader = DataLoader(\n    test_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=False,\n    num_workers=4,\n    pin_memory=torch.cuda.is_available()\n)\n\nprint(\"Train batches:\", len(train_loader))\nprint(\"Valid batches:\", len(valid_loader))\nprint(\"Test batches :\", len(test_loader))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T07:50:10.163856Z","iopub.execute_input":"2026-04-24T07:50:10.16413Z","iopub.status.idle":"2026-04-24T07:50:10.178385Z","shell.execute_reply.started":"2026-04-24T07:50:10.1641Z","shell.execute_reply":"2026-04-24T07:50:10.177611Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"batch = next(iter(train_loader))\n\nprint(\"Image shape:\", batch[\"image\"].shape)       # (B,C,H,W)\nprint(\"Binary shape:\", batch[\"binary\"].shape)     # (B,)\nprint(\"Locations shape:\", batch[\"locations\"].shape)  # (B,13)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T07:50:10.179429Z","iopub.execute_input":"2026-04-24T07:50:10.17999Z","iopub.status.idle":"2026-04-24T07:50:10.844606Z","shell.execute_reply.started":"2026-04-24T07:50:10.179968Z","shell.execute_reply":"2026-04-24T07:50:10.84384Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sample = train_dataset[0]\nstack = sample[\"image\"].numpy()\n\nfig, axes = plt.subplots(1, stack.shape[0], figsize=(3 * stack.shape[0], 3))\nfor i in range(stack.shape[0]):\n    axes[i].imshow(stack[i], cmap=\"gray\")\n    axes[i].set_title(f\"Slice {i}\")\n    axes[i].axis(\"off\")\nplt.tight_layout()\nplt.show()\n\nprint(\"Binary label:\", sample[\"binary\"].item())\nprint(\"Location labels:\", sample[\"locations\"].numpy())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T07:50:10.84659Z","iopub.execute_input":"2026-04-24T07:50:10.846935Z","iopub.status.idle":"2026-04-24T07:50:11.190503Z","shell.execute_reply.started":"2026-04-24T07:50:10.846909Z","shell.execute_reply":"2026-04-24T07:50:11.189789Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class AneurysmFastNet(nn.Module):\n    def __init__(self, in_channels=5, num_locations=13):\n        super().__init__()\n\n        self.backbone = models.resnet18(weights=models.ResNet18_Weights.DEFAULT)\n\n        old_conv = self.backbone.conv1\n        self.backbone.conv1 = nn.Conv2d(\n            in_channels=in_channels,\n            out_channels=old_conv.out_channels,\n            kernel_size=old_conv.kernel_size,\n            stride=old_conv.stride,\n            padding=old_conv.padding,\n            bias=False\n        )\n\n        with torch.no_grad():\n            if in_channels >= 3:\n                self.backbone.conv1.weight[:, :3] = old_conv.weight\n                mean_w = old_conv.weight.mean(dim=1, keepdim=True)\n                for c in range(3, in_channels):\n                    self.backbone.conv1.weight[:, c:c+1] = mean_w\n            else:\n                self.backbone.conv1.weight[:] = old_conv.weight[:, :in_channels]\n\n        feat_dim = self.backbone.fc.in_features\n        self.backbone.fc = nn.Identity()\n\n        self.binary_head = nn.Linear(feat_dim, 1)\n        self.location_head = nn.Linear(feat_dim, num_locations)\n\n    def forward(self, x):\n        feats = self.backbone(x)\n        binary_logits = self.binary_head(feats).squeeze(1)\n        location_logits = self.location_head(feats)\n\n        return {\n            \"binary_logits\": binary_logits,\n            \"location_logits\": location_logits\n        }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T07:50:11.191633Z","iopub.execute_input":"2026-04-24T07:50:11.191936Z","iopub.status.idle":"2026-04-24T07:50:11.199471Z","shell.execute_reply.started":"2026-04-24T07:50:11.191914Z","shell.execute_reply":"2026-04-24T07:50:11.198753Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"num_pos = train_df[\"label_binary\"].sum()\nnum_neg = len(train_df) - num_pos\n\npos_weight_binary = torch.tensor([num_neg / max(num_pos, 1)], dtype=torch.float32).to(DEVICE)\n\nlocation_pos = train_df[LOCATION_COLS].sum(axis=0).values\nlocation_neg = len(train_df) - location_pos\npos_weight_locations = torch.tensor(\n    location_neg / np.clip(location_pos, 1, None),\n    dtype=torch.float32\n).to(DEVICE)\n\nbinary_criterion = nn.BCEWithLogitsLoss(pos_weight=pos_weight_binary)\nlocation_criterion = nn.BCEWithLogitsLoss(pos_weight=pos_weight_locations)\n\nprint(\"Binary pos weight:\", pos_weight_binary.item())\nprint(\"Location pos weights shape:\", pos_weight_locations.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T07:50:11.200436Z","iopub.execute_input":"2026-04-24T07:50:11.200998Z","iopub.status.idle":"2026-04-24T07:50:11.217414Z","shell.execute_reply.started":"2026-04-24T07:50:11.200977Z","shell.execute_reply":"2026-04-24T07:50:11.21671Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = AneurysmFastNet(\n    in_channels=NUM_SLICES,\n    num_locations=len(LOCATION_COLS)\n).to(DEVICE)\n\noptimizer = torch.optim.Adam(model.parameters(), lr=1e-4, weight_decay=1e-5)\n\nprint(model.__class__.__name__)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T07:50:11.218429Z","iopub.execute_input":"2026-04-24T07:50:11.218783Z","iopub.status.idle":"2026-04-24T07:50:11.439376Z","shell.execute_reply.started":"2026-04-24T07:50:11.218753Z","shell.execute_reply":"2026-04-24T07:50:11.438269Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def compute_loss(outputs, batch, alpha=1.0, beta=1.0):\n    loss_binary = binary_criterion(outputs[\"binary_logits\"], batch[\"binary\"])\n    loss_locations = location_criterion(outputs[\"location_logits\"], batch[\"locations\"])\n    total_loss = alpha * loss_binary + beta * loss_locations\n\n    return total_loss, {\n        \"loss_binary\": loss_binary.item(),\n        \"loss_locations\": loss_locations.item()\n    }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T07:50:11.440593Z","iopub.execute_input":"2026-04-24T07:50:11.440945Z","iopub.status.idle":"2026-04-24T07:50:11.446143Z","shell.execute_reply.started":"2026-04-24T07:50:11.440915Z","shell.execute_reply":"2026-04-24T07:50:11.44527Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_one_epoch(model, loader, optimizer, device):\n    model.train()\n\n    total_loss = 0.0\n    total_bin = 0.0\n    total_loc = 0.0\n\n    for batch in loader:\n        images = batch[\"image\"].to(device)\n        binary = batch[\"binary\"].to(device)\n        locations = batch[\"locations\"].to(device)\n\n        optimizer.zero_grad()\n\n        outputs = model(images)\n        loss, loss_dict = compute_loss(outputs, {\n            \"binary\": binary,\n            \"locations\": locations\n        })\n\n        loss.backward()\n        optimizer.step()\n\n        total_loss += loss.item()\n        total_bin += loss_dict[\"loss_binary\"]\n        total_loc += loss_dict[\"loss_locations\"]\n\n    n = len(loader)\n    return {\n        \"loss\": total_loss / n,\n        \"loss_binary\": total_bin / n,\n        \"loss_locations\": total_loc / n\n    }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T07:50:11.447229Z","iopub.execute_input":"2026-04-24T07:50:11.447635Z","iopub.status.idle":"2026-04-24T07:50:11.462032Z","shell.execute_reply.started":"2026-04-24T07:50:11.447614Z","shell.execute_reply":"2026-04-24T07:50:11.461227Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@torch.no_grad()\ndef validate_one_epoch(model, loader, device):\n    model.eval()\n\n    total_loss = 0.0\n    total_bin = 0.0\n    total_loc = 0.0\n\n    y_true = []\n    y_prob = []\n    y_pred = []\n\n    loc_true = []\n    loc_prob = []\n\n    for batch in loader:\n        images = batch[\"image\"].to(device)\n        binary = batch[\"binary\"].to(device)\n        locations = batch[\"locations\"].to(device)\n\n        outputs = model(images)\n        loss, loss_dict = compute_loss(outputs, {\n            \"binary\": binary,\n            \"locations\": locations\n        })\n\n        total_loss += loss.item()\n        total_bin += loss_dict[\"loss_binary\"]\n        total_loc += loss_dict[\"loss_locations\"]\n\n        probs = torch.sigmoid(outputs[\"binary_logits\"]).cpu().numpy()\n        preds = (probs >= 0.5).astype(int)\n\n        y_true.extend(binary.cpu().numpy().astype(int))\n        y_prob.extend(probs)\n        y_pred.extend(preds)\n\n        lprob = torch.sigmoid(outputs[\"location_logits\"]).cpu().numpy()\n        loc_true.append(locations.cpu().numpy())\n        loc_prob.append(lprob)\n\n    y_true = np.array(y_true)\n    y_prob = np.array(y_prob)\n    y_pred = np.array(y_pred)\n\n    loc_true = np.concatenate(loc_true, axis=0)\n    loc_prob = np.concatenate(loc_prob, axis=0)\n    loc_pred = (loc_prob >= 0.5).astype(int)\n\n    metrics = {}\n    metrics[\"loss\"] = total_loss / len(loader)\n    metrics[\"loss_binary\"] = total_bin / len(loader)\n    metrics[\"loss_locations\"] = total_loc / len(loader)\n\n    metrics[\"binary_acc\"] = accuracy_score(y_true, y_pred)\n    metrics[\"binary_f1\"] = f1_score(y_true, y_pred)\n\n    try:\n        metrics[\"binary_auc\"] = roc_auc_score(y_true, y_prob)\n    except Exception:\n        metrics[\"binary_auc\"] = np.nan\n\n    metrics[\"loc_macro_f1\"] = f1_score(loc_true, loc_pred, average=\"macro\", zero_division=0)\n    metrics[\"loc_micro_f1\"] = f1_score(loc_true, loc_pred, average=\"micro\", zero_division=0)\n\n    return metrics","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T07:50:11.463006Z","iopub.execute_input":"2026-04-24T07:50:11.463627Z","iopub.status.idle":"2026-04-24T07:50:11.477259Z","shell.execute_reply.started":"2026-04-24T07:50:11.463605Z","shell.execute_reply":"2026-04-24T07:50:11.476508Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"EPOCHS = 5\nbest_auc = -np.inf\nbest_model_path = \"/kaggle/working/best_fast_model.pth\"\n\nhistory = []","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T07:50:11.478143Z","iopub.execute_input":"2026-04-24T07:50:11.478516Z","iopub.status.idle":"2026-04-24T07:50:11.493158Z","shell.execute_reply.started":"2026-04-24T07:50:11.478485Z","shell.execute_reply":"2026-04-24T07:50:11.492527Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for epoch in range(1, EPOCHS + 1):\n    train_metrics = train_one_epoch(model, train_loader, optimizer, DEVICE)\n    valid_metrics = validate_one_epoch(model, valid_loader, DEVICE)\n\n    row = {\n        \"epoch\": epoch,\n        **{f\"train_{k}\": v for k, v in train_metrics.items()},\n        **{f\"valid_{k}\": v for k, v in valid_metrics.items()},\n    }\n    history.append(row)\n\n    print(f\"\\nEpoch {epoch}/{EPOCHS}\")\n    print(\n        f\"Train Loss: {train_metrics['loss']:.4f} | \"\n        f\"Valid Loss: {valid_metrics['loss']:.4f} | \"\n        f\"Valid AUC: {valid_metrics['binary_auc']:.4f} | \"\n        f\"Valid F1: {valid_metrics['binary_f1']:.4f} | \"\n        f\"Loc Macro-F1: {valid_metrics['loc_macro_f1']:.4f}\"\n    )\n\n    current_auc = valid_metrics[\"binary_auc\"]\n    if not np.isnan(current_auc) and current_auc > best_auc:\n        best_auc = current_auc\n        torch.save(model.state_dict(), best_model_path)\n        print(\"Saved best model\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T07:50:11.495705Z","iopub.execute_input":"2026-04-24T07:50:11.495991Z","iopub.status.idle":"2026-04-24T07:52:18.758111Z","shell.execute_reply.started":"2026-04-24T07:50:11.495972Z","shell.execute_reply":"2026-04-24T07:52:18.757259Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"history_df = pd.DataFrame(history)\ndisplay(history_df)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T07:52:18.75938Z","iopub.execute_input":"2026-04-24T07:52:18.759731Z","iopub.status.idle":"2026-04-24T07:52:18.774534Z","shell.execute_reply.started":"2026-04-24T07:52:18.759703Z","shell.execute_reply":"2026-04-24T07:52:18.773733Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.figure(figsize=(8, 5))\nplt.plot(history_df[\"epoch\"], history_df[\"train_loss\"], label=\"Train Loss\")\nplt.plot(history_df[\"epoch\"], history_df[\"valid_loss\"], label=\"Valid Loss\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Loss\")\nplt.title(\"Training Curves\")\nplt.legend()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T07:52:18.775612Z","iopub.execute_input":"2026-04-24T07:52:18.775941Z","iopub.status.idle":"2026-04-24T07:52:18.945366Z","shell.execute_reply.started":"2026-04-24T07:52:18.775918Z","shell.execute_reply":"2026-04-24T07:52:18.944525Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"best_model = AneurysmFastNet(\n    in_channels=NUM_SLICES,\n    num_locations=len(LOCATION_COLS)\n).to(DEVICE)\n\nbest_model.load_state_dict(torch.load(best_model_path, map_location=DEVICE))\n\nvalid_metrics = validate_one_epoch(best_model, valid_loader, DEVICE)\ntest_metrics = validate_one_epoch(best_model, test_loader, DEVICE)\n\nprint(\"Best model validation metrics:\")\nfor k, v in valid_metrics.items():\n    print(f\"{k}: {v}\")\n\nprint(\"\\nFinal TEST metrics:\")\nfor k, v in test_metrics.items():\n    print(f\"{k}: {v}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T07:52:18.94657Z","iopub.execute_input":"2026-04-24T07:52:18.947129Z","iopub.status.idle":"2026-04-24T07:52:25.613161Z","shell.execute_reply.started":"2026-04-24T07:52:18.947105Z","shell.execute_reply":"2026-04-24T07:52:25.612113Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"history_df.to_csv(\"/kaggle/working/fast_training_history.csv\", index=False)\nprint(\"Saved fast_training_history.csv\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T07:52:25.614694Z","iopub.execute_input":"2026-04-24T07:52:25.61509Z","iopub.status.idle":"2026-04-24T07:52:25.62194Z","shell.execute_reply.started":"2026-04-24T07:52:25.61506Z","shell.execute_reply":"2026-04-24T07:52:25.621305Z"}},"outputs":[],"execution_count":null}]}