{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":99552,"databundleVersionId":13851420,"sourceType":"competition"},{"sourceId":608745,"sourceType":"modelInstanceVersion","modelInstanceId":456986,"modelId":472955}],"dockerImageVersionId":31153,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport pandas as pd\nimport nibabel as nib, torch, numpy as np\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-17T00:24:02.395792Z","iopub.execute_input":"2025-10-17T00:24:02.396079Z","iopub.status.idle":"2025-10-17T00:24:09.810966Z","shell.execute_reply.started":"2025-10-17T00:24:02.396052Z","shell.execute_reply":"2025-10-17T00:24:09.810198Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Explorando la data","metadata":{}},{"cell_type":"code","source":"folder_path = \"/kaggle/input/rsna-intracranial-aneurysm-detection/\"\ndf = pd.read_csv(folder_path + '/train_localizers.csv')\nprint(df.head())\n    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-17T00:24:09.812803Z","iopub.execute_input":"2025-10-17T00:24:09.813269Z","iopub.status.idle":"2025-10-17T00:24:09.860059Z","shell.execute_reply.started":"2025-10-17T00:24:09.813249Z","shell.execute_reply":"2025-10-17T00:24:09.859099Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\nfolder_path = \"/kaggle/input/rsna-intracranial-aneurysm-detection/\"\n# Load the two CSV files\ntrain_df = pd.read_csv(folder_path +  \"train.csv\")\nlocalizers_df = pd.read_csv(folder_path + \"train_localizers.csv\")\n\n# Find overlapping SeriesInstanceUID values\ncommon_subjects = set(train_df['SeriesInstanceUID']).intersection(\n    set(localizers_df['SeriesInstanceUID'])\n)\n\n# Convert to sorted list\ncommon_subjects_list = sorted(list(common_subjects))\n\n# Display results\nprint(f\"Found {len(common_subjects_list)} common subjects.\")\nprint(\"First 10 common subjects:\")\nfor s in common_subjects_list[:10]:\n    print(s)\n\n# Optionally, save to CSV\npd.DataFrame(common_subjects_list, columns=[\"SeriesInstanceUID\"]).to_csv(\n    \"common_subjects.csv\", index=False\n)\nprint(\"Saved to common_subjects.csv\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-17T00:24:09.863694Z","iopub.execute_input":"2025-10-17T00:24:09.863966Z","iopub.status.idle":"2025-10-17T00:24:09.928444Z","shell.execute_reply.started":"2025-10-17T00:24:09.863934Z","shell.execute_reply":"2025-10-17T00:24:09.927342Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Continue exploring the data:\nimport os\nimport random\nimport re\nimport json\nimport numpy as np\nimport pandas as pd\nimport pydicom\nimport matplotlib.pyplot as plt\n\n# -----------------------\n# Config\n# -----------------------\nTRAIN_CSV = folder_path + \"train.csv\"\nLOCALIZERS_CSV = folder_path + \"train_localizers.csv\"\nSERIES_DIR = folder_path + \"series\"   # contains series/{SeriesInstanceUID}/{SOPInstanceUID}.dcm\nN_SERIES = 5\nRANDOM_STATE = 42\n\n# -----------------------\n# Helpers\n# -----------------------\ndef parse_coordinates(coord_str):\n    \"\"\"\n    Robustly parse coordinates from various possible formats.\n    Expected to return (x, y) in pixel space.\n    Tries JSON first, then regex for two numbers.\n    \"\"\"\n    if pd.isna(coord_str):\n        return None\n    s = str(coord_str).strip()\n\n    # Try JSON-like (e.g., \"[x, y]\" or '{\"x\": x, \"y\": y}')\n    try:\n        obj = json.loads(s.replace(\"'\", '\"'))\n        if isinstance(obj, (list, tuple)) and len(obj) == 2:\n            return float(obj[0]), float(obj[1])\n        if isinstance(obj, dict) and \"x\" in obj and \"y\" in obj:\n            return float(obj[\"x\"]), float(obj[\"y\"])\n    except Exception:\n        pass\n\n    # Fallback: pull first two numbers in the string\n    nums = re.findall(r\"[-+]?\\d*\\.?\\d+(?:[eE][-+]?\\d+)?\", s)\n    if len(nums) >= 2:\n        return float(nums[0]), float(nums[1])\n\n    return None\n\ndef load_dicom_image(dcm_path):\n    \"\"\"\n    Load a DICOM image as a numpy array ready to display.\n    Handles MONOCHROME1 inversion and rescale slope/intercept.\n    \"\"\"\n    ds = pydicom.dcmread(dcm_path)\n    img = ds.pixel_array.astype(np.float32)\n\n    # Apply rescale if present\n    slope = float(getattr(ds, \"RescaleSlope\", 1.0))\n    intercept = float(getattr(ds, \"RescaleIntercept\", 0.0))\n    img = img * slope + intercept\n\n    # Handle MONOCHROME1 (invert)\n    if getattr(ds, \"PhotometricInterpretation\", \"\").upper() == \"MONOCHROME1\":\n        img = img.max() - img\n\n    # Normalize to 0..1 for display\n    if img.max() > img.min():\n        img = (img - img.min()) / (img.max() - img.min())\n\n    return ds, img\n\n# -----------------------\n# Load data\n# -----------------------\ntrain_df = pd.read_csv(TRAIN_CSV)\nloc_df = pd.read_csv(LOCALIZERS_CSV)\n\n# Keep only series that actually have localizers\nloc_series = loc_df[\"SeriesInstanceUID\"].unique()\nmeta_with_loc = train_df[train_df[\"SeriesInstanceUID\"].isin(loc_series)].copy()\n\n# Sample distinct series\nsample_series = (\n    meta_with_loc[\"SeriesInstanceUID\"]\n    .drop_duplicates()\n    .sample(min(N_SERIES, len(meta_with_loc)), random_state=RANDOM_STATE)\n    .tolist()\n)\n\nprint(\"Selected SeriesInstanceUIDs:\")\nfor uid in sample_series:\n    print(uid)\n\n# Build a per-series pick of a random SOP from localizers\nsop_choices = (\n    loc_df[loc_df[\"SeriesInstanceUID\"].isin(sample_series)]\n    .groupby(\"SeriesInstanceUID\")[\"SOPInstanceUID\"]\n    .apply(lambda s: random.Random(RANDOM_STATE).choice(list(s)))\n    .to_dict()\n)\n\n# Also grab one set of coordinates and location text per series for annotation\ncoord_map = (\n    loc_df[loc_df[\"SeriesInstanceUID\"].isin(sample_series)]\n    .groupby(\"SeriesInstanceUID\")\n    .agg({\"coordinates\": \"first\", \"location\": \"first\"})\n    .to_dict(orient=\"index\")\n)\n\n# Map SeriesInstanceUID -> (Modality, Age, Sex, AneurysmPresent)\nmeta_cols = [\"SeriesInstanceUID\", \"Modality\", \"PatientAge\", \"PatientSex\", \"Aneurysm Present\"]\nmeta_map = (\n    train_df[meta_cols]\n    .drop_duplicates(\"SeriesInstanceUID\")\n    .set_index(\"SeriesInstanceUID\")\n    .to_dict(orient=\"index\")\n)\n\n# -----------------------\n# Plot\n# -----------------------\nn = len(sample_series)\ncols = min(3, n)\nrows = int(np.ceil(n / cols))\nplt.figure(figsize=(5 * cols, 5 * rows))\n\nfor i, uid in enumerate(sample_series, 1):\n    sop = sop_choices.get(uid)\n    dcm_path = os.path.join(SERIES_DIR, uid, f\"{sop}.dcm\")\n\n    if not os.path.exists(dcm_path):\n        print(f\"⚠️ Missing file: {dcm_path}\")\n        continue\n\n    ds, img = load_dicom_image(dcm_path)\n\n    # Prepare annotation/meta\n    meta = meta_map.get(uid, {})\n    modality = meta.get(\"Modality\", \"N/A\")\n    age = meta.get(\"PatientAge\", \"N/A\")\n    sex = meta.get(\"PatientSex\", \"N/A\")\n    aneur = meta.get(\"Aneurysm Present\", np.nan)\n    aneur_str = \"Yes\" if aneur == 1 else (\"No\" if aneur == 0 else \"Unknown\")\n\n    # Coordinates (x, y) note: many datasets define (x=col, y=row)\n    coord_info = coord_map.get(uid, {})\n    coords = parse_coordinates(coord_info.get(\"coordinates\"))\n    loc_text = coord_info.get(\"location\", \"\")\n\n    ax = plt.subplot(rows, cols, i)\n    ax.imshow(img, cmap=\"gray\")\n    ax.axis(\"off\")\n\n    title = f\"{modality}, {age}, {sex}\\nAneurysm: {aneur_str}\\nSeries: {uid[:16]}…\\nSOP: {str(sop)[:16]}…\"\n    if loc_text:\n        title += f\"\\nLoc: {loc_text}\"\n    ax.set_title(title, fontsize=10)\n\n    # Overlay localization point if available\n    if coords is not None:\n        x, y = coords  # assume (x=column, y=row)\n        ax.scatter([x], [y], s=30)  # default marker; no color specified\n        # optional: label\n        # ax.text(x+3, y+3, \"aneurysm\", fontsize=8)\n\nplt.tight_layout()\nplt.show()\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-17T00:24:09.929468Z","iopub.execute_input":"2025-10-17T00:24:09.929803Z","iopub.status.idle":"2025-10-17T00:24:11.174258Z","shell.execute_reply.started":"2025-10-17T00:24:09.929772Z","shell.execute_reply":"2025-10-17T00:24:11.173381Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport re\nimport json\nimport random\nimport numpy as np\nimport pandas as pd\nimport pydicom\nimport matplotlib.pyplot as plt\nfrom matplotlib.patches import Circle\n\n# -----------------------\n# Config\n# -----------------------\nTRAIN_CSV = folder_path + \"train.csv\"\nLOCALIZERS_CSV = folder_path + \"train_localizers.csv\"\nSERIES_DIR = folder_path + \"series\"   # contains series/{SeriesInstanceUID}/{SOPInstanceUID}.dcm\nN_SERIES = 5\nRANDOM_STATE = 42\nMARK_STYLE = \"circle\"   # \"circle\" or \"x\"\nCIRCLE_RADIUS = 6       # pixels\nX_HALF_SIZE = 6         # pixels\n\n# -----------------------\n# Helpers\n# -----------------------\ndef parse_coordinates(coord_str):\n    \"\"\"\n    Parse coordinates from train_localizers 'coordinates' column.\n    Returns (x, y) in pixel coordinates (x=column, y=row) or None.\n    Tries JSON first, then regex.\n    \"\"\"\n    if pd.isna(coord_str):\n        return None\n    s = str(coord_str).strip()\n\n    # Try JSON-like formats: \"[x, y]\" or '{\"x\": x, \"y\": y}'\n    try:\n        obj = json.loads(s.replace(\"'\", '\"'))\n        if isinstance(obj, (list, tuple)) and len(obj) == 2:\n            return float(obj[0]), float(obj[1])\n        if isinstance(obj, dict) and \"x\" in obj and \"y\" in obj:\n            return float(obj[\"x\"]), float(obj[\"y\"])\n    except Exception:\n        pass\n\n    # Fallback: first two numbers\n    nums = re.findall(r\"[-+]?\\d*\\.?\\d+(?:[eE][-+]?\\d+)?\", s)\n    if len(nums) >= 2:\n        return float(nums[0]), float(nums[1])\n\n    return None\n\ndef load_dicom_image(dcm_path):\n    \"\"\"\n    Load a DICOM image as float32, apply rescale, fix MONOCHROME1,\n    and normalize to 0..1 for display.\n    \"\"\"\n    ds = pydicom.dcmread(dcm_path)\n    img = ds.pixel_array.astype(np.float32)\n\n    slope = float(getattr(ds, \"RescaleSlope\", 1.0))\n    intercept = float(getattr(ds, \"RescaleIntercept\", 0.0))\n    img = img * slope + intercept\n\n    # Invert if MONOCHROME1\n    if getattr(ds, \"PhotometricInterpretation\", \"\").upper() == \"MONOCHROME1\":\n        img = img.max() - img\n\n    # Normalize\n    rng = img.max() - img.min()\n    if rng > 0:\n        img = (img - img.min()) / rng\n\n    return ds, img\n\ndef draw_marker(ax, x, y, style=\"circle\", circle_radius=6, x_half_size=6):\n    \"\"\"\n    Draw a visible marker at (x, y) in pixel coordinates.\n    style: \"circle\" or \"x\"\n    \"\"\"\n    if style == \"circle\":\n        circ = Circle((x, y), radius=circle_radius, fill=False, linewidth=1.5)\n        ax.add_patch(circ)\n    else:  # \"x\"\n        ax.plot([x - x_half_size, x + x_half_size], [y - x_half_size, y + x_half_size], linewidth=1.5)\n        ax.plot([x - x_half_size, x + x_half_size], [y + x_half_size, y - x_half_size], linewidth=1.5)\n\n# -----------------------\n# Load data\n# -----------------------\ntrain_df = pd.read_csv(TRAIN_CSV)\nloc_df = pd.read_csv(LOCALIZERS_CSV)\n\n# Keep only series that have localizers\nseries_with_loc = loc_df[\"SeriesInstanceUID\"].unique()\nmeta_with_loc = train_df[train_df[\"SeriesInstanceUID\"].isin(series_with_loc)].copy()\n\n# Sample distinct series\nrng = random.Random(RANDOM_STATE)\nsample_series = (\n    meta_with_loc[\"SeriesInstanceUID\"]\n    .drop_duplicates()\n    .sample(min(N_SERIES, len(meta_with_loc)), random_state=RANDOM_STATE)\n    .tolist()\n)\n\nprint(\"Selected SeriesInstanceUIDs:\")\nfor uid in sample_series:\n    print(uid)\n\n# For each series, choose ONE SOPInstanceUID that has coordinates (random but deterministic)\npicked_rows = []\nfor uid in sample_series:\n    sub = loc_df[loc_df[\"SeriesInstanceUID\"] == uid].copy()\n    # Keep rows that yield valid coordinates\n    sub[\"parsed_coord\"] = sub[\"coordinates\"].apply(parse_coordinates)\n    sub = sub[~sub[\"parsed_coord\"].isna()]\n    if sub.empty:\n        continue\n    sop_choices = list(sub[\"SOPInstanceUID\"].unique())\n    sop = rng.choice(sop_choices)\n    picked_rows.append((uid, sop))\n\n# -----------------------\n# Plot each chosen slice with markers\n# -----------------------\nn = len(picked_rows)\ncols = min(3, n if n else 1)\nrows = int(np.ceil(max(1, n) / cols))\nplt.figure(figsize=(5 * cols, 5 * rows))\n\nfor i, (uid, sop) in enumerate(picked_rows, 1):\n    dcm_path = os.path.join(SERIES_DIR, uid, f\"{sop}.dcm\")\n    if not os.path.exists(dcm_path):\n        print(f\"⚠️ Missing file: {dcm_path}\")\n        continue\n\n    ds, img = load_dicom_image(dcm_path)\n\n    # gather metadata for title\n    meta = train_df.loc[train_df[\"SeriesInstanceUID\"] == uid].iloc[0]\n    modality = str(meta.get(\"Modality\", \"N/A\"))\n    age = str(meta.get(\"PatientAge\", \"N/A\"))\n    sex = str(meta.get(\"PatientSex\", \"N/A\"))\n    aneur = meta.get(\"Aneurysm Present\", np.nan)\n    aneur_str = \"Yes\" if aneur == 1 else (\"No\" if aneur == 0 else \"Unknown\")\n\n    ax = plt.subplot(rows, cols, i)\n    ax.imshow(img, cmap=\"gray\")\n    ax.axis(\"off\")\n    ax.set_title(f\"{modality}, {age}, {sex}\\nAneurysm: {aneur_str}\\nSeries: {uid[:16]}…\\nSOP: {str(sop)[:16]}…\", fontsize=10)\n\n    # Overlay ALL coordinates available for this (series, sop)\n    sub = loc_df[(loc_df[\"SeriesInstanceUID\"] == uid) & (loc_df[\"SOPInstanceUID\"] == sop)]\n    for coord_str in sub[\"coordinates\"]:\n        parsed = parse_coordinates(coord_str)\n        if parsed is None:\n            continue\n        x, y = parsed  # x=column, y=row\n        draw_marker(ax, x, y, style=MARK_STYLE, circle_radius=CIRCLE_RADIUS, x_half_size=X_HALF_SIZE)\n\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-17T00:24:11.175048Z","iopub.execute_input":"2025-10-17T00:24:11.175315Z","iopub.status.idle":"2025-10-17T00:24:12.442358Z","shell.execute_reply.started":"2025-10-17T00:24:11.175292Z","shell.execute_reply":"2025-10-17T00:24:12.441219Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Now trying to integrate with the segmentation nifti files","metadata":{}},{"cell_type":"code","source":"LABEL_MAP = {\n    1:  \"Other Posterior Circulation\",\n    2:  \"Basilar Tip\",\n    3:  \"Right Posterior Communicating Artery\",\n    4:  \"Left Posterior Communicating Artery\",\n    5:  \"Right Infraclinoid Internal Carotid Artery\",\n    6:  \"Left Infraclinoid Internal Carotid Artery\",\n    7:  \"Right Supraclinoid Internal Carotid Artery\",\n    8:  \"Left Supraclinoid Internal Carotid Artery\",\n    9:  \"Right Middle Cerebral Artery\",\n    10: \"Left Middle Cerebral Artery\",\n    11: \"Right Anterior Cerebral Artery\",\n    12: \"Left Anterior Cerebral Artery\",\n    13: \"Anterior Communicating Artery\",\n}\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-17T00:24:12.444926Z","iopub.execute_input":"2025-10-17T00:24:12.445386Z","iopub.status.idle":"2025-10-17T00:24:12.451963Z","shell.execute_reply.started":"2025-10-17T00:24:12.445358Z","shell.execute_reply":"2025-10-17T00:24:12.450791Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, glob, json, re\nimport numpy as np\nimport pandas as pd\nimport pydicom\nimport nibabel as nib\nimport matplotlib.pyplot as plt\nfrom matplotlib.patches import Patch, Circle\nfrom scipy.ndimage import zoom\n\n# -----------------------\n# Config\n# -----------------------\nTRAIN_CSV = folder_path + \"train.csv\"\nLOCALIZERS_CSV = folder_path + \"train_localizers.csv\"\nSERIES_DIR = folder_path + \"series\"          # <- set this correctly for your machine\nSEGS_DIR = folder_path + \"segmentations\"\nRANDOM_STATE = 42\nN_SERIES = 5\n\nLABEL_MAP = {\n    1:  \"Other Posterior Circulation\",\n    2:  \"Basilar Tip\",\n    3:  \"Right Posterior Communicating Artery\",\n    4:  \"Left Posterior Communicating Artery\",\n    5:  \"Right Infraclinoid Internal Carotid Artery\",\n    6:  \"Left Infraclinoid Internal Carotid Artery\",\n    7:  \"Right Supraclinoid Internal Carotid Artery\",\n    8:  \"Left Supraclinoid Internal Carotid Artery\",\n    9:  \"Right Middle Cerebral Artery\",\n    10: \"Left Middle Cerebral Artery\",\n    11: \"Right Anterior Cerebral Artery\",\n    12: \"Left Anterior Cerebral Artery\",\n    13: \"Anterior Communicating Artery\",\n}\nLABELS_TO_SHOW = list(LABEL_MAP.keys())\n\n# -----------------------\n# Helpers\n# -----------------------\ndef parse_coordinates(coord_str):\n    if pd.isna(coord_str): return None\n    s = str(coord_str).strip()\n    try:\n        obj = json.loads(s.replace(\"'\", '\"'))\n        if isinstance(obj, (list, tuple)) and len(obj) == 2: return float(obj[0]), float(obj[1])\n        if isinstance(obj, dict) and \"x\" in obj and \"y\" in obj: return float(obj[\"x\"]), float(obj[\"y\"])\n    except Exception:\n        pass\n    nums = re.findall(r\"[-+]?\\d*\\.?\\d+(?:[eE][-+]?\\d+)?\", s)\n    if len(nums) >= 2: return float(nums[0]), float(nums[1])\n    return None\n\ndef read_dicom_image(path):\n    ds = pydicom.dcmread(path)\n    img = ds.pixel_array.astype(np.float32)\n    slope = float(getattr(ds, \"RescaleSlope\", 1.0))\n    intercept = float(getattr(ds, \"RescaleIntercept\", 0.0))\n    img = img * slope + intercept\n    if getattr(ds, \"PhotometricInterpretation\", \"\").upper() == \"MONOCHROME1\":\n        img = img.max() - img\n    rng = img.max() - img.min()\n    if rng > 0: img = (img - img.min()) / rng\n    return ds, img\n\ndef series_files(series_uid):\n    \"\"\"Return a robust list of DICOM file paths for a series.\"\"\"\n    folder = os.path.join(SERIES_DIR, series_uid)\n    if not os.path.isdir(folder):\n        return []\n\n    # try common extensions\n    paths = []\n    for patt in (\"*.dcm\", \"*.DCM\", \"*\"):\n        paths.extend(sorted(glob.glob(os.path.join(folder, patt))))\n        if paths: break\n\n    # if wildcard pulled non-DICOM (json, txt), keep only files that pydicom can open\n    good = []\n    for p in paths:\n        if os.path.isdir(p):  # skip subdirs\n            continue\n        try:\n            ds = pydicom.dcmread(p, stop_before_pixels=True)\n            # require SOPInstanceUID to be present\n            _ = getattr(ds, \"SOPInstanceUID\", None)\n            good.append(p)\n        except Exception:\n            continue\n    return sorted(good)\n\ndef load_series_volume_robust(series_uid, loc_df=None):\n    \"\"\"\n    Build a 3D stack [Z,Y,X] for the given series.\n    Falls back to using SOPs from localizers if glob returns empty.\n    Returns: vol (np.ndarray) or None, sops (list), paths (list)\n    \"\"\"\n    paths = series_files(series_uid)\n\n    # Fallback: construct from localizers SOPs\n    if (not paths) and (loc_df is not None):\n        sub = loc_df[loc_df[\"SeriesInstanceUID\"] == series_uid]\n        candidates = []\n        for sop in sub[\"SOPInstanceUID\"].unique():\n            for ext in (\".dcm\", \".DCM\", \"\"):\n                p = os.path.join(SERIES_DIR, series_uid, f\"{sop}{ext}\")\n                if os.path.exists(p):\n                    candidates.append(p)\n                    break\n        paths = sorted(set(candidates))\n\n    if not paths:\n        return None, [], []\n\n    # read headers for sorting\n    info = []\n    for p in paths:\n        try:\n            ds = pydicom.dcmread(p, stop_before_pixels=True)\n            ipp = getattr(ds, \"ImagePositionPatient\", None)\n            inst = getattr(ds, \"InstanceNumber\", None)\n            sop = getattr(ds, \"SOPInstanceUID\", os.path.basename(p).replace(\".dcm\",\"\").replace(\".DCM\",\"\"))\n            info.append({\"path\": p, \"ipp\": np.array(ipp, dtype=float) if ipp is not None else None,\n                         \"inst\": int(inst) if inst is not None else None, \"sop\": sop})\n        except Exception:\n            continue\n\n    if not info:\n        return None, [], []\n\n    if all(x[\"ipp\"] is not None for x in info):\n        info.sort(key=lambda x: x[\"ipp\"][2] if len(x[\"ipp\"]) >= 3 else 0.0)\n    elif all(x[\"inst\"] is not None for x in info):\n        info.sort(key=lambda x: x[\"inst\"])\n    else:\n        info.sort(key=lambda x: x[\"path\"])\n\n    imgs, sops, used_paths = [], [], []\n    for x in info:\n        try:\n            _, img = read_dicom_image(x[\"path\"])\n            imgs.append(img); sops.append(x[\"sop\"]); used_paths.append(x[\"path\"])\n        except Exception:\n            pass\n\n    if not imgs:\n        return None, [], []\n    vol = np.stack(imgs, axis=0)\n    return vol, sops, used_paths\n\ndef find_seg_path(series_uid):\n    for ext in (\".nii.gz\", \".nii\"):\n        p = os.path.join(SEGS_DIR, f\"{series_uid}{ext}\")\n        if os.path.exists(p): return p\n    return None\n\ndef load_seg(seg_path):\n    ni = nib.load(seg_path)\n    data = ni.get_fdata().astype(np.int16)  # [X,Y,Z] typically\n    return data, ni.affine\n\ndef bring_seg_to_vol(seg_xyz, vol_zyx):\n    seg_zyx = np.transpose(seg_xyz, (2,1,0))\n    if seg_zyx.shape == vol_zyx.shape:\n        return seg_zyx\n    zoom_factors = (\n        vol_zyx.shape[0] / seg_zyx.shape[0],\n        vol_zyx.shape[1] / seg_zyx.shape[1],\n        vol_zyx.shape[2] / seg_zyx.shape[2],\n    )\n    return zoom(seg_zyx, zoom=zoom_factors, order=0)  # nearest for labels\n\ndef draw_localizers(ax, loc_rows, marker=\"circle\", radius=6, x_half=6):\n    for coord_str in loc_rows[\"coordinates\"]:\n        pt = parse_coordinates(coord_str)\n        if pt is None: continue\n        x, y = pt\n        if marker == \"circle\":\n            ax.add_patch(Circle((x, y), radius=radius, fill=False, linewidth=1.5))\n        else:\n            ax.plot([x - x_half, x + x_half], [y - x_half, y + x_half], linewidth=1.5)\n            ax.plot([x - x_half, x + x_half], [y + x_half, y - x_half], linewidth=1.5)\n\n# -----------------------\n# Main\n# -----------------------\ntrain_df = pd.read_csv(TRAIN_CSV)\nloc_df = pd.read_csv(LOCALIZERS_CSV)\n\n# choose N series that have localizers\nseries_pool = train_df[train_df[\"SeriesInstanceUID\"].isin(loc_df[\"SeriesInstanceUID\"].unique())]\nseries_uids = series_pool[\"SeriesInstanceUID\"].drop_duplicates().sample(\n    min(N_SERIES, len(series_pool)), random_state=RANDOM_STATE\n).tolist()\n\ncols = min(3, len(series_uids)); rows = int(np.ceil(len(series_uids)/cols))\nplt.figure(figsize=(5*cols, 5*rows))\n\nfor i, uid in enumerate(series_uids, 1):\n    # build volume robustly (or None)\n    vol, sops, paths = load_series_volume_robust(uid, loc_df=loc_df)\n\n    # pick a SOP from localizers\n    loc_sub = loc_df[loc_df[\"SeriesInstanceUID\"] == uid]\n    if loc_sub.empty:\n        print(f\"[skip] no localizers for {uid}\")\n        continue\n    sop = loc_sub[\"SOPInstanceUID\"].sample(1, random_state=RANDOM_STATE).iloc[0]\n\n    # find path for that SOP (fallbacks)\n    cand_paths = []\n    for ext in (\".dcm\", \".DCM\", \"\"):\n        p = os.path.join(SERIES_DIR, uid, f\"{sop}{ext}\")\n        if os.path.exists(p): cand_paths.append(p)\n    if not cand_paths and paths:\n        # try to match by SOP among discovered paths\n        for p in paths:\n            if sop in os.path.basename(p):\n                cand_paths.append(p); break\n\n    if not cand_paths:\n        print(f\"[warn] missing DICOM for SOP {sop} in {uid}\")\n        continue\n\n    dcm_path = cand_paths[0]\n    ds, img2d = read_dicom_image(dcm_path)\n\n    # get slice index if we have a stack\n    z_idx = None\n    if sops:\n        try: z_idx = sops.index(sop)\n        except ValueError: z_idx = None\n\n    # load segmentation if present and we have a stack\n    seg2d = None\n    seg_path = find_seg_path(uid)\n    if (seg_path is not None) and (vol is not None) and (z_idx is not None):\n        seg_xyz, _ = load_seg(seg_path)\n        seg_zyx = bring_seg_to_vol(seg_xyz, vol)\n        seg2d = seg_zyx[z_idx]\n\n    ax = plt.subplot(rows, cols, i)\n    ax.imshow(img2d, cmap=\"gray\"); ax.axis(\"off\")\n\n    # overlay segmentation (each label semi-transparent)\n    if seg2d is not None:\n        present = []\n        for lab in LABELS_TO_SHOW:\n            mask = (seg2d == lab)\n            if np.any(mask):\n                ax.imshow(np.ma.masked_where(~mask, mask), alpha=0.35)\n                present.append(lab)\n        if present:\n            handles = [Patch(label=f\"{lab}: {LABEL_MAP.get(lab,'Label')}\") for lab in present]\n            ax.legend(handles=handles, loc=\"lower right\", fontsize=7, frameon=True)\n\n    # overlay all localizer points for this SOP\n    draw_localizers(ax, loc_sub[loc_sub[\"SOPInstanceUID\"] == sop], marker=\"circle\", radius=6)\n\n    # Title\n    meta = train_df.loc[train_df[\"SeriesInstanceUID\"] == uid].iloc[0]\n    aneur = meta.get(\"Aneurysm Present\", np.nan)\n    aneur_str = \"Yes\" if aneur == 1 else (\"No\" if aneur == 0 else \"Unknown\")\n    ax.set_title(f\"{meta.get('Modality','N/A')}, {meta.get('PatientAge','N/A')}, {meta.get('PatientSex','N/A')}\\n\"\n                 f\"Aneurysm: {aneur_str}\\nSeries: {uid[:16]}…\\nSOP: {str(sop)[:16]}…\",\n                 fontsize=9)\n\nplt.tight_layout()\nplt.show()\n\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-17T00:24:12.45272Z","iopub.execute_input":"2025-10-17T00:24:12.453026Z","iopub.status.idle":"2025-10-17T00:25:02.340141Z","shell.execute_reply.started":"2025-10-17T00:24:12.453002Z","shell.execute_reply":"2025-10-17T00:25:02.339243Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Other way to see it","metadata":{}},{"cell_type":"code","source":"import os, glob, json, re, math\nimport numpy as np\nimport pandas as pd\nimport pydicom\nimport nibabel as nib\nfrom nibabel.processing import resample_from_to\nimport matplotlib.pyplot as plt\nfrom matplotlib.patches import Patch, Circle\nfrom scipy.ndimage import zoom\n\n# =========================\n# Config\n# =========================\nTRAIN_CSV = folder_path + \"train.csv\"\nLOCALIZERS_CSV = folder_path + \"train_localizers.csv\"\nSERIES_DIR = folder_path +  \"series\"          # series/{SeriesInstanceUID}/{SOPInstanceUID}.dcm\nSEGS_DIR = folder_path + \"segmentations\"     # segmentations/{SeriesInstanceUID}.nii(.gz)\n\nN_SERIES = 5\nRANDOM_STATE = 42\n\n# If the SOP slice has no labels, look for the nearest slice (± SEARCH_RADIUS)\nUSE_NEAREST_LABEL_SLICE = True\nSEARCH_RADIUS = 5  # slices\n\n# Marker style for localizer points\nMARK_STYLE = \"circle\"   # \"circle\" or \"x\"\nCIRCLE_RADIUS = 6\nX_HALF_SIZE = 6\n\n# Label map\nLABEL_MAP = {\n    1:  \"Other Posterior Circulation\",\n    2:  \"Basilar Tip\",\n    3:  \"Right Posterior Communicating Artery\",\n    4:  \"Left Posterior Communicating Artery\",\n    5:  \"Right Infraclinoid Internal Carotid Artery\",\n    6:  \"Left Infraclinoid Internal Carotid Artery\",\n    7:  \"Right Supraclinoid Internal Carotid Artery\",\n    8:  \"Left Supraclinoid Internal Carotid Artery\",\n    9:  \"Right Middle Cerebral Artery\",\n    10: \"Left Middle Cerebral Artery\",\n    11: \"Right Anterior Cerebral Artery\",\n    12: \"Left Anterior Cerebral Artery\",\n    13: \"Anterior Communicating Artery\",\n}\nLABELS_TO_SHOW = list(LABEL_MAP.keys())\n\n# =========================\n# Utilities\n# =========================\ndef parse_coordinates(coord_str):\n    if pd.isna(coord_str): return None\n    s = str(coord_str).strip()\n    # Try JSON-like formats\n    try:\n        obj = json.loads(s.replace(\"'\", '\"'))\n        if isinstance(obj, (list, tuple)) and len(obj) == 2:\n            return float(obj[0]), float(obj[1])\n        if isinstance(obj, dict) and \"x\" in obj and \"y\" in obj:\n            return float(obj[\"x\"]), float(obj[\"y\"])\n    except Exception:\n        pass\n    # Fallback: first two numbers in the string\n    nums = re.findall(r\"[-+]?\\d*\\.?\\d+(?:[eE][-+]?\\d+)?\", s)\n    if len(nums) >= 2:\n        return float(nums[0]), float(nums[1])\n    return None\n\ndef dicom_path_for_sop(series_uid, sop, series_dir=SERIES_DIR):\n    for ext in (\".dcm\", \".DCM\", \"\"):\n        p = os.path.join(series_dir, series_uid, f\"{sop}{ext}\")\n        if os.path.exists(p): return p\n    return None\n\ndef any_dicom_in_series(series_uid, loc_df):\n    sub = loc_df[loc_df[\"SeriesInstanceUID\"] == series_uid]\n    for sop in sub[\"SOPInstanceUID\"].unique():\n        if dicom_path_for_sop(series_uid, sop):\n            return True\n    folder = os.path.join(SERIES_DIR, series_uid)\n    if os.path.isdir(folder):\n        for patt in (\"*.dcm\", \"*.DCM\", \"*\"):\n            if glob.glob(os.path.join(folder, patt)): return True\n    return False\n\ndef read_dicom_image(path):\n    ds = pydicom.dcmread(path)\n    arr = ds.pixel_array.astype(np.float32)\n    slope = float(getattr(ds, \"RescaleSlope\", 1.0))\n    intercept = float(getattr(ds, \"RescaleIntercept\", 0.0))\n    arr = arr * slope + intercept\n    if getattr(ds, \"PhotometricInterpretation\", \"\").upper() == \"MONOCHROME1\":\n        arr = arr.max() - arr\n    rng = arr.max() - arr.min()\n    if rng > 0: arr = (arr - arr.min()) / rng\n    return ds, arr\n\ndef series_file_candidates(series_uid):\n    folder = os.path.join(SERIES_DIR, series_uid)\n    paths = []\n    if os.path.isdir(folder):\n        for patt in (\"*.dcm\", \"*.DCM\", \"*\"):\n            paths = sorted(glob.glob(os.path.join(folder, patt)))\n            if paths: break\n    # keep only readable DICOM files\n    out = []\n    for p in paths:\n        if os.path.isdir(p): continue\n        try:\n            ds = pydicom.dcmread(p, stop_before_pixels=True)\n            _ = getattr(ds, \"SOPInstanceUID\", None)\n            out.append(p)\n        except Exception:\n            pass\n    return sorted(out)\n\ndef build_dicom_stack_with_affine(series_uid):\n    \"\"\"\n    Returns:\n      vol_zyx: np.ndarray [Z,Y,X] (float32 in 0..1)\n      sops: list of SOPInstanceUIDs (sorted to match Z)\n      paths: list of file paths (sorted to match Z)\n      affine_dicom: 4x4 affine mapping voxel->[x,y,z] in patient space (mm)\n    \"\"\"\n    paths = series_file_candidates(series_uid)\n    if not paths: return None, [], [], None\n\n    # read minimal headers for sorting & affine\n    info = []\n    for p in paths:\n        try:\n            ds = pydicom.dcmread(p, stop_before_pixels=True)\n            sop = getattr(ds, \"SOPInstanceUID\", os.path.basename(p).split(\".\")[0])\n            ipp = np.array(getattr(ds, \"ImagePositionPatient\", [0,0,0]), dtype=float)\n            iop = np.array(getattr(ds, \"ImageOrientationPatient\", [1,0,0,0,1,0]), dtype=float)\n            inst = getattr(ds, \"InstanceNumber\", None)\n            info.append({\"path\": p, \"sop\": sop, \"ipp\": ipp, \"iop\": iop, \"inst\": inst})\n        except Exception:\n            pass\n    if not info: return None, [], [], None\n\n    # sort by z-position (fallback to InstanceNumber/path)\n    if all(x[\"ipp\"] is not None for x in info):\n        # sort by projection onto slice normal\n        row_cos = info[0][\"iop\"][0:3]\n        col_cos = info[0][\"iop\"][3:6]\n        normal = np.cross(row_cos, col_cos)\n        info.sort(key=lambda x: float(np.dot(x[\"ipp\"], normal)))\n    elif all(x[\"inst\"] is not None for x in info):\n        info.sort(key=lambda x: int(x[\"inst\"]))\n    else:\n        info.sort(key=lambda x: x[\"path\"])\n\n    # read pixels in sorted order\n    imgs, sops, used, headers = [], [], [], []\n    for x in info:\n        try:\n            ds, img = read_dicom_image(x[\"path\"])\n            imgs.append(img); sops.append(x[\"sop\"]); used.append(x[\"path\"]); headers.append(ds)\n        except Exception:\n            pass\n    if not imgs: return None, [], [], None\n    vol_zyx = np.stack(imgs, axis=0)  # [Z,Y,X]\n\n    # build a proper DICOM-based affine for the stack:\n    # voxel order -> (row=y, col=x, slice=z) in our vol_zyx\n    ds0 = headers[0]\n    iop = np.array(getattr(ds0, \"ImageOrientationPatient\", [1,0,0,0,1,0]), dtype=float)\n    row_cos = iop[0:3]    # direction of columns? In DICOM: first three are row direction (along row), second three column direction.\n    col_cos = iop[3:6]\n    normal = np.cross(row_cos, col_cos)\n\n    px_spacing = getattr(ds0, \"PixelSpacing\", [1.0, 1.0])\n    dy = float(px_spacing[0])   # row spacing\n    dx = float(px_spacing[1])   # col spacing\n\n    # Estimate slice spacing from adjacent IPPs if available, else SliceThickness\n    if len(headers) >= 2:\n        ipp0 = np.array(getattr(headers[0], \"ImagePositionPatient\", [0,0,0]), dtype=float)\n        ipp1 = np.array(getattr(headers[1], \"ImagePositionPatient\", [0,0,0]), dtype=float)\n        dz = abs(np.dot(ipp1 - ipp0, normal))\n        if dz == 0:\n            dz = float(getattr(ds0, \"SpacingBetweenSlices\", getattr(ds0, \"SliceThickness\", 1.0)))\n    else:\n        dz = float(getattr(ds0, \"SpacingBetweenSlices\", getattr(ds0, \"SliceThickness\", 1.0)))\n\n    # origin = IPP of first (sorted) slice\n    origin = np.array(getattr(headers[0], \"ImagePositionPatient\", [0,0,0]), dtype=float)\n\n    # We want NIfTI-style data order [X,Y,Z], so when we create the target image,\n    # we'll transpose vol to [X,Y,Z] and create the affine columns as:\n    # col 0 = X axis (columns) = col_cos * dx\n    # col 1 = Y axis (rows)    = row_cos * dy\n    # col 2 = Z axis (slices)  = normal * dz\n    # col 3 = origin\n    affine = np.eye(4, dtype=float)\n    affine[0:3, 0] = col_cos * dx\n    affine[0:3, 1] = row_cos * dy\n    affine[0:3, 2] = normal * dz\n    affine[0:3, 3] = origin\n\n    return vol_zyx, sops, used, affine\n\ndef find_seg_path(series_uid):\n    for ext in (\".nii.gz\", \".nii\"):\n        p = os.path.join(SEGS_DIR, f\"{series_uid}{ext}\")\n        if os.path.exists(p): return p\n    return None\n\ndef resample_seg_to_dicom_grid(seg_path, target_shape_zyx, target_affine_xyz):\n    \"\"\"\n    seg_path: path to NIfTI labels (label ints)\n    target_shape_zyx: DICOM volume shape [Z,Y,X]\n    target_affine_xyz: 4x4 affine for the DICOM grid with data order [X,Y,Z]\n                       (i.e., the affine we built above)\n    Return: seg_zyx on the DICOM grid (int16)\n    \"\"\"\n    seg_img = nib.load(seg_path)                    # seg data typically [X,Y,Z] in its own affine\n    # Create a NIfTI \"target\" image that represents the DICOM grid\n    # We pass a dummy zero array just to carry shape & affine; resample_from_to uses shape+affine.\n    target_shape_xyz = (target_shape_zyx[2], target_shape_zyx[1], target_shape_zyx[0])  # [X,Y,Z]\n    target_img = nib.Nifti1Image(np.zeros(target_shape_xyz, dtype=np.int16), target_affine_xyz)\n    seg_resampled_img = resample_from_to(seg_img, target_img, order=0)  # nearest neighbor preserves labels\n    seg_xyz = seg_resampled_img.get_fdata().astype(np.int16)            # [X,Y,Z] on DICOM grid\n    seg_zyx = np.transpose(seg_xyz, (2,1,0))                            # -> [Z,Y,X] to match vol\n    return seg_zyx\n\ndef draw_localizers(ax, rows, marker=\"circle\", radius=6, x_half=6):\n    for coord_str in rows[\"coordinates\"]:\n        pt = parse_coordinates(coord_str)\n        if pt is None: continue\n        x, y = pt\n        if marker == \"circle\":\n            ax.add_patch(Circle((x, y), radius=radius, fill=False, linewidth=1.5))\n        else:\n            ax.plot([x - x_half, x + x_half], [y - x_half, y + x_half], linewidth=1.5)\n            ax.plot([x - x_half, x + x_half], [y + x_half, y - x_half], linewidth=1.5)\n\ndef present_labels(seg2d, labels_whitelist):\n    if seg2d is None: return []\n    vals = np.unique(seg2d)\n    return [int(v) for v in vals if v in labels_whitelist and v != 0]\n\n# =========================\n# Main\n# =========================\ntrain_df = pd.read_csv(TRAIN_CSV)\nloc_df = pd.read_csv(LOCALIZERS_CSV)\n\n# series with localizers AND at least one DICOM present\nseries_candidates = train_df[train_df[\"SeriesInstanceUID\"].isin(loc_df[\"SeriesInstanceUID\"].unique())][\"SeriesInstanceUID\"].drop_duplicates()\navailable = [uid for uid in series_candidates if any_dicom_in_series(uid, loc_df)]\n\nif not available:\n    raise RuntimeError(\"No usable series found. Check SERIES_DIR path and file extensions.\")\n\nnp.random.seed(RANDOM_STATE)\nseries_uids = list(np.random.choice(available, size=min(N_SERIES, len(available)), replace=False))\n\ncols = min(3, len(series_uids))\nrows = int(math.ceil(len(series_uids) / cols))\nplt.figure(figsize=(5*cols, 5*rows))\n\nfor i, uid in enumerate(series_uids, 1):\n    # Build DICOM stack and affine\n    vol, sops, _, affine_xyz = build_dicom_stack_with_affine(uid)\n\n    # Choose a localizer SOP that exists on disk\n    loc_sub_all = loc_df[loc_df[\"SeriesInstanceUID\"] == uid]\n    chosen_sop = None\n    sop_path = None\n    for sop in loc_sub_all[\"SOPInstanceUID\"].unique():\n        p = dicom_path_for_sop(uid, sop)\n        if p:\n            chosen_sop = sop\n            sop_path = p\n            break\n    if sop_path is None:\n        print(f\"[skip] No readable SOP DICOM for series {uid}\")\n        continue\n\n    # Read the SOP slice (image used for display)\n    _, sop_img2d = read_dicom_image(sop_path)\n\n    # Find z index of SOP in the stack (if we have a stack)\n    z_idx = None\n    if sops:\n        try:\n            z_idx = sops.index(chosen_sop)\n        except ValueError:\n            z_idx = None\n\n    # Load and resample segmentation (if available + we have stack & affine)\n    seg2d = None\n    seg_path = find_seg_path(uid)\n    seg_zyx = None\n    if seg_path and (vol is not None) and (z_idx is not None) and (affine_xyz is not None):\n        seg_zyx = resample_seg_to_dicom_grid(seg_path, vol.shape, affine_xyz)\n        seg2d = seg_zyx[z_idx]\n\n    # If requested, pick a nearby slice with labels when SOP-slice has none\n    used_slice = \"SOP\"\n    used_z = z_idx\n    if USE_NEAREST_LABEL_SLICE and (seg_zyx is not None) and (z_idx is not None):\n        if not np.any(seg2d):\n            best_z = None\n            for dz in range(1, SEARCH_RADIUS+1):\n                for cand in (z_idx - dz, z_idx + dz):\n                    if 0 <= cand < seg_zyx.shape[0] and np.any(seg_zyx[cand]):\n                        best_z = cand\n                        break\n                if best_z is not None:\n                    break\n            if best_z is not None:\n                seg2d = seg_zyx[best_z]\n                used_slice = f\"nearest z={best_z} (SOP z={z_idx})\"\n                used_z = best_z\n\n    # Plot\n    ax = plt.subplot(rows, cols, i)\n    # Show the SOP image (what the localizer references). This keeps the coordinates meaningful.\n   # ax.imshow(sop_img2d, cmap=\"gray\")\n    #ax.axis(\"off\")\n\n    img2show = sop_img2d if (used_z is None or used_z == z_idx) else vol[used_z]\n    ax.imshow(img2show, cmap=\"gray\")\n    \n\n    # Overlay segmentation (as contours). If we switched to a different z to show labels,\n    # we draw those contours on top of the SOP image — good for a quick view, but be aware\n    # that vessels may appear slightly offset if there's slice mismatch.\n    if seg2d is None or not np.any(seg2d):\n        ax.text(5, 15, \"No segmentation visible on SOP slice\",\n                fontsize=9, color=\"w\", bbox=dict(facecolor=\"0.1\", alpha=0.6, pad=3))\n    else:\n        labs_here = present_labels(seg2d, LABELS_TO_SHOW)\n        handles = []\n        for lab in labs_here:\n            ax.contour((seg2d == lab).astype(float), levels=[0.5], linewidths=1.5)\n            handles.append(Patch(label=f\"{lab}: {LABEL_MAP.get(lab, 'Label')}\"))\n        if handles:\n            ax.legend(handles=handles, loc=\"lower right\", fontsize=7, frameon=True)\n        if used_slice != \"SOP\":\n            ax.text(5, 35, f\"Showing contours from {used_slice}\",\n                    fontsize=8, color=\"w\", bbox=dict(facecolor=\"0.1\", alpha=0.6, pad=3))\n\n    # Overlay localizer points for this SOP (coordinates are for the SOP image)\n    draw_localizers(ax, loc_sub_all[loc_sub_all[\"SOPInstanceUID\"] == chosen_sop],\n                    marker=MARK_STYLE, radius=CIRCLE_RADIUS, x_half=X_HALF_SIZE)\n\n    # Title\n    meta = train_df.loc[train_df[\"SeriesInstanceUID\"] == uid].iloc[0]\n    aneur = meta.get(\"Aneurysm Present\", np.nan)\n    aneur_str = \"Yes\" if aneur == 1 else (\"No\" if aneur == 0 else \"Unknown\")\n    ax.set_title(f\"{meta.get('Modality','N/A')}, {meta.get('PatientAge','N/A')}, {meta.get('PatientSex','N/A')}\\n\"\n                 f\"Aneurysm: {aneur_str}\\nSeries: {uid[:16]}…  SOP: {str(chosen_sop)[:16]}…\",\n                 fontsize=9)\n\nplt.tight_layout()\nplt.show()\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-17T00:25:02.343696Z","iopub.execute_input":"2025-10-17T00:25:02.343968Z","iopub.status.idle":"2025-10-17T00:25:54.517374Z","shell.execute_reply.started":"2025-10-17T00:25:02.343949Z","shell.execute_reply":"2025-10-17T00:25:54.516364Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Now trying to do do the training in the notebook","metadata":{}},{"cell_type":"code","source":"#!pip -q install pydicom nibabel scipy tqdm\n#!pip -q install scipy tqdm","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 1. Set paths to the data","metadata":{}},{"cell_type":"code","source":"folder_path = \"/kaggle/input/rsna-intracranial-aneurysm-detection/\"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from pathlib import Path\n\n# Point these to your files\nDATA_DIR = Path(\"/kaggle/input/rsna-intracranial-aneurysm-detection\")   # <-- change to your dataset folder\nSERIES_DIR = DATA_DIR/\"series\"                   # series/{SeriesInstanceUID}/{SOPInstanceUID}.dcm\nSEGS_DIR   = DATA_DIR/\"segmentations\"            # optional NIfTI labels (not required for this baseline)\nTRAIN_CSV  = DATA_DIR/\"train.csv\"\nLOCAL_CSV  = DATA_DIR/\"train_localizers.csv\"\n\nOUT_DIR    = Path(\"/kaggle/working/out_3dunet\") # models & logs\nOUT_DIR.mkdir(parents=True, exist_ok=True)\n\nprint(TRAIN_CSV.exists(), LOCAL_CSV.exists(), SERIES_DIR.exists())\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 2. Imports of the necessary libraries and creating the config class","metadata":{}},{"cell_type":"code","source":"import os, re, json, math, glob, random\nfrom dataclasses import dataclass\nfrom typing import List, Tuple, Dict\n\nimport numpy as np\nimport pandas as pd\nimport pydicom\nfrom tqdm import tqdm\nfrom scipy.ndimage import zoom, gaussian_filter, maximum_filter\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@dataclass\nclass CFG:\n    # data\n    series_dir: str = str(SERIES_DIR)\n    train_csv: str  = str(TRAIN_CSV)\n    local_csv: str  = str(LOCAL_CSV)\n    target_spacing_mm: float = 1.0\n    window_percentiles: Tuple[float, float] = (1.0, 99.0)\n    zscore: bool = True\n\n    # training\n    patch_size: Tuple[int,int,int] = (96, 128, 128)  # (Z,Y,X) — lower if OOM\n    batch_size: int = 2  # With 2 epochs gets very good loss \n    epochs: int = 2                                  # bump for better results\n    lr: float = 1e-3\n    weight_decay: float = 1e-5\n    num_workers: int = 2                              # Kaggle-safe\n    pos_fraction: float = 0.5\n    neg_distance_min: int = 32\n    gauss_sigma_vox: float = 2.0\n    amp: bool = True\n    seed: int = 42\n\n    # inference\n    sw_overlap: Tuple[int,int,int] = (32, 48, 48)\n    pred_thresh: float = 0.35\n    nms_size: int = 5\n    topk: int = 10\n\n    # output\n    out_dir: str = str(OUT_DIR)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 3. Dicom Preprocessing helpers ","metadata":{}},{"cell_type":"code","source":"def set_seed(seed=CFG.seed):\n    random.seed(seed); np.random.seed(seed)\n    torch.manual_seed(seed); torch.cuda.manual_seed_all(seed)\n\ndef percent_window(img, low=1.0, high=99.0):\n    lo, hi = np.percentile(img, [low, high])\n    if hi <= lo: hi = lo + 1.0\n    return np.clip((img - lo) / (hi - lo), 0.0, 1.0)\n\ndef read_dicom_header(path):\n    return pydicom.dcmread(path, stop_before_pixels=True)\n\ndef read_dicom_pixels(path):\n    ds = pydicom.dcmread(path)\n    arr = ds.pixel_array.astype(np.float32)\n    slope = float(getattr(ds, \"RescaleSlope\", 1.0))\n    intercept = float(getattr(ds, \"RescaleIntercept\", 0.0))\n    arr = arr * slope + intercept\n    if str(getattr(ds, \"PhotometricInterpretation\", \"\")).upper() == \"MONOCHROME1\":\n        arr = arr.max() - arr\n    return ds, arr\n\ndef series_files(series_uid, series_dir=CFG.series_dir):\n    folder = os.path.join(series_dir, series_uid)\n    if not os.path.isdir(folder): return []\n    cands = []\n    for patt in (\"*.dcm\", \"*.DCM\", \"*\"):\n        paths = glob.glob(os.path.join(folder, patt))\n        if paths:\n            cands.extend(paths); break\n    good = []\n    for p in sorted(set(cands)):\n        if os.path.isdir(p): continue\n        try:\n            ds = read_dicom_header(p)\n            _ = getattr(ds, \"SOPInstanceUID\", None)\n            good.append(p)\n        except Exception:\n            pass\n    return sorted(good)\n\ndef sort_slices_by_geometry(paths: List[str]):\n    meta = []\n    for p in paths:\n        try:\n            ds = read_dicom_header(p)\n            sop = str(getattr(ds, \"SOPInstanceUID\", os.path.basename(p).split(\".\")[0]))\n            ipp = np.array(getattr(ds, \"ImagePositionPatient\", [0,0,0]), dtype=float)\n            iop = np.array(getattr(ds, \"ImageOrientationPatient\", [1,0,0,0,1,0]), dtype=float)\n            inst = getattr(ds, \"InstanceNumber\", None)\n            meta.append(dict(path=p, sop=sop, ipp=ipp, iop=iop, inst=inst))\n        except Exception:\n            pass\n    if not meta: return []\n    have_iop = all(len(m[\"iop\"]) == 6 for m in meta)\n    have_ipp = all(len(m[\"ipp\"]) == 3 for m in meta)\n    if have_iop and have_ipp:\n        row = meta[0][\"iop\"][0:3]; col = meta[0][\"iop\"][3:6]\n        normal = np.cross(row, col)\n        meta.sort(key=lambda x: float(np.dot(x[\"ipp\"], normal)))\n    elif all(m[\"inst\"] is not None for m in meta):\n        meta.sort(key=lambda x: int(x[\"inst\"]))\n    else:\n        meta.sort(key=lambda x: x[\"path\"])\n    return meta\n\ndef build_stack(series_uid, series_dir=CFG.series_dir):\n    \"\"\"\n    Returns:\n      vol [Z,Y,X] float32 0..1 (z-scored if CFG.zscore),\n      sops list (may include 'SOP|frame' for multi-frame),\n      spacing (dz,dy,dx),\n      origin (3,), dircos (3x3)\n    \"\"\"\n    paths = series_files(series_uid, series_dir)\n    meta = sort_slices_by_geometry(paths)\n    if not meta:\n        return None, [], None, None, None\n\n    imgs, sops, headers = [], [], []\n\n    for m in meta:\n        ds, img = read_dicom_pixels(m[\"path\"])\n        sop = m[\"sop\"]\n\n        # --- normalize to 2D per \"slice\" ---\n        if img.ndim == 3:\n            # Case A: color (RGB) like [H, W, 3] or [3, H, W]\n            if (img.shape[-1] in (3,4)) and (img.ndim == 3):\n                # [H, W, C] -> grayscale\n                if img.shape[0] != ds.Rows or img.shape[1] != ds.Columns:\n                    # If layout is [C, H, W], move channel to last\n                    if img.shape[0] in (3,4) and img.shape[1] == ds.Rows and img.shape[2] == ds.Columns:\n                        img = np.moveaxis(img, 0, -1)\n                img = img.astype(np.float32)\n                img = img.mean(axis=-1)  # collapse color\n                imgs.append(img); sops.append(sop); headers.append(ds)\n\n            # Case B: multi-frame [F, H, W]\n            elif img.shape[0] > 1 and img.shape[1] == ds.Rows and img.shape[2] == ds.Columns:\n                F = img.shape[0]\n                for f in range(F):\n                    imgs.append(img[f].astype(np.float32))\n                    sops.append(f\"{sop}|{f}\")  # disambiguate frame\n                    headers.append(ds)\n            else:\n                # Unknown 3D layout: best-effort collapse along the last axis\n                if img.shape[-1] == ds.Columns:\n                    # probably [H, C, W] -> collapse middle\n                    img2 = img.mean(axis=1)\n                else:\n                    img2 = img.mean(axis=-1)\n                imgs.append(img2.astype(np.float32)); sops.append(sop); headers.append(ds)\n\n        elif img.ndim == 2:\n            imgs.append(img.astype(np.float32)); sops.append(sop); headers.append(ds)\n\n        else:\n            # 1D or >=4D – skip this file\n            continue\n\n    if not imgs:\n        return None, [], None, None, None\n\n    vol = np.stack(imgs, axis=0).astype(np.float32)  # [Z,Y,X]\n\n    # windowing then z-score\n    vol = percent_window(vol, *CFG.window_percentiles)\n    if CFG.zscore:\n        mu, sd = vol.mean(), vol.std() + 1e-6\n        vol = (vol - mu) / sd\n\n    # --- geometry ---\n    ds0 = headers[0]\n    px = getattr(ds0, \"PixelSpacing\", [1.0, 1.0])\n    dy = float(px[0]); dx = float(px[1])\n    iop = np.array(getattr(ds0, \"ImageOrientationPatient\", [1,0,0,0,1,0]), dtype=float)\n    row = iop[0:3]; col = iop[3:6]; normal = np.cross(row, col)\n\n    # slice spacing from IPP diff or thickness\n    if len(headers) >= 2:\n        ipp0 = np.array(getattr(headers[0], \"ImagePositionPatient\", [0,0,0]), dtype=float)\n        ipp1 = np.array(getattr(headers[1], \"ImagePositionPatient\", [0,0,0]), dtype=float)\n        dz = abs(np.dot(ipp1 - ipp0, normal))\n        if dz == 0:\n            dz = float(getattr(ds0, \"SpacingBetweenSlices\", getattr(ds0, \"SliceThickness\", 1.0)))\n    else:\n        dz = float(getattr(ds0, \"SpacingBetweenSlices\", getattr(ds0, \"SliceThickness\", 1.0)))\n    spacing = (dz, dy, dx)\n\n    origin = np.array(getattr(headers[0], \"ImagePositionPatient\", [0,0,0]), dtype=float)\n    dircos = np.vstack([col, row, normal]).T  # maps [x,y,z] vox * spacing -> mm\n\n    return vol, sops, spacing, origin, dircos\n\ndef resample_isotropic(vol_zyx, spacing, target_mm=CFG.target_spacing_mm):\n    dz, dy, dx = spacing\n    factors = (dz/target_mm, dy/target_mm, dx/target_mm)\n    vol_iso = zoom(vol_zyx, zoom=factors, order=1)\n    return vol_iso.astype(np.float32), (target_mm, target_mm, target_mm)\n\ndef parse_coord(s):\n    if pd.isna(s): return None\n    txt = str(s).strip()\n    try:\n        obj = json.loads(txt.replace(\"'\", '\"'))\n        if isinstance(obj, (list, tuple)) and len(obj) == 2: return float(obj[0]), float(obj[1])\n        if isinstance(obj, dict): return float(obj[\"x\"]), float(obj[\"y\"])\n    except Exception: pass\n    nums = re.findall(r\"[-+]?\\d*\\.?\\d+(?:[eE][-+]?\\d+)?\", txt)\n    if len(nums) >= 2: return float(nums[0]), float(nums[1])\n    return None\n\ndef make_heatmap(shape_zyx, points_zyx, sigma):\n    Z,Y,X = shape_zyx\n    heat = np.zeros((Z,Y,X), np.float32)\n    for (z,y,x) in points_zyx or []:\n        zc, yc, xc = int(round(z)), int(round(y)), int(round(x))\n        if 0 <= zc < Z and 0 <= yc < Y and 0 <= xc < X:\n            heat[zc, yc, xc] = 1.0\n    if sigma > 0:\n        heat = gaussian_filter(heat, sigma=sigma)\n        heat = heat / (heat.max() + 1e-6)\n    return heat\n\ndef crop_center(vol, center, size):\n    Z, Y, X = vol.shape\n    dz, dy, dx = size\n    cz, cy, cx = center\n\n    z1 = max(0, cz - dz // 2); z2 = min(Z, z1 + dz)\n    y1 = max(0, cy - dy // 2); y2 = min(Y, y1 + dy)\n    x1 = max(0, cx - dx // 2); x2 = min(X, x1 + dx)\n\n    patch = vol[z1:z2, y1:y2, x1:x2]\n    pad = (\n        (0, dz - patch.shape[0]),\n        (0, dy - patch.shape[1]),\n        (0, dx - patch.shape[2]),\n    )\n    if any(p > 0 for pair in pad for p in pair):\n        patch = np.pad(patch, pad, mode=\"constant\")\n\n    return patch.astype(np.float32), (z1, y1, x1)\n\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 4. Managing the dataset","metadata":{}},{"cell_type":"code","source":"class AneurysmPatchDataset(Dataset):\n    def __init__(self, series_ids: List[str], train_df: pd.DataFrame, loc_df: pd.DataFrame):\n        self.series_ids = series_ids\n        self.train_df = train_df\n        self.loc_df = loc_df\n        self.cache: Dict[str, Dict] = {}\n        self.points_by_series: Dict[str, List[tuple]] = {}\n\n        for uid in tqdm(series_ids, desc=\"Caching volumes\"):\n            vol, sops, spacing, origin, dircos = build_stack(uid, CFG.series_dir)\n            if vol is None:\n                continue\n            sop2z = {sop:z for z,sop in enumerate(sops)}\n            pts = []\n            for _, row in self.loc_df[self.loc_df[\"SeriesInstanceUID\"]==uid].iterrows():\n                sop = str(row[\"SOPInstanceUID\"]); xy = parse_coord(row[\"coordinates\"])\n                if xy is None or sop not in sop2z: continue\n                z = sop2z[sop]; x = float(xy[0]); y = float(xy[1])\n                pts.append((z,y,x))\n\n            vol_iso, spacing_iso = resample_isotropic(vol, spacing, CFG.target_spacing_mm)\n            dz,dy,dx = spacing\n            tz,ty,tx = dz/CFG.target_spacing_mm, dy/CFG.target_spacing_mm, dx/CFG.target_spacing_mm\n            pts_iso = [(z*tz, y*ty, x*tx) for (z,y,x) in pts]\n\n            self.cache[uid] = dict(vol=vol_iso, spacing=spacing_iso)\n            self.points_by_series[uid] = pts_iso\n\n        # sample pools\n        self.pos_samples = []\n        self.neg_samples = []\n        for uid in self.series_ids:\n            if uid not in self.cache: continue\n            for p in self.points_by_series.get(uid, []):\n                self.pos_samples.append((uid, p))\n            for _ in range(max(1, len(self.points_by_series.get(uid, [])))):\n                self.neg_samples.append(uid)\n\n        self.total_len = len(self.pos_samples) + len(self.neg_samples)\n\n    def __len__(self): return max(1, self.total_len)\n\n    def __getitem__(self, idx):\n        if random.random() < CFG.pos_fraction and self.pos_samples:\n            uid, pt = random.choice(self.pos_samples)\n            vol = self.cache[uid][\"vol\"]\n            jitter = np.array([np.random.uniform(-6,6), np.random.uniform(-12,12), np.random.uniform(-12,12)])\n            center = np.clip(np.array(pt)+jitter, [0,0,0], np.array(vol.shape)-1).astype(int)\n            patch, origin_vox = crop_center(vol, tuple(center), CFG.patch_size)\n            pz,py,px = pt; oz,oy,ox = origin_vox\n            heat = make_heatmap(CFG.patch_size, [(pz-oz, py-oy, px-ox)], CFG.gauss_sigma_vox)\n        else:\n            uid = random.choice(self.series_ids)\n            if uid not in self.cache: return self.__getitem__(idx)\n            vol = self.cache[uid][\"vol\"]; Z,Y,X = vol.shape\n            pts = self.points_by_series.get(uid, [])\n            tries=0\n            while True:\n                cz = np.random.randint(16, max(17,Z-16))\n                cy = np.random.randint(16, max(17,Y-16))\n                cx = np.random.randint(16, max(17,X-16))\n                if not pts: break\n                dmin = min(np.linalg.norm(np.array([cz,cy,cx]) - np.array(p)) for p in pts)\n                if dmin >= CFG.neg_distance_min or tries>30: break\n                tries += 1\n            patch, origin_vox = crop_center(vol, (cz,cy,cx), CFG.patch_size)\n            heat = np.zeros(CFG.patch_size, np.float32)\n\n        x = torch.from_numpy(patch[None, ...])  # [1,Z,Y,X]\n        y = torch.from_numpy(heat[None, ...])   # [1,Z,Y,X]\n        return x, y\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 4.1 Using a different function of dataloader","metadata":{}},{"cell_type":"markdown","source":"# 5. Defining the Model","metadata":{}},{"cell_type":"code","source":"# --- Lazy dataset that loads/resamples a series ONLY when it's first sampled ---\nfrom functools import lru_cache\n\n#\ndef resolve_sop_to_z(sop: str, sop2z: dict):\n    \"\"\"Exact match if present; else pick a frame of that SOP (middle by z)\"\"\"\n    if sop in sop2z:\n        return sop2z[sop]\n    cands = [(k, v) for k, v in sop2z.items() if k.startswith(sop + \"|\")]\n    if not cands:\n        return None\n    cands.sort(key=lambda kv: kv[1])\n    return cands[len(cands)//2][1]\n\n\nclass AneurysmPatchDatasetLazy(Dataset):\n    def __init__(self, series_ids: list, train_df: pd.DataFrame, loc_df: pd.DataFrame):\n        self.series_ids = series_ids\n        self.train_df = train_df\n        self.loc_df = loc_df\n\n        # Minimal index only (no DICOM I/O here)\n        self.points_table = (\n            loc_df[loc_df[\"SeriesInstanceUID\"].isin(series_ids)]\n            [[\"SeriesInstanceUID\",\"SOPInstanceUID\",\"coordinates\"]]\n            .copy()\n        )\n        # count positives per series\n        self.pos_counts = self.points_table.groupby(\"SeriesInstanceUID\").size().to_dict()\n        self.series_pool = list(series_ids)\n\n    @lru_cache(maxsize=256)\n    \n   \n    # In AneurysmPatchDatasetLazy\n    def _load_series(self, uid):\n        # 1) Build stack (handles multi-frame/color in your updated build_stack)\n        try:\n            vol, sops, spacing, origin, dircos = build_stack(uid, CFG.series_dir)\n            if vol is None:\n                return None\n        except Exception as e:\n            print(f\"[warn] build_stack failed for {uid}: {e}\")\n            return None\n    \n        # 2) Resample to isotropic\n        #    IMPORTANT: if you use DataLoader(num_workers>0), do NOT call CUDA here.\n        #    Either:\n        #       - keep CPU resampling (safe with workers), OR\n        #       - set num_workers=0 and use the GPU version.\n        USE_CPU_IN_WORKERS = False\n        if USE_CPU_IN_WORKERS:\n            vol_iso, spacing_iso = resample_isotropic_torch(vol, spacing, CFG.target_spacing_mm, device=\"cpu\")\n        else:\n            vol_iso, spacing_iso = resample_isotropic_gpu(vol, spacing, CFG.target_spacing_mm)  # requires num_workers=0\n    \n        # 3) Map localizer SOP -> z (tolerant to multi-frame with SOP|frame)\n        sop2z = {s: z for z, s in enumerate(sops)}\n        pts = []\n        sub = self.points_table[self.points_table[\"SeriesInstanceUID\"] == uid]\n        for _, r in sub.iterrows():\n            sop = str(r[\"SOPInstanceUID\"])\n            xy = parse_coord(r[\"coordinates\"])\n            if xy is None:\n                continue\n            z = resolve_sop_to_z(sop, sop2z)\n            if z is None:\n                continue\n            x, y = float(xy[0]), float(xy[1])\n            pts.append((z, y, x))\n    \n        # 4) Scale points to the iso grid (same factors we applied to Z/Y/X)\n        dz, dy, dx = spacing\n        tz, ty, tx = dz/CFG.target_spacing_mm, dy/CFG.target_spacing_mm, dx/CFG.target_spacing_mm\n        pts_iso = [(z*tz, y*ty, x*tx) for (z, y, x) in pts]\n    \n        return dict(vol=vol_iso, points=pts_iso)\n\n    def __len__(self):\n        total_pos = sum(self.pos_counts.get(uid, 0) for uid in self.series_ids)\n        return max(2*total_pos, 2000)  # arbitrary large epoch length\n\n    def __getitem__(self, idx):\n        use_pos = (np.random.rand() < CFG.pos_fraction)\n\n        if use_pos and any(self.pos_counts.get(u,0)>0 for u in self.series_ids):\n            uid = np.random.choice([u for u in self.series_ids if self.pos_counts.get(u,0)>0])\n            cache = self._load_series(uid)\n            if (cache is None) or (not cache[\"points\"]):\n                return self.__getitem__(idx)  # retry\n            vol = cache[\"vol\"]; pts = cache[\"points\"]\n            pz,py,px = pts[np.random.randint(len(pts))]\n            jitter = np.array([np.random.uniform(-6,6), np.random.uniform(-12,12), np.random.uniform(-12,12)])\n            center = np.clip(np.array([pz,py,px]) + jitter, [0,0,0], np.array(vol.shape)-1).astype(int)\n            patch, (oz,oy,ox) = crop_center(vol, tuple(center), CFG.patch_size)\n            heat = make_heatmap(CFG.patch_size, [(pz-oz, py-oy, px-ox)], CFG.gauss_sigma_vox)\n        else:\n            uid = np.random.choice(self.series_pool)\n            cache = self._load_series(uid)\n            if cache is None:\n                return self.__getitem__(idx)  # retry\n            vol = cache[\"vol\"]; pts = cache[\"points\"]\n            Z,Y,X = vol.shape\n            tries = 0\n            while True:\n                cz = np.random.randint(16, max(17, Z-16))\n                cy = np.random.randint(16, max(17, Y-16))\n                cx = np.random.randint(16, max(17, X-16))\n                if not pts: break\n                dmin = min(np.linalg.norm(np.array([cz,cy,cx]) - np.array(p)) for p in pts)\n                if dmin >= CFG.neg_distance_min or tries > 30: break\n                tries += 1\n            patch, _ = crop_center(vol, (cz,cy,cx), CFG.patch_size)\n            heat = np.zeros(CFG.patch_size, np.float32)\n\n        x = torch.from_numpy(patch[None, ...])   # [1,Z,Y,X]\n        y = torch.from_numpy(heat[None, ...])    # [1,Z,Y,X]\n        return x, y\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ConvBNReLU(nn.Module):\n    def __init__(self, in_ch, out_ch, k=3, s=1, p=1):\n        super().__init__()\n        self.conv = nn.Conv3d(in_ch, out_ch, k, s, p, bias=False)\n        self.bn = nn.BatchNorm3d(out_ch)\n        self.act = nn.ReLU(inplace=True)\n    def forward(self, x): return self.act(self.bn(self.conv(x)))\n\nclass Down(nn.Module):\n    def __init__(self, in_ch, out_ch):\n        super().__init__()\n        self.pool = nn.MaxPool3d(2)\n        self.block = nn.Sequential(ConvBNReLU(in_ch, out_ch), ConvBNReLU(out_ch, out_ch))\n    def forward(self, x): return self.block(self.pool(x))\n\nclass Up(nn.Module):\n    def __init__(self, in_ch, out_ch):\n        super().__init__()\n        self.up = nn.ConvTranspose3d(in_ch, out_ch, 2, stride=2)\n        self.block = nn.Sequential(ConvBNReLU(out_ch*2, out_ch), ConvBNReLU(out_ch, out_ch))\n    def forward(self, x, skip):\n        x = self.up(x)\n        dz = skip.shape[2]-x.shape[2]; dy = skip.shape[3]-x.shape[3]; dx = skip.shape[4]-x.shape[4]\n        x = F.pad(x, (0,dx,0,dy,0,dz))\n        x = torch.cat([skip, x], dim=1)\n        return self.block(x)\n\"\"\"\nclass UNet3D(nn.Module):\n    def __init__(self, base=16):\n        super().__init__()\n        self.inc  = nn.Sequential(ConvBNReLU(1, base), ConvBNReLU(base, base))\n        self.d1   = Down(base, base*2)\n        self.d2   = Down(base*2, base*4)\n        self.d3   = Down(base*4, base*8)\n        self.bot  = nn.Sequential(ConvBNReLU(base*8, base*16), ConvBNReLU(base*16, base*16))\n        self.u3   = Up(base*16, base*8)\n        self.u2   = Up(base*8, base*4)\n        self.u1   = Up(base*4, base*2)\n        self.outc = nn.Conv3d(base*2, 1, 1)\n    def forward(self, x):\n        x1 = self.inc(x); x2 = self.d1(x1); x3 = self.d2(x2); x4 = self.d3(x3)\n        xb = self.bot(x4)\n        x  = self.u3(xb, x4); x = self.u2(x, x3); x = self.u1(x, x2)\n        return self.outc(x)\n\"\"\"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- replace your UNet3D with this full-res version ---\nclass UNet3D(nn.Module):\n    def __init__(self, base=16):\n        super().__init__()\n        self.inc  = nn.Sequential(ConvBNReLU(1, base), ConvBNReLU(base, base))\n        self.d1   = Down(base, base*2)   # /2\n        self.d2   = Down(base*2, base*4) # /4\n        self.d3   = Down(base*4, base*8) # /8\n        self.bot  = nn.Sequential(ConvBNReLU(base*8, base*16), ConvBNReLU(base*16, base*16))\n        self.u3   = Up(base*16, base*8)  # -> /4, skip x4\n        self.u2   = Up(base*8,  base*4)  # -> /2, skip x3\n        self.u1   = Up(base*4,  base*2)  # -> /1, skip x2? (still /1 only if we add another Up)\n        self.u0   = Up(base*2,  base)    # NEW: -> full res, skip x1\n        self.outc = nn.Conv3d(base, 1, 1)\n\n    def forward(self, x):\n        x1 = self.inc(x)   # full\n        x2 = self.d1(x1)   # /2\n        x3 = self.d2(x2)   # /4\n        x4 = self.d3(x3)   # /8\n        xb = self.bot(x4)\n        x  = self.u3(xb, x4)  # /4\n        x  = self.u2(x,  x3)  # /2\n        x  = self.u1(x,  x2)  # /1? actually back to /2→ need the next up to reach full\n        x  = self.u0(x,  x1)  # full resolution\n        return self.outc(x)   # [B,1,Z,Y,X] same as input\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 6. Data Split and Data loaders","metadata":{}},{"cell_type":"code","source":"import torch.nn.functional as F\n\ndef resample_isotropic_gpu(vol_zyx: np.ndarray, spacing, target_mm: float):\n    dz, dy, dx = spacing\n    Z, Y, X = vol_zyx.shape\n    scale = np.array([dz/target_mm, dy/target_mm, dx/target_mm], dtype=np.float32)\n    Zt, Yt, Xt = np.maximum(1, np.round([Z, Y, X] * scale)).astype(int)\n\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    t = torch.from_numpy(vol_zyx).to(device=device, dtype=torch.float32).unsqueeze(0).unsqueeze(0)\n    t_iso = F.interpolate(t, size=(int(Zt), int(Yt), int(Xt)), mode=\"trilinear\", align_corners=False)\n    vol_iso = t_iso.squeeze(0).squeeze(0).detach().cpu().numpy().astype(np.float32)\n    return vol_iso, (target_mm, target_mm, target_mm)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\ndef fast_has_any_file(series_dir, uid):\n    folder = os.path.join(series_dir, uid)\n    if not os.path.isdir(folder):\n        return False\n    try:\n        with os.scandir(folder) as it:\n            for e in it:\n                if e.is_file():\n                    return True\n    except Exception:\n        pass\n    return False\n\n# series that have localizers and at least one file in the folder\nuids = loc_df[\"SeriesInstanceUID\"].unique().tolist()\nuids = [u for u in uids if fast_has_any_file(CFG.series_dir, u)]\n\nprint(\"Total usable series:\", len(uids))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"set_seed(CFG.seed)\n\n# Load CSVs first\ntrain_df = pd.read_csv(CFG.train_csv)\nloc_df   = pd.read_csv(CFG.local_csv)\n\n# Filter UIDs (fast, no DICOM decoding)\nuids = loc_df[\"SeriesInstanceUID\"].unique().tolist()\nuids = [u for u in uids if fast_has_any_file(CFG.series_dir, u)]\nprint(\"Total usable series:\", len(uids))\n\n# Split\nrng = np.random.default_rng(CFG.seed)\nval_n = max(1, int(0.1 * len(uids)))\nval_idx = set(rng.choice(len(uids), size=val_n, replace=False).tolist())\ntrain_ids = [u for i,u in enumerate(uids) if i not in val_idx]\nval_ids   = [u for i,u in enumerate(uids) if i in val_idx]\nprint(\"Train:\", len(train_ids), \"Val:\", len(val_ids))\n\n# Datasets (lazy)\nds_train = AneurysmPatchDatasetLazy(train_ids, train_df, loc_df)\nds_val   = AneurysmPatchDatasetLazy(val_ids,   train_df, loc_df)\n\n# Start with workers=0 to ensure stability; then try 2 if it's smooth\nloader_tr = DataLoader(ds_train, batch_size=CFG.batch_size, shuffle=True,\n                       num_workers=0, pin_memory=True, drop_last=True, persistent_workers=False)\nloader_va = DataLoader(ds_val,   batch_size=CFG.batch_size, shuffle=False,\n                       num_workers=0, pin_memory=True, drop_last=False, persistent_workers=False)\n\n# Smoke test: try to fetch one batch before kicking off training\nxb, yb = next(iter(loader_tr))\nprint(\"Batch OK:\", xb.shape, yb.shape)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# (optional) once stable, try a couple of workers for better throughput\n#loader_tr = DataLoader(ds_train, batch_size=CFG.batch_size, shuffle=True,\n                       #num_workers=2, pin_memory=True, drop_last=True, persistent_workers=True)\n#loader_va = DataLoader(ds_val,   batch_size=CFG.batch_size, shuffle=False,\n #                      num_workers=2, pin_memory=True, drop_last=False, persistent_workers=True)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"loader_tr = DataLoader(ds_train, batch_size=CFG.batch_size, shuffle=True,\n                       num_workers=0, pin_memory=True, drop_last=True, persistent_workers=False)\nloader_va = DataLoader(ds_val,   batch_size=CFG.batch_size, shuffle=False,\n                       num_workers=0, pin_memory=True, drop_last=False, persistent_workers=False)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 7. Training","metadata":{}},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(device)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!nvidia-smi\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- training setup (same cell where you define model/opt/loss) ---\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel = UNet3D(base=16).to(device)\n\nopt = torch.optim.AdamW(model.parameters(), lr=CFG.lr, weight_decay=CFG.weight_decay)\n\n# NEW: use torch.amp.* API (replaces torch.cuda.amp.*)\nscaler = torch.amp.GradScaler('cuda', enabled=CFG.amp)\n\ndef loss_fn(logits, target):\n    # If you kept the smaller-output model, upsample logits to target size:\n    # pred = torch.sigmoid(logits)\n    # if pred.shape[-3:] != target.shape[-3:]:\n    #     pred = F.interpolate(pred, size=target.shape[-3:], mode=\"trilinear\", align_corners=False)\n    # return F.mse_loss(pred, target)\n    return F.mse_loss(torch.sigmoid(logits), target)  # if your U-Net now outputs full res\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def print_gpu_mem(tag=\"\"):\n    if torch.cuda.is_available():\n        print(f\"{tag} alloc (MB): {torch.cuda.memory_allocated()/1024**2:.1f} | \"\n              f\"reserved (MB): {torch.cuda.memory_reserved()/1024**2:.1f}\")\n\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ===========================\n# TRAINING (safe, resumable)\n# ===========================\nimport time, datetime, random, json, os\nfrom dataclasses import asdict\nfrom collections import deque\nfrom pathlib import Path\nimport numpy as np\nimport torch\nimport torch.nn.functional as F\nfrom tqdm import tqdm\n\n# --- helpers ---\ndef fmt_time(s):  # seconds -> H:MM:SS\n    return str(datetime.timedelta(seconds=int(max(0, s))))\n\ndef print_gpu_mem(tag=\"GPU mem —\"):\n    if torch.cuda.is_available():\n        print(f\"{tag} alloc (MB): {torch.cuda.memory_allocated()/1024**2:.1f} | \"\n              f\"reserved (MB): {torch.cuda.memory_reserved()/1024**2:.1f}\")\n\ndef save_ckpt(path: Path, model, opt, scaler, epoch: int, cfg_instance):\n    path.parent.mkdir(parents=True, exist_ok=True)\n    ckpt = {\n        \"model\":     model.state_dict(),\n        \"optimizer\": opt.state_dict(),\n        \"scaler\":    scaler.state_dict() if scaler is not None else None,\n        \"epoch\":     epoch,\n        \"cfg\":       asdict(cfg_instance),  # <- SERIALIZABLE\n        \"rng\": {\n            \"torch\":  torch.get_rng_state(),\n            \"cuda\":   torch.cuda.get_rng_state_all() if torch.cuda.is_available() else None,\n            \"numpy\":  np.random.get_state(),\n            \"python\": random.getstate(),\n        },\n    }\n    torch.save(ckpt, str(path))\n\ndef load_resume(path: Path, model, opt, scaler, device):\n    ckpt = torch.load(str(path), map_location=device)\n    model.load_state_dict(ckpt[\"model\"])\n    if \"optimizer\" in ckpt and ckpt[\"optimizer\"] is not None:\n        opt.load_state_dict(ckpt[\"optimizer\"])\n    if \"scaler\" in ckpt and ckpt[\"scaler\"] is not None and scaler is not None:\n        try:\n            scaler.load_state_dict(ckpt[\"scaler\"])\n        except Exception:\n            pass  # ok if version mismatch\n    # restore RNG (optional)\n    try:\n        torch.set_rng_state(ckpt[\"rng\"][\"torch\"])\n        if torch.cuda.is_available() and ckpt[\"rng\"][\"cuda\"] is not None:\n            torch.cuda.set_rng_state_all(ckpt[\"rng\"][\"cuda\"])\n        np.random.set_state(ckpt[\"rng\"][\"numpy\"])\n        random.setstate(ckpt[\"rng\"][\"python\"])\n    except Exception:\n        pass\n    start_epoch = int(ckpt.get(\"epoch\", 0)) + 1\n    return start_epoch\n\n# --- config / setup ---\ncfg = CFG()  # make an INSTANCE (important for asdict)\nOUT_DIR = Path(cfg.out_dir)\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nmodel  = UNet3D(base=16).to(device)\nopt    = torch.optim.AdamW(model.parameters(), lr=cfg.lr, weight_decay=cfg.weight_decay)\nscaler = torch.amp.GradScaler('cuda', enabled=cfg.amp)\n\n# cuDNN speedup (ok for fixed-size patches)\ntorch.backends.cudnn.benchmark = True\n\n# --- (optional) resume ---\nRESUME_PATH = None  # e.g., Path(OUT_DIR/\"best.pt\") or OUT_DIR/\"ckpt_epoch003.pt\"\nstart_epoch = 1\nif RESUME_PATH:\n    start_epoch = load_resume(Path(RESUME_PATH), model, opt, scaler, device)\n    print(f\"Resuming from epoch {start_epoch}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ===========================\n# TRAINING (safe, resumable)\n# ===========================\nimport time, datetime, random, json, os\nfrom dataclasses import asdict\nfrom collections import deque\nfrom pathlib import Path\nimport numpy as np\nimport torch\nimport torch.nn.functional as F\nfrom tqdm import tqdm\n\n# --- helpers ---\ndef fmt_time(s):  # seconds -> H:MM:SS\n    return str(datetime.timedelta(seconds=int(max(0, s))))\n\ndef print_gpu_mem(tag=\"GPU mem —\"):\n    if torch.cuda.is_available():\n        print(f\"{tag} alloc (MB): {torch.cuda.memory_allocated()/1024**2:.1f} | \"\n              f\"reserved (MB): {torch.cuda.memory_reserved()/1024**2:.1f}\")\n\ndef save_ckpt(path: Path, model, opt, scaler, epoch: int, cfg_instance):\n    path.parent.mkdir(parents=True, exist_ok=True)\n    ckpt = {\n        \"model\":     model.state_dict(),\n        \"optimizer\": opt.state_dict(),\n        \"scaler\":    scaler.state_dict() if scaler is not None else None,\n        \"epoch\":     epoch,\n        \"cfg\":       asdict(cfg_instance),  # <- SERIALIZABLE\n        \"rng\": {\n            \"torch\":  torch.get_rng_state(),\n            \"cuda\":   torch.cuda.get_rng_state_all() if torch.cuda.is_available() else None,\n            \"numpy\":  np.random.get_state(),\n            \"python\": random.getstate(),\n        },\n    }\n    torch.save(ckpt, str(path))\n\ndef load_resume(path: Path, model, opt, scaler, device):\n    ckpt = torch.load(str(path), map_location=device)\n    model.load_state_dict(ckpt[\"model\"])\n    if \"optimizer\" in ckpt and ckpt[\"optimizer\"] is not None:\n        opt.load_state_dict(ckpt[\"optimizer\"])\n    if \"scaler\" in ckpt and ckpt[\"scaler\"] is not None and scaler is not None:\n        try:\n            scaler.load_state_dict(ckpt[\"scaler\"])\n        except Exception:\n            pass  # ok if version mismatch\n    # restore RNG (optional)\n    try:\n        torch.set_rng_state(ckpt[\"rng\"][\"torch\"])\n        if torch.cuda.is_available() and ckpt[\"rng\"][\"cuda\"] is not None:\n            torch.cuda.set_rng_state_all(ckpt[\"rng\"][\"cuda\"])\n        np.random.set_state(ckpt[\"rng\"][\"numpy\"])\n        random.setstate(ckpt[\"rng\"][\"python\"])\n    except Exception:\n        pass\n    start_epoch = int(ckpt.get(\"epoch\", 0)) + 1\n    return start_epoch\n\n# --- config / setup ---\ncfg = CFG()  # make an INSTANCE (important for asdict)\nOUT_DIR = Path(cfg.out_dir)\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nmodel  = UNet3D(base=16).to(device)\nopt    = torch.optim.AdamW(model.parameters(), lr=cfg.lr, weight_decay=cfg.weight_decay)\nscaler = torch.amp.GradScaler('cuda', enabled=cfg.amp)\n\n# cuDNN speedup \ntorch.backends.cudnn.benchmark = True\n\n# --- (optional) resume ---\nRESUME_PATH = None  # e.g., Path(OUT_DIR/\"best.pt\") or OUT_DIR/\"ckpt_epoch003.pt\"\nstart_epoch = 1\nif RESUME_PATH:\n    start_epoch = load_resume(Path(RESUME_PATH), model, opt, scaler, device)\n    print(f\"Resuming from epoch {start_epoch}\")\n\n# --- training loop ---\nbest_val = float('inf')\navg_win  = 50  # moving average window for ETA\nfor epoch in range(start_epoch, cfg.epochs + 1):\n    # ----- TRAIN -----\n    model.train(); tr_losses = []\n    step_times = deque(maxlen=avg_win)\n    pbar = tqdm(enumerate(loader_tr, 1), total=len(loader_tr), desc=f\"Epoch {epoch}/{cfg.epochs} [train]\")\n\n    for step, (x, y) in pbar:\n        t0 = time.perf_counter()\n\n        x = x.to(device, non_blocking=True)\n        y = y.to(device, non_blocking=True)\n\n        opt.zero_grad(set_to_none=True)\n\n        with torch.amp.autocast('cuda', enabled=cfg.amp):\n            logits = model(x)\n            loss   = loss_fn(logits, y)\n\n        scaler.scale(loss).backward()\n        scaler.step(opt)\n        scaler.update()\n\n        tr_losses.append(loss.item())\n\n        # ETA computation\n        step_times.append(time.perf_counter() - t0)\n        avg_step = sum(step_times)/len(step_times)\n        steps_left_epoch = len(loader_tr) - step\n        eta_ep  = avg_step * steps_left_epoch\n        eta_all = eta_ep + avg_step * len(loader_tr) * (cfg.epochs - epoch)\n\n        pbar.set_postfix(loss=float(np.mean(tr_losses)),\n                         eta_ep=fmt_time(eta_ep),\n                         eta_all=fmt_time(eta_all))\n\n        # Optional GPU mem print (every 100 steps)\n        if torch.cuda.is_available() and (step % 100 == 0):\n            print_gpu_mem()\n\n    # ----- VALID -----\n    model.eval(); va_losses = []\n    with torch.no_grad():\n        for x, y in tqdm(loader_va, desc=f\"Epoch {epoch}/{cfg.epochs} [valid]\"):\n            x = x.to(device, non_blocking=True)\n            y = y.to(device, non_blocking=True)\n            with torch.amp.autocast('cuda', enabled=cfg.amp):\n                logits = model(x)\n                va_losses.append(loss_fn(logits, y).item())\n\n    va = float(np.mean(va_losses)) if va_losses else float('nan')\n    print(f\"epoch {epoch}: train={np.mean(tr_losses):.4f}  valid={va:.4f}\")\n\n    # ----- CHECKPOINTS -----\n    try:\n        save_ckpt(OUT_DIR / f\"ckpt_epoch{epoch:03d}.pt\", model, opt, scaler, epoch, cfg)\n        if va < best_val:\n            best_val = va\n            save_ckpt(OUT_DIR / \"best.pt\", model, opt, scaler, epoch, cfg)\n            print(\"  -> saved best.pt\")\n    except TypeError as e:\n        # ultra-safe fallback if something non-serializable sneaks in\n        print(\"[warn] checkpoint save failed, saving weights only:\", e)\n        torch.save(model.state_dict(), str(OUT_DIR / f\"weights_epoch{epoch:03d}.pt\"))\n        if va < best_val:\n            best_val = va\n            torch.save(model.state_dict(), str(OUT_DIR / \"best_weights.pt\"))\n            print(\"  -> saved best_weights.pt\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#import os, torch\n#os.makedirs(OUT_DIR, exist_ok=True)\n#safe_path = str(OUT_DIR / \"best.pt\")  \n#torch.save(model.state_dict(), safe_path)\n#print(\"Saved weights to\", safe_path)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 8. Inference","metadata":{}},{"cell_type":"code","source":"#print(folder_path)\n#print(OUT_DIR)\n## I have to do this because the code broke in the second epoch, and I got the best.pt from the first epoch \n#!ls /kaggle/input/unet_3d/pytorch/default/1\n#!cp /kaggle/input/unet_3d/pytorch/default/1/best.pt /kaggle/working/out_3dunet/","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#CKPT_PATH = \"/kaggle/input/unet_3d/pytorch/default/1/best.pt\"\nCKPT_PATH = \"/kaggle/working/out_3dunet/best.pt\"\n\n# 1) recreate the exact model\nmodel = UNet3D(base=16)\n\n# 2) robust checkpoint loading (handles plain state_dict, dict with 'model', and DDP 'module.' prefix)\ndef load_ckpt_into(model, path):\n    ckpt = torch.load(path, map_location=\"cpu\", weights_only=False)  # <-- key change\n    # cases:\n    if isinstance(ckpt, dict) and \"state_dict\" in ckpt:\n        state = ckpt[\"state_dict\"]\n    elif isinstance(ckpt, dict) and \"model\" in ckpt:\n        state = ckpt[\"model\"]\n    else:\n        state = ckpt  # assume plain state_dict\n\n    # strip \"module.\" if saved from DDP\n    if any(k.startswith(\"module.\") for k in state.keys()):\n        state = {k.replace(\"module.\", \"\", 1): v for k, v in state.items()}\n\n    missing, unexpected = model.load_state_dict(state, strict=False)\n    print(f\"Loaded weights. missing={len(missing)}, unexpected={len(unexpected)}\")\n    if missing:   print(\"  missing:\", missing[:5], \"...\")\n    if unexpected: print(\"  unexpected:\", unexpected[:5], \"...\")\n    return model\n\nmodel = load_ckpt_into(model, CKPT_PATH).to(device).eval()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(folder_path)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# Full Inference + Submission\n# =========================\nimport os, glob, numpy as np, pandas as pd, torch\nfrom tqdm import tqdm\n\n# --------- 0) Sanity checks (do not redefine CFG) ---------\ntry:\n    _ = CFG.series_dir\n    _ = CFG.patch_size\n    _ = CFG.sw_overlap\nexcept Exception as e:\n    raise RuntimeError(\"CFG is not the original dataclass. Re-run the cell where you defined it.\") from e\n\ntry:\n    _ = model.eval()\nexcept NameError:\n    raise RuntimeError(\"`model` is not defined. Re-run the cell that loads best.pt into `model`.\")\n\ntry:\n    device\nexcept NameError:\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# --------- 1) Find test IDs and the correct series directory ---------\nROOT = \"/kaggle/input/rsna-intracranial-aneurysm-detection\"\nEVAL_ROOT   = os.path.join(ROOT, \"kaggle_evaluation\")\nEVAL_SERIES = os.path.join(EVAL_ROOT, \"series\")\nEVAL_TEST   = os.path.join(EVAL_ROOT, \"test.csv\")\nOFF_SERIES  = os.path.join(ROOT, \"test_images\")\nOFF_SAMPLE  = os.path.join(ROOT, \"sample_submission.csv\")\n\nif os.path.isdir(EVAL_SERIES) and os.path.isfile(EVAL_TEST):\n    mode = \"eval_pack\"\n    SERIES_DIR = EVAL_SERIES\n    df_ids = pd.read_csv(EVAL_TEST)\n    id_candidates = [\"id\",\"series_id\",\"SeriesInstanceUID\",\"seriesinstanceuid\"]\n    id_col = next((c for c in df_ids.columns\n                   if c in id_candidates or c.lower() in [s.lower() for s in id_candidates]), None)\n    assert id_col is not None, f\"Could not find an ID column in {EVAL_TEST} (columns={list(df_ids.columns)})\"\n    target_col = \"aneurysm\"           # required by RSNA scorer\n    test_ids = df_ids[id_col].tolist()\nelse:\n    mode = \"official_test\"\n    SERIES_DIR = OFF_SERIES\n    sample = pd.read_csv(OFF_SAMPLE)\n    id_col = [c for c in sample.columns if c.lower() in (\"id\",\"series_id\",\"seriesinstanceuid\")][0]\n    target_col = [c for c in sample.columns if c != id_col][0]\n    test_ids = sample[id_col].tolist()\n\nprint(f\"Mode: {mode}\")\nprint(f\"Series dir: {SERIES_DIR}\")\nprint(f\"id_col: {id_col}  | target_col: {target_col}  | #test={len(test_ids)}\")\n\n# --------- 2) Heatmap -> single probability helper ---------\ndef series_score_from_heatmap(heat: np.ndarray, method=\"topk_mean\", k=20):\n    if method == \"max\":\n        return float(heat.max())\n    if method == \"percentile\":\n        return float(np.percentile(heat, 99.9))\n    if method == \"mean\":\n        return float(heat.mean())\n    flat = heat.ravel()\n    k = min(k, flat.size)\n    return float(np.partition(flat, -k)[-k:].mean())\n\n# --------- 3) Inference loop ---------\ntorch.backends.cudnn.benchmark = True\ntorch.backends.cuda.matmul.allow_tf32 = True\ntorch.backends.cudnn.allow_tf32 = True\n\npreds = []\nfor uid in tqdm(test_ids, desc=\"Inference\"):\n    # load series -> volume (Z,Y,X), spacing (dz,dy,dx)\n    vol, sops, spacing, origin, dircos = build_stack(uid, SERIES_DIR)\n    if vol is None:\n        preds.append(0.0)\n        continue\n\n    # resample to the spacing used in training\n    vol_iso, _ = resample_isotropic(vol, spacing, CFG.target_spacing_mm)\n\n    # sliding-window prediction (model should return logits; SW applies sigmoid)\n    heat = sliding_window_predict(\n        model=model,\n        vol=vol_iso.astype(np.float32),\n        patch=CFG.patch_size,          # (Z,Y,X) from your CFG\n        overlap=CFG.sw_overlap,        # from your CFG\n        device=device,\n        amp=CFG.amp\n        #batch_chunks=2,                # try 2–4 for extra speed if memory allows\n    )\n\n    \n\n    # reduce heatmap to one probability per series\n    prob = series_score_from_heatmap(heat, method=\"topk_mean\", k=max(1, CFG.topk))\n    preds.append(prob)\n\n# --------- 4) Build submission in the required format ---------\nOUT_CSV = \"/kaggle/working/submission.csv\"\n\nif mode == \"eval_pack\":\n    sub = pd.DataFrame({id_col: test_ids, \"aneurysm\": np.clip(preds, 0.0, 1.0)})\nelse:\n    sub = sample.copy()\n    sub[target_col] = np.clip(preds, 0.0, 1.0)\n\nsub.to_csv(OUT_CSV, index=False)\nprint(sub.head())\nprint(\"✅ Submission saved to:\", OUT_CSV)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# make sure columns are exactly what the comp expects\n# e.g., id column (SeriesInstanceUID/series_id/id) + prediction column ('aneurysm')\nprint(sub.dtypes)  # quick sanity check\n\n# write parquet (required)\nout_parquet = \"/kaggle/working/submission.parquet\"\nsub.to_parquet(out_parquet, index=False)  # uses pyarrow on Kaggle\n\n# (optional) keep csv for your own inspection\nout_csv = \"/kaggle/working/submission.csv\"\nsub.to_csv(out_csv, index=False)\n\nimport os\nprint(\"Files created:\", os.path.exists(out_parquet), os.path.exists(out_csv))\n","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}