{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","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":"nvidiaTeslaT4","dataSources":[{"sourceId":113558,"databundleVersionId":14174843,"sourceType":"competition"}],"dockerImageVersionId":31154,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# ********************\n# PART 1: Imports & Environment Setup\n# ********************\n\n# Basic libraries\nimport os, json, math, random, warnings, glob\nfrom pathlib import Path\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image, ImageFile\nImageFile.LOAD_TRUNCATED_IMAGES = True\n\n# Computer vision & deep learning\nimport cv2\nimport torch\nimport torch.nn as nn\nimport torchvision\nfrom torchvision.transforms import functional as TF\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\n\n# Utilities\nfrom sklearn.model_selection import StratifiedShuffleSplit\nfrom skimage.feature import peak_local_max\nfrom skimage.segmentation import watershed\nfrom scipy.optimize import linear_sum_assignment\nfrom tqdm import tqdm\n\n# Configuration\nwarnings.filterwarnings('ignore', category=UserWarning)\ntorch.backends.cudnn.benchmark = True\n\n# Device setup\nDEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint('Torch:', torch.__version__, '| CUDA:', torch.cuda.is_available())\nif torch.cuda.is_available():\n    print('GPU:', torch.cuda.get_device_name(0))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T05:59:31.072323Z","iopub.execute_input":"2025-10-29T05:59:31.072658Z","iopub.status.idle":"2025-10-29T05:59:31.079339Z","shell.execute_reply.started":"2025-10-29T05:59:31.072636Z","shell.execute_reply":"2025-10-29T05:59:31.078623Z"}},"outputs":[{"name":"stdout","text":"Torch: 2.6.0+cu124 | CUDA: True\nGPU: Tesla T4\n","output_type":"stream"}],"execution_count":3},{"cell_type":"code","source":"# ********************\n# PART 2: Paths & Directory Setup\n# ********************\n\n# Main dataset and output paths\nCOMP_DIR = \"/kaggle/input/recodai-luc-scientific-image-forgery-detection\"\nTRAIN_DIR = f\"{COMP_DIR}/train_images\"\nMASK_DIR = f\"{COMP_DIR}/train_masks\"\nOUT_DIR = \"/kaggle/working\"\n\n# Create output directory if not exists\nos.makedirs(OUT_DIR, exist_ok=True)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T05:59:40.350567Z","iopub.execute_input":"2025-10-29T05:59:40.351078Z","iopub.status.idle":"2025-10-29T05:59:40.354967Z","shell.execute_reply.started":"2025-10-29T05:59:40.351057Z","shell.execute_reply":"2025-10-29T05:59:40.354216Z"}},"outputs":[],"execution_count":4},{"cell_type":"code","source":"# ********************\n# PART 3: Training Configuration & Random Seed\n# ********************\n\n# Training hyperparameters\nSEED = 42\nIMAGE_SIZE = 512\nBATCH_SIZE = 6\nEPOCHS = 10\nBASE_LR = 3e-4\nWEIGHT_DECAY = 1e-4\nWARMUP_EPOCHS = 1\n\n# Self-similarity parameters\nFEAT_RADIUS = 4  # search radius in feature map (~8x8 input px)\nUSE_RESNET_FEATURES = True  # fallback to Sobel if pretrained unavailable\n\n# Post-processing defaults\nPP_THR = 0.5\nPP_MIN_AREA = 64\nPP_MIN_DIST = 5\n\n# Seed function for reproducibility\ndef set_seed(s=SEED):\n    random.seed(s)\n    np.random.seed(s)\n    torch.manual_seed(s)\n    torch.cuda.manual_seed_all(s)\n\n# Apply seed\nset_seed()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T05:59:52.138193Z","iopub.execute_input":"2025-10-29T05:59:52.138857Z","iopub.status.idle":"2025-10-29T05:59:52.144424Z","shell.execute_reply.started":"2025-10-29T05:59:52.138833Z","shell.execute_reply":"2025-10-29T05:59:52.143775Z"}},"outputs":[],"execution_count":5},{"cell_type":"code","source":"# =====================\n# PART 1: Mask utilities\n# - load_mask_instances: supports various .npy shapes (dict, list, 2D, 3D)\n# =====================\nfrom typing import List, Optional\n\ndef load_mask_instances(mask_path: str) -> List[np.ndarray]:\n    \"\"\"\n    Load a mask file (commonly .npy) and return a list of binary instance masks.\n    Each instance is a 2D uint8 array with values 0/1.\n\n    Supports:\n    - plain 2D array (single mask or label map)\n    - 3D array (stack of masks)\n    - dict-like npy with a 'masks' key or first key\n    - list/tuple of masks\n    \"\"\"\n    m = np.load(mask_path, allow_pickle=True)\n    insts: List[np.ndarray] = []\n\n    # If saved as dict-like (common in some pipelines)\n    if isinstance(m, dict):\n        key = 'masks' if 'masks' in m else list(m.keys())[0]\n        m = np.asarray(m[key])\n\n    # If a single 2D mask / labeled mask\n    if isinstance(m, np.ndarray) and m.ndim == 2:\n        # connected components: convert any positive value to binary and separate instances\n        n, lab = cv2.connectedComponents((m > 0).astype(np.uint8))\n        for k in range(1, n):\n            insts.append((lab == k).astype(np.uint8))\n\n    # If stacked masks (N x H x W) or (C x H x W)\n    elif isinstance(m, np.ndarray) and m.ndim == 3:\n        for i in range(m.shape[0]):\n            a = m[i]\n            # if single channel 2D or multi-channel per-instance\n            mm = (a > 0).astype(np.uint8) if a.ndim == 2 else ((a > 0).sum(axis=0) > 0).astype(np.uint8)\n            if mm.sum() > 0:\n                insts.append(mm)\n\n    # If it's a list/tuple of mask arrays\n    else:\n        if isinstance(m, (list, tuple)):\n            for a in m:\n                a = np.asarray(a)\n                mm = (a > 0).astype(np.uint8) if a.ndim == 2 else ((a > 0).sum(axis=0) > 0).astype(np.uint8)\n                if mm.sum() > 0:\n                    insts.append(mm)\n\n    return insts\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T06:04:01.075479Z","iopub.execute_input":"2025-10-29T06:04:01.076228Z","iopub.status.idle":"2025-10-29T06:04:01.092177Z","shell.execute_reply.started":"2025-10-29T06:04:01.076165Z","shell.execute_reply":"2025-10-29T06:04:01.091519Z"}},"outputs":[],"execution_count":7},{"cell_type":"code","source":"# =====================\n# PART 2: Union mask & resizing helper\n# - load_union_mask: produces a single combined binary mask from instances\n# =====================\n\ndef load_union_mask(mask_path: str, target_hw: Optional[tuple] = None) -> Optional[np.ndarray]:\n    \"\"\"\n    Load instance masks and return a single unioned binary mask (H x W uint8).\n    If target_hw (h, w) is provided, the result is resized to that shape.\n    \"\"\"\n    insts = load_mask_instances(mask_path)\n    if len(insts) == 0:\n        return None\n\n    # Determine canvas size (max H and W among instances)\n    H = max(i.shape[0] for i in insts)\n    W = max(i.shape[1] for i in insts)\n    union = np.zeros((H, W), np.uint8)\n\n    # OR-combine each instance into the union canvas (resize if needed)\n    for m in insts:\n        h, w = m.shape\n        if (h, w) != (H, W):\n            m = cv2.resize(m, (W, H), interpolation=cv2.INTER_NEAREST)\n        union |= (m > 0).astype(np.uint8)\n\n    # If caller wants a specific target shape (h, w)\n    if target_hw is not None and union.shape != tuple(target_hw):\n        union = cv2.resize(union, (target_hw[1], target_hw[0]), interpolation=cv2.INTER_NEAREST)\n\n    return union\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T06:04:24.044725Z","iopub.execute_input":"2025-10-29T06:04:24.045284Z","iopub.status.idle":"2025-10-29T06:04:24.05191Z","shell.execute_reply.started":"2025-10-29T06:04:24.045259Z","shell.execute_reply":"2025-10-29T06:04:24.050903Z"}},"outputs":[],"execution_count":8},{"cell_type":"code","source":"# =====================\n# PART 3: Build items (image list) and quick usage\n# - build_items: collects images from authentic/forged subfolders or flat dir\n# - small utility: show_item to preview image+mask overlay\n# =====================\n\ndef build_items(train_dir: str, mask_dir: str):\n    exts = (\"*.png\", \"*.jpg\", \"*.jpeg\", \"*.tif\", \"*.tiff\", \"*.bmp\")\n    items = []\n    train_dir = Path(train_dir)\n    auth_dir = train_dir / 'authentic'\n    forg_dir = train_dir / 'forged'\n\n    def add_from_dir(img_dir: Path, label: int):\n        # iterate every supported extension\n        for ext in exts:\n            for p in sorted(img_dir.glob(ext)):\n                case_id = p.stem\n                mask_path = None\n                if label == 1 and Path(mask_dir).exists():\n                    cand = Path(mask_dir) / f\"{case_id}.npy\"\n                    if cand.exists():\n                        mask_path = str(cand)\n                items.append({\"path\": str(p), \"case_id\": case_id, \"label\": label, \"mask_path\": mask_path})\n\n    # If dataset follows authentic/forged subfolder layout\n    if auth_dir.exists() and forg_dir.exists():\n        add_from_dir(auth_dir, 0)\n        add_from_dir(forg_dir, 1)\n    else:\n        # Flat fallback: assume all are authentic by default (label 0)\n        for ext in exts:\n            for p in sorted(train_dir.glob(ext)):\n                items.append({\"path\": str(p), \"case_id\": p.stem, \"label\": 0, \"mask_path\": None})\n    return items\n\n# -------------------------\n# Example: build and inspect\n# -------------------------\nitems = build_items(TRAIN_DIR, MASK_DIR)\nlabels = np.array([it['label'] for it in items], dtype=np.int64)\nprint('Items:', len(items), '| forged:', labels.sum(), '| authentic:', (1 - labels).sum())\n\n# -------------------------\n# Optional helper to visualize one item (image + union mask overlay)\n# -------------------------\ndef show_item(idx: int, items_list: list, out_path: str = None):\n    \"\"\"\n    Visualize an item with its union mask overlay.\n    If out_path is given, saves overlay image there. Otherwise returns the overlay array.\n    \"\"\"\n    it = items_list[idx]\n    img = cv2.imread(it['path'], cv2.IMREAD_COLOR)\n    if img is None:\n        raise FileNotFoundError(f\"Image not found: {it['path']}\")\n\n    overlay = img.copy()\n    if it.get('mask_path'):\n        union = load_union_mask(it['mask_path'], target_hw=img.shape[:2])\n        if union is not None:\n            # create red mask overlay (alpha blend)\n            red = np.zeros_like(img)\n            red[:, :, 2] = (union * 255).astype(np.uint8)\n            overlay = cv2.addWeighted(img, 0.7, red, 0.3, 0)\n\n    if out_path:\n        cv2.imwrite(out_path, overlay)\n        print(f\"Saved preview to {out_path}\")\n    else:\n        # return BGR array\n        return overlay\n\n# usage example: save preview for first item (if any)\nif len(items) > 0:\n    preview_path = os.path.join(OUT_DIR, \"preview_0.png\")\n    show_item(0, items, out_path=preview_path)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T06:04:32.91435Z","iopub.execute_input":"2025-10-29T06:04:32.914659Z","iopub.status.idle":"2025-10-29T06:04:34.549392Z","shell.execute_reply.started":"2025-10-29T06:04:32.914636Z","shell.execute_reply":"2025-10-29T06:04:34.548501Z"}},"outputs":[{"name":"stdout","text":"Items: 5128 | forged: 2751 | authentic: 2377\nSaved preview to /kaggle/working/preview_0.png\n","output_type":"stream"}],"execution_count":9},{"cell_type":"code","source":"# =====================\n# Stratified 80/20 split\n# =====================\nif len(items)==0:\n    raise RuntimeError('No training images found. Check COMP_DIR/TRAIN_DIR.')\nsss = StratifiedShuffleSplit(n_splits=1, test_size=0.20, random_state=SEED)\ntrain_idx, valid_idx = next(sss.split(np.arange(len(labels)), labels))\ntrain_items = [items[i] for i in train_idx]\nvalid_items = [items[i] for i in valid_idx]\nprint(f'Train {len(train_items)} | Valid {len(valid_items)}')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T06:05:57.182013Z","iopub.execute_input":"2025-10-29T06:05:57.182294Z","iopub.status.idle":"2025-10-29T06:05:57.222524Z","shell.execute_reply.started":"2025-10-29T06:05:57.182274Z","shell.execute_reply":"2025-10-29T06:05:57.221583Z"}},"outputs":[{"name":"stdout","text":"Train 4102 | Valid 1026\n","output_type":"stream"}],"execution_count":11},{"cell_type":"code","source":"# =====================\n# PART 5.1: ResNet18 Feature Extractor (for self-similarity)\n# =====================\n\nimport torch\nimport torch.nn as nn\nimport torchvision\n\nclass ResNet18_Features(nn.Module):\n    def __init__(self):\n        super().__init__()\n        try:\n            # Load pretrained ResNet18 for feature extraction\n            m = torchvision.models.resnet18(weights='IMAGENET1K_V1')\n            print(' ResNet18 pretrained weights loaded for self-similarity features.')\n        except Exception as e:\n            print('ResNet18 weights not available offline. Using random initialization.')\n            print('   Reason:', e)\n            m = torchvision.models.resnet18(weights=None)\n\n        # Use only early layers to extract low-level texture/edge features\n        self.stem = nn.Sequential(m.conv1, m.bn1, m.relu, m.maxpool)\n        self.layer1 = m.layer1\n        self.layer2 = m.layer2\n\n        # Freeze parameters (no training)\n        for p in self.parameters():\n            p.requires_grad_(False)\n\n        # Set model to evaluation mode\n        self.eval()\n\n        # Register mean/std buffers for normalization (ImageNet)\n        self.register_buffer('mean', torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1))\n        self.register_buffer('std', torch.tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1))\n\n    @torch.no_grad()\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        \"\"\"\n        Forward pass through ResNet18 stem and first two layers.\n        Args:\n            x: [B,3,H,W] float tensor in range [0,1]\n        Returns:\n            Feature map [B,C,H/8,W/8]\n        \"\"\"\n        x = (x - self.mean.to(x.device)) / self.std.to(x.device)\n        x = self.stem(x)\n        x = self.layer1(x)\n        x = self.layer2(x)\n        return x\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T06:06:36.938796Z","iopub.execute_input":"2025-10-29T06:06:36.93908Z","iopub.status.idle":"2025-10-29T06:06:36.947135Z","shell.execute_reply.started":"2025-10-29T06:06:36.939059Z","shell.execute_reply":"2025-10-29T06:06:36.946143Z"}},"outputs":[],"execution_count":12},{"cell_type":"code","source":"# =====================\n# PART 5.2: Sobel Feature Extractor (Fallback if ResNet unavailable)\n# =====================\n\nclass Sobel_Features(nn.Module):\n    def __init__(self):\n        super().__init__()\n        # Define Sobel kernels for edge detection\n        kx = torch.tensor([[1, 0, -1],\n                           [2, 0, -2],\n                           [1, 0, -1]], dtype=torch.float32).view(1, 1, 3, 3)\n        ky = torch.tensor([[1, 2, 1],\n                           [0, 0, 0],\n                           [-1, -2, -1]], dtype=torch.float32).view(1, 1, 3, 3)\n\n        self.register_buffer('kx', kx)\n        self.register_buffer('ky', ky)\n        self.pool = nn.AvgPool2d(8)  # downsample ~8×\n\n    @torch.no_grad()\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        \"\"\"\n        Simple Sobel feature extraction (edge magnitude + directions).\n        Args:\n            x: [B,3,H,W] float tensor in range [0,1]\n        Returns:\n            Feature map [B,3,H/8,W/8]\n        \"\"\"\n        # Convert RGB → grayscale\n        gray = 0.2989 * x[:, 0:1] + 0.5870 * x[:, 1:2] + 0.1140 * x[:, 2:3]\n        gx = torch.conv2d(gray, self.kx.to(x.device), padding=1)\n        gy = torch.conv2d(gray, self.ky.to(x.device), padding=1)\n        mag = torch.sqrt(gx * gx + gy * gy + 1e-6)\n\n        feat = torch.cat([gx, gy, mag], dim=1)\n        return self.pool(feat)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T06:06:53.370006Z","iopub.execute_input":"2025-10-29T06:06:53.370827Z","iopub.status.idle":"2025-10-29T06:06:53.378466Z","shell.execute_reply.started":"2025-10-29T06:06:53.3708Z","shell.execute_reply":"2025-10-29T06:06:53.377506Z"}},"outputs":[],"execution_count":13},{"cell_type":"code","source":"# =====================\n# PART 5.3: Self-Similarity Map Computation\n# =====================\n\n# Try to initialize ResNet18 features; fallback to Sobel if it fails\ntry:\n    FEAT_EXTRACTOR = ResNet18_Features().to(DEVICE)\n    FEAT_MODE = 'resnet'\nexcept Exception as e:\n    print('⚠️ Falling back to Sobel features due to error:', e)\n    FEAT_EXTRACTOR = Sobel_Features().to(DEVICE)\n    FEAT_MODE = 'sobel'\n\nprint('🔍 Self-similarity feature mode:', FEAT_MODE.upper())\n\n@torch.no_grad()\ndef selfsim_map(imgs: torch.Tensor, radius: int = FEAT_RADIUS) -> torch.Tensor:\n    \"\"\"\n    Compute per-image best self-similarity map (excluding zero shift).\n\n    Args:\n        imgs: [B,3,H,W] float tensor in [0,1]\n        radius: search radius in feature map (default=FEAT_RADIUS)\n    Returns:\n        sim_up: [B,1,H,W] float tensor in [0,1]\n    \"\"\"\n    # Extract features\n    feats = FEAT_EXTRACTOR(imgs)  # [B,C,Hf,Wf]\n    \n    # Normalize along channel dimension\n    feats = feats / (feats.norm(dim=1, keepdim=True) + 1e-6)\n    B, C, Hf, Wf = feats.shape\n    maxcorr = torch.full((B, Hf, Wf), -1.0, device=imgs.device)\n\n    # Search local neighborhood shifts\n    for dy in range(-radius, radius + 1):\n        for dx in range(-radius, radius + 1):\n            if dy == 0 and dx == 0:\n                continue\n            shifted = torch.roll(feats, shifts=(dy, dx), dims=(2, 3))\n            corr = (feats * shifted).sum(dim=1)  # [B,Hf,Wf]\n            maxcorr = torch.maximum(maxcorr, corr)\n\n    # Map from [-1,1] to [0,1]\n    maxcorr = (maxcorr + 1.0) * 0.5\n\n    # Upsample to input image size\n    sim_up = torch.nn.functional.interpolate(\n        maxcorr.unsqueeze(1),\n        size=imgs.shape[-2:],\n        mode='bilinear',\n        align_corners=False\n    )\n\n    return sim_up.clamp(0, 1)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T06:07:07.471949Z","iopub.execute_input":"2025-10-29T06:07:07.472745Z","iopub.status.idle":"2025-10-29T06:07:08.511537Z","shell.execute_reply.started":"2025-10-29T06:07:07.472718Z","shell.execute_reply":"2025-10-29T06:07:08.510651Z"}},"outputs":[{"name":"stderr","text":"Downloading: \"https://download.pytorch.org/models/resnet18-f37072fd.pth\" to /root/.cache/torch/hub/checkpoints/resnet18-f37072fd.pth\n100%|██████████| 44.7M/44.7M [00:00<00:00, 146MB/s] \n","output_type":"stream"},{"name":"stdout","text":" ResNet18 pretrained weights loaded for self-similarity features.\n🔍 Self-similarity feature mode: RESNET\n","output_type":"stream"}],"execution_count":14},{"cell_type":"code","source":"# Example dummy batch\ndummy = torch.rand(2, 3, 256, 256).to(DEVICE)  # 2 RGB images in [0,1]\nsim_map = selfsim_map(dummy)\nprint(sim_map.shape)  # → torch.Size([2, 1, 256, 256])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T06:07:22.022594Z","iopub.execute_input":"2025-10-29T06:07:22.023119Z","iopub.status.idle":"2025-10-29T06:07:22.862218Z","shell.execute_reply.started":"2025-10-29T06:07:22.023097Z","shell.execute_reply":"2025-10-29T06:07:22.8614Z"}},"outputs":[{"name":"stdout","text":"torch.Size([2, 1, 256, 256])\n","output_type":"stream"}],"execution_count":15},{"cell_type":"code","source":"# =====================\n# PART 6: Dataset & DataLoaders (Photometric Augmentations Only)\n# =====================\n\nimport random\nimport numpy as np\nimport torch\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nfrom torchvision import transforms as T\nimport torchvision.transforms.functional as TF\nfrom PIL import Image\n\n# Helper: Load mask and resize to target dimensions\ndef load_union_mask(mask_path, target_hw):\n    \"\"\"\n    Loads the mask for forged region, resizes to (H, W), and returns a binary array.\n    \"\"\"\n    mask = Image.open(mask_path).convert('L')\n    mask = mask.resize((target_hw[1], target_hw[0]), resample=Image.NEAREST)\n    mask = np.array(mask, dtype=np.uint8)\n    mask = (mask > 127).astype(np.uint8)\n    return mask\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T06:07:57.084616Z","iopub.execute_input":"2025-10-29T06:07:57.084914Z","iopub.status.idle":"2025-10-29T06:07:57.090943Z","shell.execute_reply.started":"2025-10-29T06:07:57.084891Z","shell.execute_reply":"2025-10-29T06:07:57.089968Z"}},"outputs":[],"execution_count":16},{"cell_type":"code","source":"class ForgeryDataset(Dataset):\n    \"\"\"\n    Dataset for image forgery detection.\n    Loads (image, mask, label, case_id) tuples and applies augmentations.\n    \"\"\"\n\n    def __init__(self, items, image_size=512, is_train=True):\n        self.items = items\n        self.image_size = image_size\n        self.is_train = is_train\n\n    def __len__(self):\n        return len(self.items)\n\n    # ---------------------\n    # Photometric Augmentations (brightness, contrast, blur)\n    # ---------------------\n    def _aug_photometric(self, img):\n        if random.random() < 0.5:\n            img = TF.adjust_brightness(img, 0.9 + 0.2 * random.random())\n            img = TF.adjust_contrast(img, 0.9 + 0.2 * random.random())\n        if random.random() < 0.3:\n            img = TF.gaussian_blur(img, kernel_size=3)\n        return img\n\n    # ---------------------\n    # Main sample loading\n    # ---------------------\n    def __getitem__(self, idx):\n        it = self.items[idx]\n        img = Image.open(it['path']).convert('RGB')\n        W, H = img.size\n\n        # Load mask (only for forged samples)\n        if it['label'] == 1 and it.get('mask_path') is not None:\n            union = load_union_mask(it['mask_path'], target_hw=(H, W))\n        else:\n            union = np.zeros((H, W), np.uint8)\n\n        # Resize image and mask\n        img = img.resize((self.image_size, self.image_size), resample=Image.BILINEAR)\n        mask = Image.fromarray(union * 255).resize((self.image_size, self.image_size), resample=Image.NEAREST)\n\n        # ---------------------\n        # Apply random augmentations for training\n        # ---------------------\n        if self.is_train:\n            if random.random() < 0.5:\n                img = TF.hflip(img)\n                mask = TF.hflip(mask)\n            if random.random() < 0.2:\n                img = TF.vflip(img)\n                mask = TF.vflip(mask)\n            img = self._aug_photometric(img)\n\n        # ---------------------\n        # Convert to tensors\n        # ---------------------\n        img_t = TF.to_tensor(img)  # [3,H,W]\n        mask_t = torch.from_numpy(np.array(mask, dtype=np.uint8) // 255).float().unsqueeze(0)\n        label_t = torch.tensor([it['label']], dtype=torch.float32)\n\n        return img_t, mask_t, label_t, it['case_id']\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T06:08:06.349976Z","iopub.execute_input":"2025-10-29T06:08:06.350601Z","iopub.status.idle":"2025-10-29T06:08:06.360866Z","shell.execute_reply.started":"2025-10-29T06:08:06.350575Z","shell.execute_reply":"2025-10-29T06:08:06.359995Z"}},"outputs":[],"execution_count":17},{"cell_type":"code","source":"def load_union_mask(mask_path, target_hw=None):\n    \"\"\"\n    Loads a union mask from either:\n      - a .npy file (NumPy array, scientific mask)\n      - or an image file (.png, .jpg, etc.)\n    and resizes to target height/width if given.\n    \"\"\"\n    # Case 1: .npy mask (scientific format)\n    if mask_path.lower().endswith('.npy'):\n        m = np.load(mask_path, allow_pickle=True)\n\n        # Handle multiple mask types (dict, array, list)\n        if isinstance(m, dict):\n            key = 'masks' if 'masks' in m else list(m.keys())[0]\n            m = m[key]\n\n        if isinstance(m, (list, tuple)):\n            m = np.array(m)\n\n        # Convert 3D/2D to binary union mask\n        if m.ndim == 3:\n            union = (m > 0).sum(axis=0).astype(np.uint8)\n        else:\n            union = (m > 0).astype(np.uint8)\n\n    # Case 2: Image-based mask\n    else:\n        mask = Image.open(mask_path).convert('L')\n        union = np.array(mask, dtype=np.uint8)\n        union = (union > 127).astype(np.uint8)\n\n    # Resize if needed\n    if target_hw is not None and union.shape != tuple(target_hw):\n        union = cv2.resize(union, (target_hw[1], target_hw[0]), interpolation=cv2.INTER_NEAREST)\n\n    return union\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T06:09:08.200331Z","iopub.execute_input":"2025-10-29T06:09:08.201441Z","iopub.status.idle":"2025-10-29T06:09:08.209987Z","shell.execute_reply.started":"2025-10-29T06:09:08.201405Z","shell.execute_reply":"2025-10-29T06:09:08.208997Z"}},"outputs":[],"execution_count":19},{"cell_type":"code","source":"b = next(iter(train_loader))\nprint('Batch shapes:', b[0].shape, b[1].shape)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T06:09:20.977157Z","iopub.execute_input":"2025-10-29T06:09:20.977981Z","iopub.status.idle":"2025-10-29T06:09:22.680028Z","shell.execute_reply.started":"2025-10-29T06:09:20.977951Z","shell.execute_reply":"2025-10-29T06:09:22.679021Z"}},"outputs":[{"name":"stdout","text":"Batch shapes: torch.Size([6, 3, 512, 512]) torch.Size([6, 1, 512, 512])\n","output_type":"stream"}],"execution_count":20},{"cell_type":"code","source":"def adapt_first_conv_to_4ch(model):\n    \"\"\"\n    Adapt the first convolution layer of DeepLabV3 (ResNet backbone) to accept 4 channels.\n    Copies pretrained weights for first 3 channels and averages for the 4th.\n    \"\"\"\n    bb = model.backbone\n    body = getattr(bb, 'body', None)  # Some models wrap the resnet in 'body'\n    if body is None:\n        body = bb\n\n    conv1 = body.conv1\n    new_conv = nn.Conv2d(\n        4, conv1.out_channels,\n        kernel_size=conv1.kernel_size,\n        stride=conv1.stride,\n        padding=conv1.padding,\n        bias=False\n    )\n\n    with torch.no_grad():\n        if conv1.weight.shape[1] == 3:\n            new_conv.weight[:, :3] = conv1.weight\n            # Initialize 4th channel as mean of existing 3\n            new_conv.weight[:, 3:4] = conv1.weight.mean(dim=1, keepdim=True)\n        else:\n            # Unexpected, copy what we can\n            new_conv.weight[:, :min(4, conv1.weight.shape[1])] = conv1.weight[:, :min(4, conv1.weight.shape[1])]\n\n    body.conv1 = new_conv\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T06:10:33.706343Z","iopub.execute_input":"2025-10-29T06:10:33.707151Z","iopub.status.idle":"2025-10-29T06:10:33.713523Z","shell.execute_reply.started":"2025-10-29T06:10:33.707121Z","shell.execute_reply":"2025-10-29T06:10:33.712642Z"}},"outputs":[],"execution_count":21},{"cell_type":"code","source":"class BCEDiceLoss(nn.Module):\n    \"\"\"\n    Combined Binary Cross-Entropy (BCE) + Dice loss for segmentation.\n    \"\"\"\n\n    def __init__(self, pos_weight=None, eps=1e-6):\n        super().__init__()\n        self.eps = eps\n        self.bce = nn.BCEWithLogitsLoss(pos_weight=pos_weight)\n\n    def forward(self, logits, targets):\n        bce = self.bce(logits, targets)\n        probs = torch.sigmoid(logits)\n        inter = (probs * targets).sum(dim=(2, 3))\n        union = probs.sum(dim=(2, 3)) + targets.sum(dim=(2, 3)) + self.eps\n        dice = 1 - (2 * inter + self.eps) / union\n        return bce + dice.mean()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T06:10:47.375254Z","iopub.execute_input":"2025-10-29T06:10:47.37586Z","iopub.status.idle":"2025-10-29T06:10:47.382352Z","shell.execute_reply.started":"2025-10-29T06:10:47.375836Z","shell.execute_reply":"2025-10-29T06:10:47.381378Z"}},"outputs":[],"execution_count":22},{"cell_type":"code","source":"def build_model_4ch():\n    \"\"\"\n    Builds DeepLabV3-ResNet50 with single-channel output (for mask)\n    and 4-channel input adapted for self-similarity maps + RGB.\n    \"\"\"\n    m = torchvision.models.segmentation.deeplabv3_resnet50(weights='DEFAULT', aux_loss=True)\n\n    # Replace classifier heads to single-channel output\n    m.classifier[4] = nn.Conv2d(256, 1, kernel_size=1)\n    if m.aux_classifier is not None:\n        m.aux_classifier[4] = nn.Conv2d(256, 1, kernel_size=1)\n\n    # Adapt first conv to 4-channel input\n    adapt_first_conv_to_4ch(m)\n    return m\n\n# Initialize model\nmodel = build_model_4ch().to(DEVICE)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T06:11:00.215633Z","iopub.execute_input":"2025-10-29T06:11:00.216455Z","iopub.status.idle":"2025-10-29T06:11:02.139639Z","shell.execute_reply.started":"2025-10-29T06:11:00.216427Z","shell.execute_reply":"2025-10-29T06:11:02.138878Z"}},"outputs":[{"name":"stderr","text":"Downloading: \"https://download.pytorch.org/models/deeplabv3_resnet50_coco-cd0a2569.pth\" to /root/.cache/torch/hub/checkpoints/deeplabv3_resnet50_coco-cd0a2569.pth\n100%|██████████| 161M/161M [00:00<00:00, 175MB/s]  \n","output_type":"stream"}],"execution_count":23},{"cell_type":"code","source":"# Positive weight for BCE to balance forged pixels\npos_weight = torch.tensor([8.0], device=DEVICE)\n\n# Loss function\ncriterion = BCEDiceLoss(pos_weight=pos_weight)\n\n# Optimizer\noptimizer = torch.optim.AdamW(model.parameters(), lr=BASE_LR, weight_decay=WEIGHT_DECAY)\n\n# Cosine LR scheduler with warmup\ndef lr_lambda(epoch):\n    if epoch < WARMUP_EPOCHS:\n        return float(epoch + 1) / float(max(1, WARMUP_EPOCHS))\n    progress = (epoch - WARMUP_EPOCHS) / max(1, EPOCHS - WARMUP_EPOCHS)\n    return 0.5 * (1.0 + math.cos(math.pi * progress))\n\nscheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda=lr_lambda)\n\n# Automatic Mixed Precision\nscaler = torch.cuda.amp.GradScaler(enabled=torch.cuda.is_available())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T06:11:15.80636Z","iopub.execute_input":"2025-10-29T06:11:15.806675Z","iopub.status.idle":"2025-10-29T06:11:15.816073Z","shell.execute_reply.started":"2025-10-29T06:11:15.806653Z","shell.execute_reply":"2025-10-29T06:11:15.815086Z"}},"outputs":[{"name":"stderr","text":"/tmp/ipykernel_37/2012720799.py:20: FutureWarning: `torch.cuda.amp.GradScaler(args...)` is deprecated. Please use `torch.amp.GradScaler('cuda', args...)` instead.\n  scaler = torch.cuda.amp.GradScaler(enabled=torch.cuda.is_available())\n","output_type":"stream"}],"execution_count":24},{"cell_type":"code","source":"from skimage.feature import peak_local_max\nfrom skimage.segmentation import watershed\nimport cv2\nimport numpy as np\n\ndef prob_to_instances(prob: np.ndarray, thr=0.5, min_area=64, min_dist=5):\n    \"\"\"\n    Converts a probability map into individual binary instances using:\n      - Thresholding\n      - Distance transform + local maxima\n      - Watershed segmentation\n    \"\"\"\n    binm = (prob >= thr).astype(np.uint8)\n    if binm.sum() == 0:\n        return []\n\n    # Distance transform for watershed\n    dist = cv2.distanceTransform((binm*255).astype(np.uint8), cv2.DIST_L2, 3)\n    coords = peak_local_max(dist, min_distance=min_dist, labels=binm)\n\n    # Watershed markers\n    markers = np.zeros_like(binm, np.int32)\n    for i, (y, x) in enumerate(coords, start=1):\n        markers[y, x] = i\n\n    labels = watershed(-dist, markers, mask=binm.astype(bool))\n    \n    # Collect instances\n    insts = []\n    for k in range(1, labels.max()+1):\n        m = (labels == k).astype(np.uint8)\n        if m.sum() >= min_area:\n            insts.append(m)\n    return insts\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T06:12:07.640784Z","iopub.execute_input":"2025-10-29T06:12:07.641506Z","iopub.status.idle":"2025-10-29T06:12:07.648402Z","shell.execute_reply.started":"2025-10-29T06:12:07.641479Z","shell.execute_reply":"2025-10-29T06:12:07.647438Z"}},"outputs":[],"execution_count":26},{"cell_type":"code","source":"def split_instances_from_union(union_mask: np.ndarray, min_area=64):\n    \"\"\"\n    Splits a union binary mask into separate connected components.\n    \"\"\"\n    n, lab = cv2.connectedComponents((union_mask > 0).astype(np.uint8))\n    insts = []\n    for k in range(1, n):\n        m = (lab == k).astype(np.uint8)\n        if m.sum() >= min_area:\n            insts.append(m)\n    return insts\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T06:12:23.960878Z","iopub.execute_input":"2025-10-29T06:12:23.961222Z","iopub.status.idle":"2025-10-29T06:12:23.966038Z","shell.execute_reply.started":"2025-10-29T06:12:23.961202Z","shell.execute_reply":"2025-10-29T06:12:23.965078Z"}},"outputs":[],"execution_count":27},{"cell_type":"code","source":"def pixel_f1(a: np.ndarray, b: np.ndarray):\n    \"\"\"\n    Computes pixel-level F1 between two binary masks.\n    \"\"\"\n    a = (a > 0).ravel()\n    b = (b > 0).ravel()\n    tp = np.sum((a==1) & (b==1))\n    fp = np.sum((a==1) & (b==0))\n    fn = np.sum((a==0) & (b==1))\n    p = tp / (tp + fp) if tp + fp > 0 else 0.0\n    r = tp / (tp + fn) if tp + fn > 0 else 0.0\n    return 2 * p * r / (p + r) if (p + r) > 0 else 0.0\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T06:12:33.405666Z","iopub.execute_input":"2025-10-29T06:12:33.406092Z","iopub.status.idle":"2025-10-29T06:12:33.414251Z","shell.execute_reply.started":"2025-10-29T06:12:33.406057Z","shell.execute_reply":"2025-10-29T06:12:33.413068Z"}},"outputs":[],"execution_count":28},{"cell_type":"code","source":"from scipy.optimize import linear_sum_assignment\n\ndef of1_for_image(pred_insts, gt_insts):\n    \"\"\"\n    Computes object-level F1 (oF1) for a single image.\n    Uses Hungarian matching for best assignment.\n    \"\"\"\n    if len(gt_insts) == 0 and len(pred_insts) == 0:\n        return 1.0\n    if len(gt_insts) == 0 or len(pred_insts) == 0:\n        return 0.0\n\n    n_p, n_g = len(pred_insts), len(gt_insts)\n    M = np.zeros((max(n_p, n_g), max(n_p, n_g)), np.float32)\n\n    for i in range(n_p):\n        for j in range(n_g):\n            M[i, j] = pixel_f1(pred_insts[i], gt_insts[j])\n\n    r, c = linear_sum_assignment(-M)\n    matched = M[r, c]\n    k = min(n_p, n_g)\n    score = matched[:k].mean() if k > 0 else 0.0\n\n    # Penalty for extra predictions\n    penalty = n_g / max(n_p, n_g)\n    return float(score * penalty)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T06:12:46.054852Z","iopub.execute_input":"2025-10-29T06:12:46.055567Z","iopub.status.idle":"2025-10-29T06:12:46.061792Z","shell.execute_reply.started":"2025-10-29T06:12:46.055543Z","shell.execute_reply":"2025-10-29T06:12:46.0608Z"}},"outputs":[],"execution_count":29},{"cell_type":"code","source":"import torch\n\n@torch.no_grad()\ndef evaluate_of1(model, loader, thr=0.5, min_area=64, min_dist=5):\n    \"\"\"\n    Computes mean object-level F1 for a dataset.\n    Concatenates RGB + self-similarity map as 4th channel input.\n    \"\"\"\n    model.eval()\n    scores = []\n\n    for imgs, masks, labels, ids in loader:\n        imgs = imgs.to(DEVICE)\n\n        # Build self-similarity channel\n        sim = selfsim_map(imgs)  # [B,1,H,W]\n        x4 = torch.cat([imgs, sim], dim=1)  # [B,4,H,W]\n\n        out = model(x4)\n        probs = torch.sigmoid(out['out']).cpu().numpy()\n        masks_np = masks.cpu().numpy()\n        B = probs.shape[0]\n\n        for b in range(B):\n            prob = probs[b, 0]\n            gt_union = (masks_np[b, 0] > 0).astype(np.uint8)\n            pred_insts = prob_to_instances(prob, thr=thr, min_area=min_area, min_dist=min_dist)\n            gt_insts = split_instances_from_union(gt_union, min_area=min_area)\n            scores.append(of1_for_image(pred_insts, gt_insts))\n\n    return float(np.mean(scores)) if scores else 0.0\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T06:12:57.498482Z","iopub.execute_input":"2025-10-29T06:12:57.499225Z","iopub.status.idle":"2025-10-29T06:12:57.506051Z","shell.execute_reply.started":"2025-10-29T06:12:57.499179Z","shell.execute_reply":"2025-10-29T06:12:57.505135Z"}},"outputs":[],"execution_count":30},{"cell_type":"code","source":"def train_one_epoch(model, loader, optimizer, scaler=None, print_every=100):\n    \"\"\"\n    Train the model for one epoch with optional AMP (scaler).\n    Builds the 4th self-similarity channel on-the-fly.\n    \"\"\"\n    model.train()\n    running_loss = 0.0\n\n    for it, (imgs, masks, _, _) in enumerate(loader, start=1):\n        imgs = imgs.to(DEVICE)\n        masks = masks.to(DEVICE)\n\n        # Build 4th channel (self-similarity map)\n        with torch.no_grad():\n            sim = selfsim_map(imgs)  # [B,1,H,W]\n        x4 = torch.cat([imgs, sim], dim=1)  # [B,4,H,W]\n\n        optimizer.zero_grad(set_to_none=True)\n\n        with torch.cuda.amp.autocast(enabled=(scaler is not None)):\n            out = model(x4)\n            logits = out['out']\n\n            # Main + auxiliary loss\n            loss_main = criterion(logits, masks)\n            loss_aux = criterion(out['aux'], masks) if out.get('aux') is not None else 0.0\n            loss = loss_main + (0.4 * loss_aux if isinstance(loss_aux, torch.Tensor) else 0.0)\n\n        # Backprop\n        if scaler is not None:\n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n        else:\n            loss.backward()\n            optimizer.step()\n\n        running_loss += float(loss.item())\n\n        if it % print_every == 0:\n            print(f\"[train] it {it:04d} | loss {running_loss/it:.4f}\")\n\n    return running_loss / max(1, it)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T06:13:57.176138Z","iopub.execute_input":"2025-10-29T06:13:57.176965Z","iopub.status.idle":"2025-10-29T06:13:57.184632Z","shell.execute_reply.started":"2025-10-29T06:13:57.176938Z","shell.execute_reply":"2025-10-29T06:13:57.18371Z"}},"outputs":[],"execution_count":31},{"cell_type":"code","source":"def lr_lambda(epoch):\n    \"\"\"\n    Cosine decay with linear warmup.\n    \"\"\"\n    if epoch < WARMUP_EPOCHS:\n        return float(epoch + 1) / max(1, WARMUP_EPOCHS)\n    progress = (epoch - WARMUP_EPOCHS) / max(1, EPOCHS - WARMUP_EPOCHS)\n    return 0.5 * (1.0 + math.cos(math.pi * progress))\n\nscheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda=lr_lambda)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T06:14:09.511201Z","iopub.execute_input":"2025-10-29T06:14:09.511983Z","iopub.status.idle":"2025-10-29T06:14:09.516553Z","shell.execute_reply.started":"2025-10-29T06:14:09.511955Z","shell.execute_reply":"2025-10-29T06:14:09.515695Z"}},"outputs":[],"execution_count":32},{"cell_type":"code","source":"best_of1 = -1.0\nbest_pp = {'thr': PP_THR, 'min_area': PP_MIN_AREA, 'min_dist': PP_MIN_DIST}\nckpt_path = f\"{OUT_DIR}/deeplabv3_resnet50_selfsim.pt\"\n\nfor epoch in range(1, EPOCHS + 1):\n    print(f\"\\n==== Epoch {epoch}/{EPOCHS} ====\")\n\n    # Train\n    train_loss = train_one_epoch(model, train_loader, optimizer, scaler=scaler, print_every=50)\n    \n    # Scheduler step\n    scheduler.step()\n\n    # Validate\n    valid_of1 = evaluate_of1(model, valid_loader,\n                             thr=best_pp['thr'],\n                             min_area=best_pp['min_area'],\n                             min_dist=best_pp['min_dist'])\n\n    print(f\"Epoch {epoch:02d} | train_loss={train_loss:.4f} | valid_oF1={valid_of1:.4f} | LR={scheduler.get_last_lr()[0]:.2e}\")\n\n    # Save best model\n    if valid_of1 > best_of1:\n        best_of1 = valid_of1\n        torch.save({\n            'state_dict': model.state_dict(),\n            'image_size': IMAGE_SIZE,\n            'arch': 'deeplabv3_resnet50_4ch'\n        }, ckpt_path)\n        print(f\"  New best saved → {ckpt_path} (oF1={best_of1:.4f})\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T06:14:18.418596Z","iopub.execute_input":"2025-10-29T06:14:18.418876Z","execution_failed":"2025-10-29T06:26:06.614Z"}},"outputs":[{"name":"stdout","text":"\n==== Epoch 1/10 ====\n","output_type":"stream"},{"name":"stderr","text":"/tmp/ipykernel_37/2056855907.py:20: FutureWarning: `torch.cuda.amp.autocast(args...)` is deprecated. Please use `torch.amp.autocast('cuda', args...)` instead.\n  with torch.cuda.amp.autocast(enabled=(scaler is not None)):\n","output_type":"stream"},{"name":"stdout","text":"[train] it 0050 | loss 2.2025\n[train] it 0100 | loss 2.1185\n[train] it 0150 | loss 2.0666\n[train] it 0200 | loss 2.0270\n[train] it 0250 | loss 1.9915\n[train] it 0300 | loss 1.9650\n[train] it 0350 | loss 1.9466\n[train] it 0400 | loss 1.9335\n[train] it 0450 | loss 1.9193\n[train] it 0500 | loss 1.9145\n[train] it 0550 | loss 1.9033\n[train] it 0600 | loss 1.8964\n[train] it 0650 | loss 1.8828\nEpoch 01 | train_loss=1.8816 | valid_oF1=0.0284 | LR=3.00e-04\n ✅ New best saved → /kaggle/working/deeplabv3_resnet50_selfsim.pt (oF1=0.0284)\n\n==== Epoch 2/10 ====\n[train] it 0050 | loss 1.8062\n","output_type":"stream"}],"execution_count":null},{"cell_type":"code","source":"# ---------------------\n# Train one epoch\n# ---------------------\ndef train_one_epoch(model, loader, optimizer, scaler=None, print_every=100):\n    model.train()\n    running_loss = 0.0\n\n    for it, (imgs, masks, _, _) in enumerate(loader, start=1):\n        imgs = imgs.to(DEVICE)\n        masks = masks.to(DEVICE)\n\n        # Build 4th self-similarity channel on-the-fly\n        with torch.no_grad():\n            sim = selfsim_map(imgs)  # [B,1,H,W]\n        x4 = torch.cat([imgs, sim], dim=1)  # [B,4,H,W]\n\n        optimizer.zero_grad(set_to_none=True)\n\n        # Forward + loss\n        with torch.cuda.amp.autocast(enabled=(scaler is not None)):\n            out = model(x4)\n            logits = out['out']\n\n            loss_main = criterion(logits, masks)\n            loss_aux = criterion(out['aux'], masks) if out.get('aux') is not None else 0.0\n            loss = loss_main + (0.4 * loss_aux if isinstance(loss_aux, torch.Tensor) else 0.0)\n\n        # Backprop\n        if scaler is not None:\n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n        else:\n            loss.backward()\n            optimizer.step()\n\n        running_loss += float(loss.item())\n\n        if it % print_every == 0:\n            print(f\"[train] it {it:04d} | loss {running_loss/it:.4f}\")\n\n    return running_loss / max(1, it)\n\n\n# ---------------------\n# Learning rate schedule\n# ---------------------\ndef lr_lambda(epoch):\n    if epoch < WARMUP_EPOCHS:\n        return float(epoch + 1) / max(1, WARMUP_EPOCHS)\n    progress = (epoch - WARMUP_EPOCHS) / max(1, EPOCHS - WARMUP_EPOCHS)\n    return 0.5 * (1.0 + math.cos(math.pi * progress))\n\n\nscheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda=lr_lambda)\n\n\n# ---------------------\n# Full training loop\n# ---------------------\nbest_of1 = -1.0\nbest_pp = {'thr': PP_THR, 'min_area': PP_MIN_AREA, 'min_dist': PP_MIN_DIST}\nckpt_path = f\"{OUT_DIR}/deeplabv3_resnet50_selfsim.pt\"\n\nfor epoch in range(1, EPOCHS + 1):\n    print(f\"\\n==== Epoch {epoch}/{EPOCHS} ====\")\n\n    # Train\n    tr_loss = train_one_epoch(model, train_loader, optimizer, scaler=scaler, print_every=50)\n\n    # Update LR\n    scheduler.step()\n\n    # Validate with oF1\n    va_of1 = evaluate_of1(\n        model, valid_loader,\n        thr=best_pp['thr'],\n        min_area=best_pp['min_area'],\n        min_dist=best_pp['min_dist']\n    )\n\n    print(f\"Epoch {epoch:02d} | train_loss={tr_loss:.4f} | valid_oF1={va_of1:.4f} | LR={scheduler.get_last_lr()[0]:.2e}\")\n\n    # Save best model\n    if va_of1 > best_of1:\n        best_of1 = va_of1\n        torch.save({\n            'state_dict': model.state_dict(),\n            'image_size': IMAGE_SIZE,\n            'arch': 'deeplabv3_resnet50_4ch'\n        }, ckpt_path)\n        print(f\" New best saved → {ckpt_path} (oF1={best_of1:.4f})\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T06:26:47.481069Z","iopub.execute_input":"2025-10-29T06:26:47.481252Z","iopub.status.idle":"2025-10-29T06:26:47.568876Z","shell.execute_reply.started":"2025-10-29T06:26:47.481235Z","shell.execute_reply":"2025-10-29T06:26:47.567966Z"}},"outputs":[{"traceback":["\u001b[0;31m---------------------------------------------------------------------------\u001b[0m","\u001b[0;31mNameError\u001b[0m                                 Traceback (most recent call last)","\u001b[0;32m/tmp/ipykernel_37/1679414374.py\u001b[0m in \u001b[0;36m<cell line: 0>\u001b[0;34m()\u001b[0m\n\u001b[1;32m     53\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m     54\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m---> 55\u001b[0;31m \u001b[0mscheduler\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mtorch\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0moptim\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mlr_scheduler\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mLambdaLR\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0moptimizer\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mlr_lambda\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0mlr_lambda\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m     56\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m     57\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;31mNameError\u001b[0m: name 'torch' is not defined"],"ename":"NameError","evalue":"name 'torch' is not defined","output_type":"error"}],"execution_count":1},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -----------------------------\n# Train one epoch\n# -----------------------------\ndef train_one_epoch(model, loader, optimizer, scaler=None, print_every=100):\n    model.train()\n    running_loss = 0.0\n\n    for it, (imgs, masks, _, _) in enumerate(loader, start=1):\n        imgs = imgs.to(DEVICE)\n        masks = masks.to(DEVICE)\n\n        # Build 4th self-similarity channel on-the-fly\n        with torch.no_grad():\n            sim = selfsim_map(imgs)  # [B,1,H,W]\n        x4 = torch.cat([imgs, sim], dim=1)  # [B,4,H,W]\n\n        optimizer.zero_grad(set_to_none=True)\n\n        # Forward + loss\n        with torch.cuda.amp.autocast(enabled=(scaler is not None)):\n            out = model(x4)\n            logits = out['out']\n\n            loss_main = criterion(logits, masks)\n            loss_aux = criterion(out['aux'], masks) if out.get('aux') is not None else 0.0\n            loss = loss_main + (0.4 * loss_aux if isinstance(loss_aux, torch.Tensor) else 0.0)\n\n        # Backprop\n        if scaler is not None:\n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n        else:\n            loss.backward()\n            optimizer.step()\n\n        running_loss += float(loss.item())\n\n        if it % print_every == 0:\n            print(f\"[train] it {it:04d} | loss {running_loss/it:.4f}\")\n\n    return running_loss / max(1, it)\n\n\n# -----------------------------\n# Learning rate schedule\n# -----------------------------\ndef lr_lambda(epoch):\n    if epoch < WARMUP_EPOCHS:\n        return float(epoch + 1) / max(1, WARMUP_EPOCHS)\n    progress = (epoch - WARMUP_EPOCHS) / max(1, EPOCHS - WARMUP_EPOCHS)\n    return 0.5 * (1.0 + math.cos(math.pi * progress))\n\n\nscheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda=lr_lambda)\n\n\n# -----------------------------\n# Full training loop\n# -----------------------------\nbest_of1 = -1.0\nbest_pp = {'thr': PP_THR, 'min_area': PP_MIN_AREA, 'min_dist': PP_MIN_DIST}\nckpt_path = f\"{OUT_DIR}/deeplabv3_resnet50_selfsim.pt\"\n\nfor epoch in range(1, EPOCHS + 1):\n    print(f\"\\n==== Epoch {epoch}/{EPOCHS} ====\")\n\n    # Train\n    tr_loss = train_one_epoch(model, train_loader, optimizer, scaler=scaler, print_every=50)\n\n    # Update learning rate\n    scheduler.step()\n\n    # Validate with oF1 metric\n    va_of1 = evaluate_of1(\n        model, valid_loader,\n        thr=best_pp['thr'],\n        min_area=best_pp['min_area'],\n        min_dist=best_pp['min_dist']\n    )\n\n    print(f\"Epoch {epoch:02d} | train_loss={tr_loss:.4f} | valid_oF1={va_of1:.4f} | LR={scheduler.get_last_lr()[0]:.2e}\")\n\n    # Save best model\n    if va_of1 > best_of1:\n        best_of1 = va_of1\n        torch.save({\n            'state_dict': model.state_dict(),\n            'image_size': IMAGE_SIZE,\n            'arch': 'deeplabv3_resnet50_4ch'\n        }, ckpt_path)\n        print(f\" ✅ New best saved → {ckpt_path} (oF1={best_of1:.4f})\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-26T16:10:17.80752Z","iopub.execute_input":"2025-10-26T16:10:17.807816Z","iopub.status.idle":"2025-10-26T16:21:17.978652Z","shell.execute_reply.started":"2025-10-26T16:10:17.807788Z","shell.execute_reply":"2025-10-26T16:21:17.977732Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from time import time\nfrom tqdm import tqdm\nimport json\n\n# -----------------------------\n# Collect validation probabilities\n# -----------------------------\n@torch.no_grad()\ndef collect_val_probs(model, loader):\n    \"\"\"Run the model once over validation; return list of dicts with prob map and GT.\"\"\"\n    model.eval()\n    cached = []\n\n    for imgs, masks, labels, ids in tqdm(loader, desc=\"Infer val (cache)\", ncols=90):\n        imgs = imgs.to(DEVICE)\n        sim = selfsim_map(imgs)  # [B,1,H,W]\n        x4 = torch.cat([imgs, sim], dim=1)  # [B,4,H,W]\n        out = model(x4)\n        probs = torch.sigmoid(out[\"out\"]).float().cpu().numpy()\n        masks_np = masks.cpu().numpy()\n\n        for b in range(probs.shape[0]):\n            cached.append({\n                \"case_id\": ids[b],\n                \"prob\": probs[b, 0],             # HxW float32 in [0,1]\n                \"gt_union\": (masks_np[b, 0] > 0).astype(np.uint8)  # HxW uint8 {0,1}\n            })\n\n    print(f\"Collected {len(cached)} val prob maps.\")\n    return cached\n\n\n# -----------------------------\n# Quick threshold calibration\n# -----------------------------\ndef quick_calibrate_threshold(cached, thr_values, min_area=64, min_dist=5, subset_size=150, seed=42):\n    \"\"\"Sweep only threshold on a small random subset to pick thr quickly.\"\"\"\n    rng = np.random.default_rng(seed)\n\n    if len(cached) == 0:\n        raise RuntimeError(\"No cached validation predictions. Run collect_val_probs first.\")\n\n    idx = rng.choice(len(cached), size=min(subset_size, len(cached)), replace=False)\n    cached_small = [cached[i] for i in idx]\n\n    best_score, best_thr = -1.0, None\n    for thr in tqdm(thr_values, desc=\"Sweep thr\", ncols=90):\n        scores = []\n        for e in cached_small:\n            pred_insts = prob_to_instances(e[\"prob\"], thr=thr, min_area=min_area, min_dist=min_dist)\n            gt_insts = split_instances_from_union(e[\"gt_union\"], min_area=min_area)\n            scores.append(of1_for_image(pred_insts, gt_insts))\n        s = float(np.mean(scores)) if scores else 0.0\n        if s > best_score:\n            best_score, best_thr = s, thr\n\n    return best_score, best_thr\n\n\n# -----------------------------\n# Run quick post-processing calibration\n# -----------------------------\nprint(\"\\nCalibrating post-processing (quick, threshold-only)...\")\nt0 = time()\n\n# 1) Cache val probs once (reuse if already have `cached_val`)\ntry:\n    cached_val\nexcept NameError:\n    cached_val = collect_val_probs(model, valid_loader)\n\n# 2) Pick fixed structural params, sweep threshold only\nQUICK_MIN_AREA = 64\nQUICK_MIN_DIST = 5\nTHR_GRID = np.linspace(0.30, 0.65, 9)  # 9 thresholds: 0.30, 0.35, ..., 0.65\nSUBSET_SIZE = 150                       # reduce/increase for speed vs. robustness\n\nbest_score, best_thr = quick_calibrate_threshold(\n    cached_val, THR_GRID,\n    min_area=QUICK_MIN_AREA, min_dist=QUICK_MIN_DIST,\n    subset_size=SUBSET_SIZE, seed=42\n)\n\nbest_pp = {\n    \"thr\": float(best_thr),\n    \"min_area\": int(QUICK_MIN_AREA),\n    \"min_dist\": int(QUICK_MIN_DIST)\n}\n\n# Save calibration to file\nwith open(f\"{OUT_DIR}/postproc_selfsim.json\", \"w\") as f:\n    json.dump(best_pp, f)\n\nt1 = time()\nprint(f\"Quick calib → best oF1={best_score:.4f} @ thr={best_thr:.3f}, \"\n      f\"min_area={QUICK_MIN_AREA}, min_dist={QUICK_MIN_DIST} | elapsed: {t1-t0:.2f}s\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-26T16:31:20.17296Z","iopub.execute_input":"2025-10-26T16:31:20.173292Z","iopub.status.idle":"2025-10-26T16:34:31.162315Z","shell.execute_reply.started":"2025-10-26T16:31:20.173268Z","shell.execute_reply":"2025-10-26T16:34:31.161511Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =====================\n# Optional: Export TorchScript (logits-only wrapper)\n# =====================\nts_path = f\"{OUT_DIR}/deeplabv3_selfsim_ts.pt\"\n\ntry:\n    # Load trained 4-channel model\n    export_model = build_model_4ch().to(DEVICE)\n    state = torch.load(f\"{OUT_DIR}/deeplabv3_resnet50_selfsim.pt\", map_location=DEVICE)\n    export_model.load_state_dict(state['state_dict'])\n    export_model.eval()\n\n    # Wrap model to return only logits\n    class DeeplabOut(nn.Module):\n        def __init__(self, m):\n            super().__init__()\n            self.m = m\n            for p in self.m.parameters():\n                p.requires_grad_(False)\n\n        def forward(self, x):\n            return self.m(x)['out']\n\n    wrapper = DeeplabOut(export_model).to(DEVICE)\n\n    # Test with dummy input and export\n    ex = torch.randn(1, 4, IMAGE_SIZE, IMAGE_SIZE, device=DEVICE)\n    scripted = torch.jit.script(wrapper)\n    _ = scripted(ex)\n    scripted.save(ts_path)\n    print(' Saved TorchScript model:', ts_path)\n\nexcept Exception as e:\n    print(' TorchScript export skipped:', e)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}