{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"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":"nvidiaTeslaT4","dataSources":[{"sourceId":99552,"databundleVersionId":13190393,"sourceType":"competition"},{"sourceId":12637336,"sourceType":"datasetVersion","datasetId":7981664}],"dockerImageVersionId":31090,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **IA Detection: 8-Frame Image Training Pipeline**","metadata":{}},{"cell_type":"markdown","source":"https://www.kaggle.com/code/stpeteishii/ia-detection-8-frame-image-training-pipeline\n\nhttps://www.kaggle.com/code/stpeteishii/ia-detection-8-frame-image-inference-pipeline","metadata":{}},{"cell_type":"markdown","source":"\n---\n\n\n**Introduction:**\nThis notebook implements the **training phase** of an eight-frame medical image classification pipeline.\nIt takes raw DICOM/PNG series as input, processes them into standardized image volumes, and trains a deep learning model using **patient-level cross-validation**.\n\nKey features of this training notebook:\n\n* **Data preparation**: Reads series metadata, extracts 8 representative frames per study, and applies preprocessing (windowing, CLAHE, normalization).\n* **Augmentation**: Uses strong spatial, intensity, and dropout augmentations for robust model learning.\n* **Loss functions**: Combines weighted BCE and focal loss to prioritize rare but critical cases.\n* **Training loop**: Implements mixed-precision training, gradient accumulation, and batch skipping for stability.\n* **Validation**: Tracks per-class AUC and combined scores across folds.\n\n**Purpose:**\nThis notebook is solely focused on **model training and validation**.\nA separate notebook should be used for **inference and submission generation** to maintain a clean workflow and avoid data leakage.\n\n---\n\n\n","metadata":{}},{"cell_type":"code","source":"# ------------------------------------------------------------------\n# Imports & Reproducibility\n# ------------------------------------------------------------------\nimport os\nimport random\nimport warnings\nimport gc\nimport pathlib\nimport numpy as np\nimport pandas as pd\nfrom typing import List, Tuple, Dict\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.optim import AdamW\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\nimport timm\nfrom albumentations import (\n    Compose, RandomBrightnessContrast, Blur, CLAHE,\n    HorizontalFlip, Rotate, ShiftScaleRotate,\n    ElasticTransform, GridDistortion,\n    RandomGamma, GaussNoise, ISONoise,\n    CoarseDropout, Normalize\n)\nfrom albumentations.pytorch import ToTensorV2\nfrom sklearn.model_selection import GroupKFold\nfrom sklearn.metrics import roc_auc_score\nimport pydicom\nimport functools \nfrom tqdm import tqdm\nimport glob\n\nwarnings.filterwarnings(\"ignore\")\nnp.random.seed(42); random.seed(42); torch.manual_seed(42)\n\nimport os\nimport warnings\nfrom dataclasses import dataclass\n\nimport os\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ======================\n# Configuration\n# ======================\nfrom dataclasses import dataclass\n\n@dataclass(frozen=False)\nclass Config:\n    DATA_DIR      : str = \"/kaggle/input/rsna-2025-intracranial-aneurysm-png-224x224\"\n    OUTPUT_DIR    : str = \"/kaggle/working\"  \n    field_with_default: str = \"default_value\"\n    CVT_PNG_DIR   : str = os.path.join(DATA_DIR, \"cvt_png\")\n    SERIES_MAP    : str = os.path.join(DATA_DIR, \"series_index_mapping.csv\")\n    LOCALIZERS    : str = os.path.join(DATA_DIR, \"train_localizers_with_relative.csv\")\n    TRAIN_CSV     : str = \"/kaggle/input/rsna-intracranial-aneurysm-detection/train.csv\"\n\n    NUM_FRAMES    : int = 8\n    IMG_SIZE      : int = 224\n    NUM_CLASSES   : int = 14\n\n    BATCH_SIZE    : int = 1\n    ACCUM_STEPS   : int = 5\n    EPOCHS        : int = 8\n    LR            : float = 5e-5\n    PATIENCE      : int = 3\n\n    NUM_FOLDS      : int = 5\n    NUM_FOLD       : int = 0\n    USE_GROUP_CV   : bool = True\n    USE_CLAMPH     : bool = True\n    USE_STRONG_AUG : bool = True\n    USE_IMPROVED_LOSS : bool = True\n    HIDDEN_DIM     : int = 512\n    CACHE_SIZE     : int = 5000\n    NUM_WORKERS    : int = 4\n    MODEL_NAME_BACKBONE: str = \"efficientnet_b0\"\n    BATCH_SIZE     : int = 8 \n\n    # --- New: weight for the “aneurysm (last) class” --------------------\n    AneurysmWeight : float = 3.0\n\n# Create a *single* configuration instance – frozen dataclass becomes immutable.\nconfig = Config()\n\n# Device helper\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ======================\n# Focal Loss implementation (self‑contained)\n# ======================\nclass FocalLoss(nn.Module):\n    def __init__(self, gamma: float = 2.0, alpha: float = 0.25, reduction: str = \"mean\"):\n        super().__init__()\n        self.gamma = gamma\n        self.alpha = alpha\n        self.reduction = reduction\n\n    def forward(self, logits: torch.Tensor, targets: torch.Tensor) -> torch.Tensor:\n        # Apply the sigmoid first → probabilities p in [0,1]\n        probs = torch.sigmoid(logits)\n\n        # Compute the cross‑entropy component\n        ce_loss = F.binary_cross_entropy_with_logits(logits, targets, reduction=\"none\")\n\n        # Modulating factor: (1 - p)^γ for positives, p^γ for negatives\n        p_t = probs * targets + (1 - probs) * (1 - targets)\n        mod_factor = (1 - p_t) ** self.gamma\n\n        # α-balancing factor\n        alpha_factor = self.alpha * targets + (1 - self.alpha) * (1 - targets)\n\n        # Final focal loss\n        focal = alpha_factor * mod_factor * ce_loss\n\n        if self.reduction == \"mean\":\n            return focal.mean()\n        elif self.reduction == \"sum\":\n            return focal.sum()\n        else:  # 'none'\n            return focal","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ======================\n# ImprovedLoss (combines weighted BCE + focal)\n# ======================\nclass ImprovedLoss(nn.Module):\n    def __init__(self,\n                 weight_aneurysm: float | None = None,\n                 focal_weight: float = 0.3):\n        super().__init__()\n\n        # ----------------------\n        # 1) Configuration\n        # ----------------------\n        self.num_classes   = config.NUM_CLASSES\n        self.focal_weight  = focal_weight\n        self.focal = FocalLoss(gamma=2.0, alpha=0.25, reduction=\"mean\")\n\n        # ----------------------\n        # 2) Per‑class weight vector\n        # ----------------------\n        # Start with all ones and then scale the aneurysm class.\n        self.w = torch.ones(self.num_classes, device=device)\n\n        if weight_aneurysm is None:\n            weight_aneurysm = config.AneurysmWeight\n\n        self.w[-1] = weight_aneurysm\n\n    # ------------------------------------------------------------------\n    # forward\n    # ------------------------------------------------------------------\n    def forward(self, out: torch.Tensor, lab: torch.Tensor) -> torch.Tensor:\n        # BCE with logits (reduction='none' → per‑element)\n        bce = F.binary_cross_entropy_with_logits(out, lab, reduction=\"none\")\n        w_bce = (bce * self.w).mean()          # weighted mean over all entries\n\n        # Focal loss (already mean‑aggregated inside FocalLoss)\n        fl = self.focal(out, lab)\n\n        # Blend the two components\n        return (1.0 - self.focal_weight) * w_bce + self.focal_weight * fl","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ======================\n# Simple fallback (weighted BCE only)\n# ======================\ndef simple_weighted_bce(out: torch.Tensor, lab: torch.Tensor) -> torch.Tensor:\n    w = torch.ones(config.NUM_CLASSES, device=device)\n    w[-1] = config.AneurysmWeight\n    return (F.binary_cross_entropy_with_logits(out, lab, reduction='none') * w).mean()\n\n# Create the right criterion – the user chose it via the config flag\ncriterion = ImprovedLoss() if config.USE_IMPROVED_LOSS else simple_weighted_bce\n\n# ------------------------------------------------------------------\n# Load tabular data\n# ------------------------------------------------------------------\ntrain_meta = pd.read_csv(config.TRAIN_CSV)\nseries_map = pd.read_csv(config.SERIES_MAP)\nloc_df    = pd.read_csv(config.LOCALIZERS)\nTARGETS   = [\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    'Aneurysm Present'\n]\n\n# ------------------------------------------------------------------\n# Utility functions\n# ------------------------------------------------------------------\nMOD_W_WP = {\n    'CT'  : (40.,80.),\n    'CTA' : (50.,350.),\n    'MRA' : (600.,1200.),\n    'MRI' : (40.,80.),\n    'MR'  : (40.,80.)\n}\n\ndef window(img, center, width):\n    lo, hi = center-width/2., center+width/2.\n    return ((np.clip(img,lo,hi)-lo)/(hi-lo+1e-7)*255).astype(np.uint8)\n\ndef clahe(img, mod):\n    clip, grid = (3.,(8,8)) if mod in ['CTA','MRA'] else (2.,(8,8))\n    if mod in ['MRI','MR']:\n        img = cv2.cvtColor(img.astype(np.uint8),cv2.COLOR_BGR2GRAY) if img.ndim==3 else img\n        img = cv2.createCLAHE(clipLimit=clip, tileGridSize=grid).apply(img)\n        img = (np.power(img/255.,0.9)*255).astype(np.uint8)\n        return img\n    return cv2.createCLAHE(clipLimit=clip, tileGridSize=grid).apply(img)\n\ndef robust_norm(vol):\n    p1,p99 = np.percentile(vol, [1,99])\n    vol = np.clip(vol,p1,p99)\n    return ((vol-p1)/(p99-p1+1e-7)*255).astype(np.uint8)\n\ndef build_3c(vol):\n    mid  = vol[vol.shape[0]//2]\n    mip  = vol.max(0)\n    stdp = vol.std(0)\n    stdp = ((np.clip(stdp, *np.percentile(stdp,[5,95]))-stdp.min())/(stdp.max()-stdp.min()+1e-7)*255).astype(np.uint8)\n    return np.stack([mid,mip,stdp],2)\n\ndef smart_slice(paths):\n    n = len(paths)\n    if n<=config.NUM_FRAMES:\n        return (paths+paths*config.NUM_FRAMES)[:config.NUM_FRAMES]\n    start = max(0, int(0.1*n))\n    step  = max(1,(n-start)//config.NUM_FRAMES)\n    idxs  = list(range(start, min(n,start+step*config.NUM_FRAMES),step))\n    while len(idxs)<config.NUM_FRAMES:\n        idxs.append(idxs[-1])        # duplicate last\n    return [paths[i] for i in idxs]\n\n@functools.lru_cache(5000)\ndef patient_group(series_uid):\n    dicom_dir = pathlib.Path(f\"/kaggle/input/rsna-intracranial-aneurysm-detection/series/{series_uid}\")\n    if dicom_dir.exists():\n        for f in dicom_dir.glob(\"*.dcm\"):   # first file is enough\n            try:\n                ds = pydicom.dcmread(f, stop_before_pixels=True, force=True)\n                return ds.get('StudyInstanceUID', ds.get('PatientID','X'))\n            except: pass\n    return f\"fallback_{series_uid[:32]}\"\n\n# ------------------------------------------------------------------\n# Build 8‑frame path mapping\n# ------------------------------------------------------------------\ndef build_frame_map():\n    frame_map: Dict[str,List[str]] = {}\n    for uid in tqdm(train_meta['SeriesInstanceUID'].unique(), ascii=True):\n        src = series_map[series_map['SeriesInstanceUID']==uid]\n        if src.empty:\n            frame_map[uid] = []\n            continue\n        # try disease directories\n        found = []\n        row = train_meta[train_meta['SeriesInstanceUID']==uid].iloc[0]\n        for col in TARGETS[:-1]:\n            if row[col]==1:\n                loc = os.path.join(config.CVT_PNG_DIR, col.replace('/','_'), uid)\n                if os.path.isdir(loc):\n                    p = sorted(glob.glob(os.path.join(loc,\"*.png\")))\n                    if p: found=p; break\n        # fallback to raw dict\n        if not found:\n            n = len(src)\n            found = [f\"dummy_path_{i:04d}.png\" for i in range(n)]\n        frame_map[uid] = smart_slice(found)\n    return frame_map\n\nFRAME_MAP = build_frame_map()\nprint(f\"Frame map ready for {len(FRAME_MAP)} series\")\n\n# ------------------------------------------------------------------\n# Cross‑validation split (GroupKFold)\n# ------------------------------------------------------------------\ngroup_ids = [patient_group(u) for u in train_meta['SeriesInstanceUID']]\ntrain_meta['group'] = group_ids\nif config.USE_GROUP_CV:\n    splits = GroupKFold(config.NUM_FOLDS).split(train_meta, train_meta['Aneurysm Present'], train_meta['group'])\nelse:\n    from sklearn.model_selection import StratifiedKFold\n    splits = StratifiedKFold(config.NUM_FOLDS).split(train_meta, train_meta['Aneurysm Present'])\ntrain_idx, val_idx = list(splits)[config.NUM_FOLD]\ntrain_df, val_df = train_meta.iloc[train_idx], train_meta.iloc[val_idx]\nprint(f\"Fold{config.NUM_FOLD}: {len(train_df)} train / {len(val_df)} val\")\n\n# ------------------------------------------------------------------\n# Augmentations\n# ------------------------------------------------------------------\ndef compose_aug(is_train):\n    if is_train and config.USE_STRONG_AUG:\n        aug = Compose([\n            Rotate(limit=15, p=0.7),\n            HorizontalFlip(p=0.5),\n            ShiftScaleRotate(shift_limit=0.1,scale_limit=0.1,rotate_limit=10,p=0.6),\n            ElasticTransform(alpha=50,sigma=5,p=0.3),\n            GridDistortion(num_steps=3,distort_limit=0.1,p=0.3),\n            RandomBrightnessContrast(brightness_limit=0.2,contrast_limit=0.2,p=0.6),\n            CLAHE(clipLimit=2.0,tile_grid_size=(8,8),p=0.5),\n            RandomGamma(gamma_limit=(80,120),p=0.4),\n            GaussNoise(var_limit=(10,80),p=0.4),\n            ISONoise(color_shift=(0.01,0.05),intensity=(0.1,0.5),p=0.3),\n            Blur(blur_limit=3,p=0.2),\n            CoarseDropout(max_holes=8,max_height=32,max_width=32,p=0.3),\n            Normalize(mean=[0.485,0.456,0.406],std=[0.229,0.224,0.225]),\n            ToTensorV2()\n        ])\n    else:\n        aug = Compose([\n            Normalize(mean=[0.485,0.456,0.406],std=[0.229,0.224,0.225]),\n            ToTensorV2()\n        ])\n    return aug\n\ntrain_aug = compose_aug(True)\nval_aug   = compose_aug(False)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ------------------------------------------------------------------\n# Dataset\n# ------------------------------------------------------------------\nclass EightFrameDataset(Dataset):\n    def __init__(self, df, frame_map, augment, is_train=True):\n        self.df     = df.reset_index(drop=True)\n        self.map    = frame_map\n        self.aug    = augment\n        self.is_train = is_train\n        self.cache  = {}\n        self.cache_keys = []\n        self.cache_max = config.CACHE_SIZE\n\n    def __len__(self): return len(self.df)\n\n    #fixed 08-09 01:03\n    def __getitem__(self, idx):\n        if idx in self.cache: \n            return self.cache[idx]\n        \n        row = self.df.iloc[idx]\n        uid = row.SeriesInstanceUID\n        paths = self.map[uid]\n        \n        if paths and paths[0].startswith('dummy_path'):\n            vol = self.dicom_volume(uid, row.Modality, paths)\n        else:\n            vol = self.png_volume(paths)\n        \n        vol = robust_norm(vol)\n        img = build_3c(vol)\n        img = self.aug(image=img)[\"image\"]\n        \n        # FIX: Handle the tensor conversion properly\n        try:\n            # Get target values\n            target_values = row[TARGETS].values\n            \n            # Check if it's an object array and needs conversion\n            if target_values.dtype == 'object':\n                # Convert to numeric, handling non-numeric values\n                import pandas as pd\n                target_series = pd.Series(target_values)\n                target_numeric = pd.to_numeric(target_series, errors='coerce').fillna(0.0)\n                lbl = torch.tensor(target_numeric.values.astype(np.float32), dtype=torch.float32)\n            else:\n                # Already numeric, convert directly\n                lbl = torch.tensor(target_values.astype(np.float32), dtype=torch.float32)\n                \n        except Exception as e:\n            print(f\"Error converting targets at index {idx}: {e}\")\n            print(f\"Target values: {target_values}\")\n            print(f\"Target types: {[type(x) for x in target_values]}\")\n            # Fallback: create zero tensor with correct shape\n            lbl = torch.zeros(len(TARGETS), dtype=torch.float32)\n        \n        meta = torch.tensor([self.age(row), self.sex(row)], dtype=torch.float32)\n        out = (img, lbl, meta)\n        self.add_to_cache(idx, out)\n        return out\n\n    \n\n    def add_to_cache(self,i,item):\n        if len(self.cache)>=self.cache_max:\n            f=self.cache_keys.pop(0); del self.cache[f]\n        self.cache[i]=item; self.cache_keys.append(i)\n\n    @staticmethod\n    def age(row):\n        age = row.get('PatientAge',50)\n        if pd.isna(age): age=50\n        if isinstance(age,str): age=int(''.join(filter(str.isdigit,age[:3])) or 50)\n        return min(float(age)/100.,1.)\n\n    @staticmethod\n    def sex(row):\n        return 1.0 if row.get('PatientSex','M')=='M' else 0.0\n\n    def png_volume(self,paths):\n        imgs=[]\n        for p in paths:\n            try:\n                im = cv2.imread(p,cv2.IMREAD_GRAYSCALE)\n                if im is not None:\n                    im=cv2.resize(im,(config.IMG_SIZE,config.IMG_SIZE),interpolation=cv2.INTER_AREA)\n                else: im=np.zeros((config.IMG_SIZE,config.IMG_SIZE),dtype=np.uint8)\n            except: im=np.zeros((config.IMG_SIZE,config.IMG_SIZE),dtype=np.uint8)\n            imgs.append(im)\n        return np.stack(imgs,0)\n\n\n    def dicom_volume(self, uid, mod, paths):\n        # -----------------------------------------------------------------\n        # 1) Grab the series rows from the global `series_map`\n        # -----------------------------------------------------------------\n        rel_idx = series_map[series_map['SeriesInstanceUID'] == uid]['relative_index'].values\n        rows = (\n            series_map[\n                (series_map['SeriesInstanceUID'] == uid) &\n                (series_map['relative_index'].isin(rel_idx))\n            ]\n            .sort_values('relative_index')\n        )\n    \n        vols = []\n        for _, d in rows.iterrows():\n            try:\n                ds = pydicom.dcmread(d.dicom_filename, force=True)\n                img = ds.pixel_array.astype(np.float32)\n    \n                # Convert color images to grayscale, if needed\n                if img.ndim == 3 and img.shape[-1] == 3:\n                    img = cv2.cvtColor(img.astype(np.uint8), cv2.COLOR_BGR2GRAY).astype(np.float32)\n                elif img.ndim == 3:\n                    img = img[:, :, 0]  # sensible fallback\n    \n                # Apply rescale if the DICOM contains it\n                if hasattr(ds, 'RescaleSlope') and hasattr(ds, 'RescaleIntercept'):\n                    img = img * ds.RescaleSlope + ds.RescaleIntercept\n    \n                # ----------- FIXED PART ------------------------------------\n                cen, wid = MOD_W_WP.get(mod, (40.0, 80.0))\n                img = window(img, cen, wid)\n    \n                # Optional CLAHE (Contrast‑Limited Adaptive Histogram Equalisation)\n                img = clahe(img, mod)\n    \n                # Resize to model input size\n                last = cv2.resize(\n                    img, (config.IMG_SIZE, config.IMG_SIZE),\n                    interpolation=cv2.INTER_AREA\n                )\n                vols.append(last)\n    \n            except Exception as exc:\n                # In production you might want to log the exception.\n                # For now we just use a blank image so that the array shape stays consistent.\n                vols.append(np.zeros((config.IMG_SIZE, config.IMG_SIZE), dtype=np.uint8))\n    \n        # Pad or truncate to the required number of frames (config.NUM_FRAMES)\n        s = vols[:config.NUM_FRAMES]\n        if len(s) < config.NUM_FRAMES:\n            s += vols[:config.NUM_FRAMES - len(s)]\n        return np.stack(s, 0)\n\n\n    \n    def dicom_volume(self,uid,mod,paths):\n        # paths are dummy placeholders – we simply parse series_map\n        rel_idx = series_map[series_map['SeriesInstanceUID']==uid]['relative_index'].values\n        rows = series_map[(series_map['SeriesInstanceUID']==uid)&\n                          (series_map['relative_index'].isin(rel_idx))].sort_values('relative_index')\n        vols=[]\n        for _,d in rows.iterrows():\n            try:\n                ds=pydicom.dcmread(d.dicom_filename,force=True)\n                img=ds.pixel_array.astype(np.float32)\n                if img.ndim==3 and img.shape[-1]==3:\n                    img=cv2.cvtColor(img.astype(np.uint8),cv2.COLOR_BGR2GRAY).astype(np.float32)\n                else:\n                    img=img[:,:,0]\n                if hasattr(ds,'RescaleSlope'): img=img*ds.RescaleSlope+ds.RescaleIntercept\n                cen, wid = MOD_W_WP.get(mod, (40.0, 80.0))        \n                img=window(img,cen,wid)\n                img=clahe(img,mod)\n                last=cv2.resize(img,(config.IMG_SIZE,config.IMG_SIZE),interpolation=cv2.INTER_AREA)\n                vols.append(last)\n            except: vols.append(np.zeros((config.IMG_SIZE,config.IMG_SIZE),dtype=np.uint8))\n        # pad/truncate to 8\n        s=vols[:config.NUM_FRAMES]\n        if len(s)<config.NUM_FRAMES:\n            s+=vols[:config.NUM_FRAMES-len(s)]\n        return np.stack(s,0)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ------------------------------------------------------------------\n# DataLoaders\n# ------------------------------------------------------------------\ntrain_ds = EightFrameDataset(train_df,FRAME_MAP,train_aug,True)\nval_ds   = EightFrameDataset(val_df,FRAME_MAP,val_aug,False)\n\nconfig.BATCH_SIZE = max(2, config.BATCH_SIZE)\n\ntrain_loader = DataLoader(\n    train_ds,\n    batch_size=config.BATCH_SIZE,\n    shuffle=True,\n    num_workers=config.NUM_WORKERS,\n    drop_last=True  # ADD THIS LINE\n)\n\nval_loader = DataLoader(\n    val_ds,\n    batch_size=config.BATCH_SIZE,\n    shuffle=False,\n    num_workers=config.NUM_WORKERS,\n    drop_last=True  # ADD THIS LINE\n)\n\n# IMMEDIATE FIX 2: Modify your training loop to skip size-1 batches\ndef train_one():\n    \"\"\"Modified training function that skips batches with size 1\"\"\"\n    model.train()\n    torch.cuda.empty_cache()\n    epoch_loss = 0.0\n    valid_batches = 0\n\n    for it, (img, lab, meta) in enumerate(train_loader):\n        # SKIP BATCHES WITH SIZE 1\n        if img.size(0) == 1:\n            print(f\"Skipping batch {it} with size 1 to avoid BatchNorm error\")\n            continue\n            \n        # ---- device placement -------------------------------------------------\n        img = img.to(device, non_blocking=True)\n        lab = lab.to(device, non_blocking=True)\n\n        if not isinstance(meta, torch.Tensor):\n            meta = torch.tensor(meta).to(device, non_blocking=True)\n        else:\n            meta = meta.to(device, non_blocking=True)\n\n        # ---- forward/backward ---------------------------------------------------\n        with torch.cuda.amp.autocast():\n            out = model(img, meta)\n            loss = criterion(out, lab) / config.ACCUM_STEPS\n\n        scaler.scale(loss).backward()\n        epoch_loss += loss.item() * config.ACCUM_STEPS\n        valid_batches += 1\n\n        # ---- step when we hit an accumulation boundary -----------------------\n        if valid_batches % config.ACCUM_STEPS == 0 or it == len(train_loader) - 1:\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad()\n\n    return epoch_loss / max(valid_batches, 1)\n\ndef validate_one():\n    \"\"\"Modified validation function that skips batches with size 1\"\"\"\n    model.eval()\n    torch.cuda.empty_cache()\n    epoch_loss = 0.0\n    out_all = []\n    lab_all = []\n\n    with torch.no_grad():\n        for img, lab, meta in val_loader:\n            # SKIP BATCHES WITH SIZE 1\n            if img.size(0) == 1:\n                print(\"Skipping validation batch with size 1\")\n                continue\n                \n            img = img.to(device, non_blocking=True)\n            lab = lab.to(device, non_blocking=True)\n\n            if not isinstance(meta, torch.Tensor):\n                meta = torch.tensor(meta).to(device, non_blocking=True)\n            else:\n                meta = meta.to(device, non_blocking=True)\n\n            with torch.cuda.amp.autocast():\n                out = model(img, meta)\n                loss = criterion(out, lab)\n\n            epoch_loss += loss.item()\n            out_all.append(torch.sigmoid(out).cpu().numpy())\n            lab_all.append(lab.cpu().numpy())\n\n    # ----- metrics ------------------------------------------------------------\n    if len(out_all) == 0:  # Handle case where all batches were skipped\n        return float('inf'), 0.0\n        \n    out_all = np.concatenate(out_all, 0)\n    lab_all = np.concatenate(lab_all, 0)\n\n    aucs = []\n    for i in range(config.NUM_CLASSES):\n        if len(np.unique(lab_all[:, i])) > 1:\n            aucs.append(roc_auc_score(lab_all[:, i], out_all[:, i]))\n        else:\n            aucs.append(0.5)\n\n    auc_mixed = np.mean(aucs[:-1])\n    auc_aneurysm = np.mean(aucs[-1:])\n    score = (auc_mixed + auc_aneurysm) / 2.0\n\n    return epoch_loss / len(out_all), score\n\n# IMMEDIATE FIX 3: Set BatchNorm to eval mode during training\ndef set_batchnorm_eval(model):\n    \"\"\"Set all BatchNorm layers to eval mode\"\"\"\n    for module in model.modules():\n        if isinstance(module, (torch.nn.BatchNorm1d, torch.nn.BatchNorm2d, torch.nn.BatchNorm3d)):\n            module.eval()\n\n# Modified training function using BatchNorm in eval mode\ndef train_one_bn_eval():\n    \"\"\"Training with BatchNorm in eval mode\"\"\"\n    model.train()\n    set_batchnorm_eval(model)  # Keep BatchNorm in eval mode\n    torch.cuda.empty_cache()\n    epoch_loss = 0.0\n\n    for it, (img, lab, meta) in enumerate(train_loader):\n        # ---- device placement -------------------------------------------------\n        img = img.to(device, non_blocking=True)\n        lab = lab.to(device, non_blocking=True)\n\n        if not isinstance(meta, torch.Tensor):\n            meta = torch.tensor(meta).to(device, non_blocking=True)\n        else:\n            meta = meta.to(device, non_blocking=True)\n\n        # ---- forward/backward ---------------------------------------------------\n        with torch.cuda.amp.autocast():\n            out = model(img, meta)\n            loss = criterion(out, lab) / config.ACCUM_STEPS\n\n        scaler.scale(loss).backward()\n        epoch_loss += loss.item() * config.ACCUM_STEPS\n\n        # ---- step when we hit an accumulation boundary -----------------------\n        if (it + 1) % config.ACCUM_STEPS == 0 or it == len(train_loader) - 1:\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad()\n\n    return epoch_loss / len(train_loader)\n\n# IMMEDIATE FIX 4: Check your current batch sizes\ndef debug_batch_sizes():\n    \"\"\"Debug function to see what batch sizes you're getting\"\"\"\n    print(\"Checking batch sizes in train_loader:\")\n    for i, (img, lab, meta) in enumerate(train_loader):\n        print(f\"Batch {i}: img.shape={img.shape}\")\n        if i >= 5:  # Check first 5 batches\n            break\n    \n    print(\"Checking batch sizes in val_loader:\")\n    for i, (img, lab, meta) in enumerate(val_loader):\n        print(f\"Batch {i}: img.shape={img.shape}\")\n        if i >= 5:\n            break\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ------------------------------------------------------------------\n# Model\n# ------------------------------------------------------------------\nclass AFNet(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.backbone=timm.create_model(config.MODEL_NAME_BACKBONE,\n                                        pretrained=True, num_classes=0,\n                                        global_pool='avg')\n        dim=self.backbone.num_features\n        self.meta_fc=nn.Sequential(\n            nn.Linear(2,16),nn.ReLU(),nn.Dropout(0.2),\n            nn.Linear(16,32),nn.ReLU()\n        )\n        self.head=nn.Sequential(\n            nn.Linear(dim+32,config.HIDDEN_DIM),\n            nn.BatchNorm1d(config.HIDDEN_DIM),nn.ReLU(),nn.Dropout(0.3),\n            nn.Linear(config.HIDDEN_DIM,config.HIDDEN_DIM//2),\n            nn.BatchNorm1d(config.HIDDEN_DIM//2),nn.ReLU(),nn.Dropout(0.3),\n            nn.Linear(config.HIDDEN_DIM//2,config.NUM_CLASSES)\n        )\n    def forward(self,img,meta):\n        feat=self.backbone(img)\n        meta=self.meta_fc(meta)\n        feat=torch.cat([feat,meta],1)\n        return self.head(feat)\n\ndevice=torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nmodel=AFNet().to(device)\n\n# ------------------------------------------------------------------\n# Loss & Optimiser\n# ------------------------------------------------------------------\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nclass ImprovedLoss(nn.Module):\n    def __init__(self,\n                 weight_aneurysm: float | None = None,\n                 focal_weight: float = 0.3):\n        super().__init__()\n\n        # ---- collect settings ------------------------------------------------\n        self.num_classes = config.NUM_CLASSES\n        self.focal_weight = focal_weight\n        self.focal = FocalLoss()\n\n        # ---- create per‑class weight vector -----------------------------------\n        # Start with a 1‑vector and scale the last entry\n        self.w = torch.ones(self.num_classes, device=device)   # ← device defined elsewhere\n\n        if weight_aneurysm is None:\n            # Default fallback to the config value\n            weight_aneurysm = config.AneurysmWeight\n        self.w[-1] = weight_aneurysm              # last class = aneurysm\n\n    # --------------------------------------------------------------------------\n    # forward\n    # --------------------------------------------------------------------------\n    def forward(self, out: torch.Tensor, lab: torch.Tensor) -> torch.Tensor:\n        bce = F.binary_cross_entropy_with_logits(out, lab, reduction='none')\n        weighted_bce = (bce * self.w).mean()\n        focal_loss = self.focal(out, lab).mean()\n        return (1.0 - self.focal_weight) * weighted_bce + self.focal_weight * focal_loss\n\n\ncriterion=ImprovedLoss() if config.USE_IMPROVED_LOSS else None\n# fallback simple weighted BCE\nif criterion is None:\n    def simple_lc(out,lab):\n        w=torch.ones(config.NUM_CLASSES,device=device); w[-1]=3.0\n        return (F.binary_cross_entropy_with_logits(out,lab,reduction='none')*w).mean()\n    criterion=simple_lc\n\noptimizer=AdamW(model.parameters(),lr=config.LR,weight_decay=1e-4)\nscheduler=CosineAnnealingLR(optimizer,config.EPOCHS,eta_min=1e-6)\nscaler=torch.cuda.amp.GradScaler()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ------------------------------------------------------------------\n# Training helpers\n# ------------------------------------------------------------------\nimport torch\nimport torch.nn.functional as F\nfrom sklearn.metrics import roc_auc_score\nimport numpy as np\n\n# ------------------------------------------------------------------\n# Training helpers\n# ------------------------------------------------------------------\n\ndef _accumulate_grad(optimizer, scaler, accum_steps, it):\n    \"\"\"Step the optimizer and reset gradients every `accum_steps` steps.\"\"\"\n    scaler.step(optimizer)\n    scaler.update()\n    optimizer.zero_grad()\n    # At epoch end also step if we hit the last batch\n    if (it + 1) % accum_steps == 0:\n        return\n    else:\n        # If this was the last batch (len % accum_steps != 0)\n        if it == len(train_loader) - 1:\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad()\n\n\ndef train_one():\n    model.train()\n    torch.cuda.empty_cache()          # once per epoch only\n    epoch_loss = 0.0\n\n    for it, (img, lab, meta) in enumerate(train_loader):\n        # ---- device placement -------------------------------------------------\n        img = img.to(device, non_blocking=True)\n        lab = lab.to(device, non_blocking=True)\n\n        # Nothing guaranteed to be a tensor, so guard against that\n        if not isinstance(meta, torch.Tensor):\n            meta = torch.tensor(meta).to(device, non_blocking=True)\n        else:\n            meta = meta.to(device, non_blocking=True)\n\n        # ---- forward/backward ---------------------------------------------------\n        with torch.cuda.amp.autocast():\n            out = model(img, meta)\n            loss = criterion(out, lab) / config.ACCUM_STEPS\n\n        scaler.scale(loss).backward()\n        epoch_loss += loss.item() * config.ACCUM_STEPS   # back to unscaled loss\n\n        # ---- step when we hit an accumulation boundary -----------------------\n        if (it + 1) % config.ACCUM_STEPS == 0 or it == len(train_loader) - 1:\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad()\n\n    return epoch_loss / len(train_loader)\n\n\ndef validate_one():\n    model.eval()\n    torch.cuda.empty_cache()          # once per epoch only\n    epoch_loss = 0.0\n    out_all = []\n    lab_all = []\n\n    with torch.no_grad():\n        for img, lab, meta in val_loader:\n            img = img.to(device, non_blocking=True)\n            lab = lab.to(device, non_blocking=True)\n\n            if not isinstance(meta, torch.Tensor):\n                meta = torch.tensor(meta).to(device, non_blocking=True)\n            else:\n                meta = meta.to(device, non_blocking=True)\n\n            with torch.cuda.amp.autocast():\n                out = model(img, meta)\n                loss = criterion(out, lab)\n\n            epoch_loss += loss.item()\n\n            out_all.append(torch.sigmoid(out).cpu().numpy())\n            lab_all.append(lab.cpu().numpy())\n\n    # ----- metrics ------------------------------------------------------------\n    out_all = np.concatenate(out_all, 0)\n    lab_all = np.concatenate(lab_all, 0)\n\n    aucs = []\n    for i in range(config.NUM_CLASSES):\n        if len(np.unique(lab_all[:, i])) > 1:\n            aucs.append(roc_auc_score(lab_all[:, i], out_all[:, i]))\n        else:\n            aucs.append(0.5)                        # constant baseline\n\n    auc_mixed = np.mean(aucs[:-1])                    # all but aneurysm\n    auc_aneurysm = np.mean(aucs[-1:])                # aneurysm only\n    score = (auc_mixed + auc_aneurysm) / 2.0\n\n    return epoch_loss / len(val_loader), score","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ------------------------------------------------------------------\n# Main loop\n# ------------------------------------------------------------------\nbest=0; patience=0\nfor epoch in range(config.EPOCHS):\n    print(f\"\\nEpoch {epoch+1}/{config.EPOCHS}\")\n    tl= train_one()\n    vl, scr= validate_one()\n    scheduler.step()\n    print(f\"train:{tl:.6f} val:{vl:.6f} score:{scr:.6f}\")\n    if scr>best:\n        best=scr; patience=0; torch.save({'model':model.state_dict(),\n                                          'optimizer':optimizer.state_dict(),\n                                          'epoch':epoch+1,\n                                          'score':scr}, os.path.join(config.OUTPUT_DIR,f\"{config.MODEL_NAME_BACKBONE}_best.pth\"))\n    else: patience+=1\n    if patience>=config.PATIENCE: break\n    torch.cuda.empty_cache()\n\nprint(f\"\\nBest score {best:.6f}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Overall Pipline\n\n---\n\n**1. Setup and Configuration**\n\n* Define global settings (data paths, model hyperparameters, augmentation flags).\n* Fix random seeds for reproducibility.\n* Select device (GPU/CPU).\n\n**2. Loss Functions**\n\n* Implement **Focal Loss** for class imbalance handling.\n* Create an **Improved Loss** combining weighted BCE and focal loss, with extra weight for the aneurysm class.\n\n**3. Data Preparation**\n\n* Load metadata CSV files (training labels, series mapping, localizers).\n* Build a mapping from each imaging series to 8 representative frames (PNG or DICOM).\n* Group patients for cross-validation (GroupKFold to avoid leakage).\n\n**4. Image Processing Utilities**\n\n* Functions for windowing, CLAHE, robust normalization, and creating 3-channel composite images from volume slices.\n* Logic for slicing volumes to the desired frame count.\n\n**5. Dataset Class** (`EightFrameDataset`)\n\n* Loads and processes 8-frame volumes (PNG or DICOM).\n* Applies augmentations.\n* Generates both image tensors and metadata features (age, sex).\n* Caches loaded samples to speed up training.\n\n**6. Data Augmentation**\n\n* Strong augmentations (rotation, scaling, elastic transforms, noise, dropout, etc.) for training.\n* Minimal normalization-only pipeline for validation.\n\n**7. DataLoaders**\n\n* Create training and validation DataLoaders with multi-worker loading and batch skipping logic for small batches (to avoid BatchNorm errors).\n\n**8. Model Training & Validation Loops**\n\n* **train\\_one()**: Mixed precision training, gradient accumulation, skips problematic batches, updates optimizer.\n* **validate\\_one()**: Runs inference, computes average loss, calculates per-class AUC, and combines into final score.\n* Optional: Keep BatchNorm layers in eval mode during training for stability.\n\n**9. Cross-Validation Execution**\n\n* For each fold: train and validate model, track best score, save outputs.\n\n---\n\n","metadata":{}},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}