{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":113558,"databundleVersionId":14456136,"sourceType":"competition"}],"dockerImageVersionId":31193,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# Cell 1 – installs (run ONCE per runtime)\n!pip install -q albumentations==1.4.3\n!pip install -q segmentation-models-pytorch --no-deps\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip uninstall -y numpy scipy\n!pip install numpy==1.26.4\n!pip install scipy==1.11.4\n!pip install albumentations==1.4.3\n!pip install segmentation-models-pytorch\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-23T09:55:11.514726Z","iopub.execute_input":"2025-11-23T09:55:11.514987Z","iopub.status.idle":"2025-11-23T09:55:40.176606Z","shell.execute_reply.started":"2025-11-23T09:55:11.514969Z","shell.execute_reply":"2025-11-23T09:55:40.175864Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip uninstall -y numpy scipy matplotlib albumentations\n\n!pip install \"numpy==1.26.4\" \"scipy==1.11.4\" \"matplotlib==3.8.4\"\n\n# install albumentations WITHOUT touching numpy/scipy\n!pip install \"albumentations==1.4.3\" --no-deps\n\n# install SMP WITHOUT deps as well\n!pip install \"segmentation-models-pytorch==0.5.0\" --no-deps\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-23T10:00:33.529941Z","iopub.execute_input":"2025-11-23T10:00:33.530224Z","iopub.status.idle":"2025-11-23T10:01:26.101387Z","shell.execute_reply.started":"2025-11-23T10:00:33.530202Z","shell.execute_reply":"2025-11-23T10:01:26.100376Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport scipy\nimport matplotlib\nimport matplotlib.pyplot as plt\nimport albumentations\nimport segmentation_models_pytorch as smp\n\nprint(\"numpy:\", np.__version__)\nprint(\"scipy:\", scipy.__version__)\nprint(\"matplotlib:\", matplotlib.__version__)\nprint(\"albumentations:\", albumentations.__version__)\nprint(\"smp:\", smp.__version__)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-23T10:01:31.299665Z","iopub.execute_input":"2025-11-23T10:01:31.300709Z","iopub.status.idle":"2025-11-23T10:01:39.008122Z","shell.execute_reply.started":"2025-11-23T10:01:31.300678Z","shell.execute_reply":"2025-11-23T10:01:39.007321Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport albumentations\nimport segmentation_models_pytorch as smp\n\nprint(\"numpy:\", np.__version__)\nprint(\"albumentations:\", albumentations.__version__)\nprint(\"smp:\", smp.__version__)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-23T10:02:59.747352Z","iopub.execute_input":"2025-11-23T10:02:59.747661Z","iopub.status.idle":"2025-11-23T10:02:59.753166Z","shell.execute_reply.started":"2025-11-23T10:02:59.747639Z","shell.execute_reply":"2025-11-23T10:02:59.752116Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 1 – Imports & basic config (NO sklearn)\n# ============================================================\nimport os\nfrom pathlib import Path\nimport random\nimport gc\n\nimport numpy as np\nimport pandas as pd\nimport cv2\nfrom PIL import Image\nimport matplotlib.pyplot as plt\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nimport segmentation_models_pytorch as smp\nfrom tqdm.auto import tqdm\n\nSEED = 42\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\ntorch.cuda.manual_seed_all(SEED)\n\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nprint(\"Device:\", DEVICE)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-23T10:07:18.860241Z","iopub.execute_input":"2025-11-23T10:07:18.861024Z","iopub.status.idle":"2025-11-23T10:07:18.965883Z","shell.execute_reply.started":"2025-11-23T10:07:18.860996Z","shell.execute_reply":"2025-11-23T10:07:18.965211Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 2 – Dataset paths (UPDATE DATA_ROOT)\n# ============================================================\n# TODO: Change this path to your actual dataset root\nDATA_ROOT = Path(\"/kaggle/input/recodai-luc-scientific-image-forgery-detection\")\n\nTRAIN_IMG_DIR = DATA_ROOT / \"train_images\"\nTRAIN_MASK_DIR = DATA_ROOT / \"train_masks\"          # .npy files\nSUPP_IMG_DIR   = DATA_ROOT / \"supplemental_images\"\nSUPP_MASK_DIR  = DATA_ROOT / \"supplemental_masks\"\nTEST_IMG_DIR   = DATA_ROOT / \"test_images\"\nSAMPLE_SUB_PATH = DATA_ROOT / \"sample_submission.csv\"\n\nFORGED_DIR    = TRAIN_IMG_DIR / \"forged\"\nAUTHENTIC_DIR = TRAIN_IMG_DIR / \"authentic\"\n\nforged_paths = sorted(FORGED_DIR.glob(\"*.png\"))\nauth_paths   = sorted(AUTHENTIC_DIR.glob(\"*.png\"))\n\nprint(\"Forged images:\", len(forged_paths))\nprint(\"Authentic images:\", len(auth_paths))\nprint(\"Total train images:\", len(forged_paths) + len(auth_paths))\n\nprint(\"\\nSample mask files in train_masks:\")\nfor p in sorted(TRAIN_MASK_DIR.glob(\"*\"))[:10]:\n    print(\" \", p.name)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-23T10:07:44.412357Z","iopub.execute_input":"2025-11-23T10:07:44.412658Z","iopub.status.idle":"2025-11-23T10:07:44.818608Z","shell.execute_reply.started":"2025-11-23T10:07:44.412636Z","shell.execute_reply":"2025-11-23T10:07:44.817953Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 3 – Build train_df with labels\n# ============================================================\ntrain_records = []\n\nfor p in forged_paths:\n    case_id = p.stem  # \"10.png\" -> \"10\"\n    train_records.append({\"case_id\": case_id, \"path\": p, \"label\": 1})  # forged\n\nfor p in auth_paths:\n    case_id = p.stem\n    train_records.append({\"case_id\": case_id, \"path\": p, \"label\": 0})  # authentic\n\ntrain_df = pd.DataFrame(train_records)\nprint(train_df.head())\nprint(\"\\nLabel counts:\\n\", train_df[\"label\"].value_counts())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-23T10:08:18.859949Z","iopub.execute_input":"2025-11-23T10:08:18.860233Z","iopub.status.idle":"2025-11-23T10:08:18.899310Z","shell.execute_reply.started":"2025-11-23T10:08:18.860212Z","shell.execute_reply":"2025-11-23T10:08:18.898623Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 4 – Image size analysis (optional but useful)\n# ============================================================\ndef get_image_size(path):\n    with Image.open(path) as img:\n        return img.size  # (w, h)\n\nsizes = train_df[\"path\"].map(get_image_size)\ntrain_df[\"width\"] = sizes.map(lambda s: s[0])\ntrain_df[\"height\"] = sizes.map(lambda s: s[1])\n\nprint(train_df[[\"width\", \"height\"]].describe())\n\nplt.figure(figsize=(6,4))\nplt.hist(train_df[\"width\"], bins=30, alpha=0.6, label=\"width\")\nplt.hist(train_df[\"height\"], bins=30, alpha=0.6, label=\"height\")\nplt.legend()\nplt.title(\"Image size distribution\")\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-23T10:08:36.267367Z","iopub.execute_input":"2025-11-23T10:08:36.267655Z","iopub.status.idle":"2025-11-23T10:09:12.714451Z","shell.execute_reply.started":"2025-11-23T10:08:36.267634Z","shell.execute_reply":"2025-11-23T10:09:12.713861Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 5 – Visualize some forged & authentic images\n# ============================================================\ndef show_examples(df, label, n=6, title=\"\"):\n    subset = df[df.label == label].sample(min(n, len(df[df.label == label])))\n    cols = n // 2\n    fig, axes = plt.subplots(2, cols, figsize=(3*cols, 6))\n    axes = axes.flatten()\n    for ax, (_, row) in zip(axes, subset.iterrows()):\n        img = cv2.imread(str(row[\"path\"]))\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        ax.imshow(img)\n        ax.set_title(f\"{row['case_id']} (label={row['label']})\")\n        ax.axis(\"off\")\n    plt.suptitle(title)\n    plt.tight_layout()\n    plt.show()\n\nshow_examples(train_df, 1, n=6, title=\"Forged examples\")\nshow_examples(train_df, 0, n=6, title=\"Authentic examples\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-23T10:09:29.867410Z","iopub.execute_input":"2025-11-23T10:09:29.868088Z","iopub.status.idle":"2025-11-23T10:09:33.179348Z","shell.execute_reply.started":"2025-11-23T10:09:29.868066Z","shell.execute_reply":"2025-11-23T10:09:33.178513Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 6 – .npy mask loader\n# ============================================================\ndef build_mask_for_case(case_id: str, mask_dir: Path, image_size=None):\n    \"\"\"\n    Loads .npy mask for a given case_id and returns a binary mask.\n    - case_id: '10', '10015', ...\n    - mask_dir: path to train_masks\n    - image_size: (H, W) for resizing to match image\n    \"\"\"\n    case_id = str(case_id)\n    mask_path = mask_dir / f\"{case_id}.npy\"\n\n    if not mask_path.exists():\n        # no mask => authentic\n        if image_size is None:\n            return None\n        return np.zeros(image_size, dtype=np.uint8)\n\n    arr = np.load(mask_path)  # can be HxW or HxWxN\n\n    if arr.ndim == 3:\n        # merge channels/regions\n        arr = arr.max(axis=-1)\n\n    mask = (arr > 0).astype(np.uint8)\n\n    if image_size is not None and mask.shape != image_size:\n        h, w = image_size\n        mask = cv2.resize(mask, (w, h), interpolation=cv2.INTER_NEAREST)\n\n    return mask\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-23T10:10:13.532360Z","iopub.execute_input":"2025-11-23T10:10:13.532977Z","iopub.status.idle":"2025-11-23T10:10:13.538180Z","shell.execute_reply.started":"2025-11-23T10:10:13.532954Z","shell.execute_reply":"2025-11-23T10:10:13.537498Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 7 – Sanity check: overlay a few masks on forged images\n# ============================================================\ndef show_image_with_mask(row, alpha=0.4):\n    img = cv2.imread(str(row[\"path\"]))\n    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    h, w = img.shape[:2]\n\n    mask = build_mask_for_case(row[\"case_id\"], TRAIN_MASK_DIR, image_size=(h, w))\n    if mask is None:\n        print(f\"No mask for {row['case_id']}\")\n        return\n\n    fig, ax = plt.subplots(1, 1, figsize=(6, 6))\n    ax.imshow(img)\n    ax.imshow(mask, alpha=alpha, cmap=\"Reds\")\n    ax.set_title(f\"case_id {row['case_id']}\")\n    ax.axis(\"off\")\n    plt.show()\n\nsample_forged = train_df[train_df.label == 1].sample(3, random_state=SEED)\nfor _, row in sample_forged.iterrows():\n    show_image_with_mask(row)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-23T10:10:17.700437Z","iopub.execute_input":"2025-11-23T10:10:17.700907Z","iopub.status.idle":"2025-11-23T10:10:20.372047Z","shell.execute_reply.started":"2025-11-23T10:10:17.700884Z","shell.execute_reply":"2025-11-23T10:10:20.371241Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 8 – Albumentations transforms (Albumentations 2.0.8)\n# ============================================================\nIMG_SIZE = 512  # try 512 first; can go 768 if GPU memory is fine\n\ntrain_transform = A.Compose([\n    A.Resize(IMG_SIZE, IMG_SIZE),\n\n    A.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.5),\n    A.RandomRotate90(p=0.5),\n\n    A.GaussNoise(noise_limit=(10, 50), p=0.3),\n    A.ISONoise(p=0.3),\n\n    A.OneOf([\n        A.MotionBlur(blur_limit=(3, 5), p=1.0),\n        A.GaussianBlur(blur_limit=(3, 5), p=1.0),\n        A.MedianBlur(blur_limit=3, p=1.0),\n    ], p=0.3),\n\n    A.RandomBrightnessContrast(p=0.4),\n\n    A.ElasticTransform(\n        alpha=1.0,\n        sigma=50.0,\n        approximate=True,\n        p=0.2,\n    ),\n\n    A.Normalize(mean=(0.485, 0.456, 0.406),\n                std=(0.229, 0.224, 0.225)),\n    ToTensorV2(),\n])\n\nvalid_transform = A.Compose([\n    A.Resize(IMG_SIZE, IMG_SIZE),\n    A.Normalize(mean=(0.485, 0.456, 0.406),\n                std=(0.229, 0.224, 0.225)),\n    ToTensorV2(),\n])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-23T10:11:00.728445Z","iopub.execute_input":"2025-11-23T10:11:00.728924Z","iopub.status.idle":"2025-11-23T10:11:00.743141Z","shell.execute_reply.started":"2025-11-23T10:11:00.728901Z","shell.execute_reply":"2025-11-23T10:11:00.742513Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 9 – Dataset class\n# ============================================================\nclass ForgerySegmentationDataset(Dataset):\n    def __init__(self, df, mask_dir, transform=None):\n        self.df = df.reset_index(drop=True)\n        self.mask_dir = mask_dir\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img_path = row[\"path\"]\n        case_id = row[\"case_id\"]\n\n        img = cv2.imread(str(img_path))\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        h, w = img.shape[:2]\n\n        mask = build_mask_for_case(case_id, self.mask_dir, image_size=(h, w))\n        if mask is None:\n            mask = np.zeros((h, w), dtype=np.uint8)\n\n        if self.transform is not None:\n            augmented = self.transform(image=img, mask=mask)\n            img = augmented[\"image\"]\n            mask = augmented[\"mask\"]\n\n        if isinstance(mask, torch.Tensor):\n            mask = mask.unsqueeze(0).float()\n        else:\n            mask = torch.from_numpy(mask).unsqueeze(0).float()\n\n        return img, mask\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-23T10:11:15.765055Z","iopub.execute_input":"2025-11-23T10:11:15.765334Z","iopub.status.idle":"2025-11-23T10:11:15.772179Z","shell.execute_reply.started":"2025-11-23T10:11:15.765316Z","shell.execute_reply":"2025-11-23T10:11:15.771463Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 10 – Manual stratified train/val split (NO sklearn)\n# ============================================================\ntrain_df[\"has_mask\"] = train_df[\"label\"]  # forged=1, authentic=0\n\ndef stratified_split(df, label_col, test_size=0.2, random_state=42):\n    train_parts = []\n    val_parts = []\n    rng = np.random.RandomState(random_state)\n\n    for label, group in df.groupby(label_col):\n        n_total = len(group)\n        n_val = int(round(test_size * n_total))\n        n_val = max(1, min(n_val, n_total - 1))  # keep at least 1 in train & val\n\n        val_idx = rng.choice(group.index.values, size=n_val, replace=False)\n        val_part = group.loc[val_idx]\n        train_part = group.drop(val_idx)\n\n        train_parts.append(train_part)\n        val_parts.append(val_part)\n\n    train_df_tr = pd.concat(train_parts).sample(frac=1.0, random_state=random_state).reset_index(drop=True)\n    train_df_val = pd.concat(val_parts).sample(frac=1.0, random_state=random_state).reset_index(drop=True)\n    return train_df_tr, train_df_val\n\ntrain_df_tr, train_df_val = stratified_split(train_df, \"has_mask\", test_size=0.2, random_state=SEED)\n\nprint(\"Train size:\", len(train_df_tr))\nprint(\"Valid size:\", len(train_df_val))\nprint(\"Train mask ratio:\\n\", train_df_tr[\"has_mask\"].value_counts(normalize=True))\nprint(\"Valid mask ratio:\\n\", train_df_val[\"has_mask\"].value_counts(normalize=True))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-23T10:12:29.451081Z","iopub.execute_input":"2025-11-23T10:12:29.451386Z","iopub.status.idle":"2025-11-23T10:12:29.472259Z","shell.execute_reply.started":"2025-11-23T10:12:29.451365Z","shell.execute_reply":"2025-11-23T10:12:29.471605Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 11 – Dataloaders\n# ============================================================\nBATCH_SIZE = 4  # adjust based on GPU\n\ntrain_dataset = ForgerySegmentationDataset(train_df_tr, TRAIN_MASK_DIR, transform=train_transform)\nvalid_dataset = ForgerySegmentationDataset(train_df_val, TRAIN_MASK_DIR, transform=valid_transform)\n\ntrain_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE,\n                          shuffle=True, num_workers=4, pin_memory=True)\nvalid_loader = DataLoader(valid_dataset, batch_size=BATCH_SIZE,\n                          shuffle=False, num_workers=4, pin_memory=True)\n\nprint(\"Train batches:\", len(train_loader), \"Valid batches:\", len(valid_loader))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-23T10:13:03.394222Z","iopub.execute_input":"2025-11-23T10:13:03.394509Z","iopub.status.idle":"2025-11-23T10:13:03.402093Z","shell.execute_reply.started":"2025-11-23T10:13:03.394488Z","shell.execute_reply":"2025-11-23T10:13:03.401402Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 12 – Batch sanity check: images + masks\n# ============================================================\nimgs, masks = next(iter(train_loader))\nprint(\"Image batch shape:\", imgs.shape)\nprint(\"Mask batch shape:\", masks.shape)\nprint(\"Unique mask values:\", masks.unique())\n\nimgs_np = imgs.permute(0,2,3,1).cpu().numpy()\nmasks_np = masks.squeeze(1).cpu().numpy()\n\nplt.figure(figsize=(10,10))\nfor i in range(min(4, imgs_np.shape[0])):\n    plt.subplot(4,2,2*i+1)\n    plt.imshow(imgs_np[i])\n    plt.title(\"Image\")\n    plt.axis(\"off\")\n\n    plt.subplot(4,2,2*i+2)\n    plt.imshow(masks_np[i], cmap=\"gray\")\n    plt.title(\"Mask\")\n    plt.axis(\"off\")\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-23T10:13:22.495119Z","iopub.execute_input":"2025-11-23T10:13:22.495611Z","iopub.status.idle":"2025-11-23T10:13:24.535868Z","shell.execute_reply.started":"2025-11-23T10:13:22.495586Z","shell.execute_reply":"2025-11-23T10:13:24.535048Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 13 – U-Net model (SMP)\n# ============================================================\nENCODER = \"resnet34\"\nENCODER_WEIGHTS = \"imagenet\"\n\nmodel = smp.Unet(\n    encoder_name=ENCODER,\n    encoder_weights=ENCODER_WEIGHTS,\n    in_channels=3,\n    classes=1\n).to(DEVICE)\n\nparams = sum(p.numel() for p in model.parameters() if p.requires_grad)\nprint(f\"Trainable parameters: {params/1e6:.2f}M\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-23T10:13:42.092605Z","iopub.execute_input":"2025-11-23T10:13:42.093302Z","iopub.status.idle":"2025-11-23T10:13:44.248887Z","shell.execute_reply.started":"2025-11-23T10:13:42.093273Z","shell.execute_reply":"2025-11-23T10:13:44.248091Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 14 – Losses & metric (BCE + Dice)\n# ============================================================\nclass DiceLoss(nn.Module):\n    def __init__(self, smooth=1.0):\n        super().__init__()\n        self.smooth = smooth\n\n    def forward(self, logits, targets):\n        probs = torch.sigmoid(logits)\n        targets = targets.float()\n        num = 2 * (probs * targets).sum(dim=(2,3)) + self.smooth\n        den = probs.sum(dim=(2,3)) + targets.sum(dim=(2,3)) + self.smooth\n        dice = 1 - (num / den)\n        return dice.mean()\n\nbce_loss = nn.BCEWithLogitsLoss()\ndice_loss = DiceLoss()\n\ndef combined_loss(logits, targets, bce_weight=0.5, dice_weight=0.5):\n    return bce_weight * bce_loss(logits, targets) + dice_weight * dice_loss(logits, targets)\n\ndef dice_coef(logits, targets, threshold=0.5):\n    probs = torch.sigmoid(logits)\n    preds = (probs > threshold).float()\n    targets = targets.float()\n    num = 2 * (preds * targets).sum(dim=(2,3))\n    den = preds.sum(dim=(2,3)) + targets.sum(dim=(2,3)) + 1e-7\n    dice = (num / den).mean().item()\n    return dice\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-23T10:13:59.223436Z","iopub.execute_input":"2025-11-23T10:13:59.224084Z","iopub.status.idle":"2025-11-23T10:13:59.230907Z","shell.execute_reply.started":"2025-11-23T10:13:59.224057Z","shell.execute_reply":"2025-11-23T10:13:59.230291Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 15 – Optimizer & scheduler\n# ============================================================\nLEARNING_RATE = 1e-4\n\noptimizer = torch.optim.AdamW(model.parameters(), lr=LEARNING_RATE, weight_decay=1e-6)\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer, mode='max', factor=0.5, patience=2, verbose=True\n)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-23T10:14:12.225959Z","iopub.execute_input":"2025-11-23T10:14:12.226542Z","iopub.status.idle":"2025-11-23T10:14:12.238241Z","shell.execute_reply.started":"2025-11-23T10:14:12.226522Z","shell.execute_reply":"2025-11-23T10:14:12.237585Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 16 – Train & valid loops\n# ============================================================\ndef train_one_epoch(model, loader, optimizer):\n    model.train()\n    epoch_loss = 0.0\n    epoch_dice = 0.0\n    \n    for imgs, masks in tqdm(loader, desc=\"Train\", leave=False):\n        imgs = imgs.to(DEVICE)\n        masks = masks.to(DEVICE)\n\n        optimizer.zero_grad()\n        logits = model(imgs)\n        loss = combined_loss(logits, masks)\n        loss.backward()\n        optimizer.step()\n\n        epoch_loss += loss.item() * imgs.size(0)\n        epoch_dice += dice_coef(logits.detach(), masks.detach()) * imgs.size(0)\n    \n    n = len(loader.dataset)\n    return epoch_loss / n, epoch_dice / n\n\n@torch.no_grad()\ndef valid_one_epoch(model, loader):\n    model.eval()\n    epoch_loss = 0.0\n    epoch_dice = 0.0\n    \n    for imgs, masks in tqdm(loader, desc=\"Valid\", leave=False):\n        imgs = imgs.to(DEVICE)\n        masks = masks.to(DEVICE)\n\n        logits = model(imgs)\n        loss = combined_loss(logits, masks)\n\n        epoch_loss += loss.item() * imgs.size(0)\n        epoch_dice += dice_coef(logits, masks) * imgs.size(0)\n    \n    n = len(loader.dataset)\n    return epoch_loss / n, epoch_dice / n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-23T10:14:31.673979Z","iopub.execute_input":"2025-11-23T10:14:31.674283Z","iopub.status.idle":"2025-11-23T10:14:31.681694Z","shell.execute_reply.started":"2025-11-23T10:14:31.674261Z","shell.execute_reply":"2025-11-23T10:14:31.680682Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 17 – Training loop\n# ============================================================\nEPOCHS = 25\nbest_dice = 0.0\nbest_model_path = \"/kaggle/working/unet_resnet34_best.pth\"\n\nfor epoch in range(1, EPOCHS + 1):\n    print(f\"\\nEpoch {epoch}/{EPOCHS}\")\n    train_loss, train_dice = train_one_epoch(model, train_loader, optimizer)\n    val_loss, val_dice = valid_one_epoch(model, valid_loader)\n    \n    scheduler.step(val_dice)\n    \n    print(f\"  Train loss: {train_loss:.4f} | Train dice: {train_dice:.4f}\")\n    print(f\"  Valid loss: {val_loss:.4f} | Valid dice: {val_dice:.4f}\")\n    \n    if val_dice > best_dice:\n        best_dice = val_dice\n        torch.save(model.state_dict(), best_model_path)\n        print(f\"  🔥 New best model saved with dice {best_dice:.4f}\")\n    \n    gc.collect()\n    if DEVICE == \"cuda\":\n        torch.cuda.empty_cache()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-23T10:14:52.137318Z","iopub.execute_input":"2025-11-23T10:14:52.137929Z","iopub.status.idle":"2025-11-23T12:20:25.662555Z","shell.execute_reply.started":"2025-11-23T10:14:52.137904Z","shell.execute_reply":"2025-11-23T12:20:25.661347Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 18 – RLE encode\n# ============================================================\ndef rle_encode(mask):\n    \"\"\"\n    mask: 2D numpy array, values {0,1}\n    returns run-length encoding as space-separated string\n    \"\"\"\n    pixels = mask.T.flatten()\n    pixels = np.concatenate([[0], pixels, [0]])\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 1\n    runs[1::2] = runs[1::2] - runs[::2]\n    return \" \".join(str(x) for x in runs)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-23T12:21:52.807186Z","iopub.execute_input":"2025-11-23T12:21:52.808000Z","iopub.status.idle":"2025-11-23T12:21:52.812965Z","shell.execute_reply.started":"2025-11-23T12:21:52.807973Z","shell.execute_reply":"2025-11-23T12:21:52.812157Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 19 – Predict mask probability for a single test image\n# ============================================================\n@torch.no_grad()\ndef predict_mask(model, img_path):\n    model.eval()\n    \n    img = cv2.imread(str(img_path))\n    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    orig_h, orig_w = img.shape[:2]\n    \n    dummy_mask = np.zeros((orig_h, orig_w), dtype=np.uint8)\n    transformed = valid_transform(image=img, mask=dummy_mask)\n    img_t = transformed[\"image\"].unsqueeze(0).to(DEVICE)\n    \n    logits = model(img_t)\n    probs = torch.sigmoid(logits)[0,0].cpu().numpy()  # [H,W] (resized)\n    \n    prob_full = cv2.resize(probs, (orig_w, orig_h), interpolation=cv2.INTER_LINEAR)\n    return prob_full\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-23T12:22:12.827115Z","iopub.execute_input":"2025-11-23T12:22:12.827643Z","iopub.status.idle":"2025-11-23T12:22:12.832913Z","shell.execute_reply.started":"2025-11-23T12:22:12.827622Z","shell.execute_reply":"2025-11-23T12:22:12.832179Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 20 – Convert prob map to annotation (\"authentic\" or RLE)\n# ============================================================\nPROB_THRESHOLD = 0.5\nAREA_THRESHOLD = 100  # tune this on validation set\n\ndef mask_to_annotation(prob_map):\n    bin_mask = (prob_map > PROB_THRESHOLD).astype(np.uint8)\n    area = bin_mask.sum()\n    if area < AREA_THRESHOLD:\n        return \"authentic\"\n    else:\n        return rle_encode(bin_mask)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-23T12:22:42.763666Z","iopub.execute_input":"2025-11-23T12:22:42.764245Z","iopub.status.idle":"2025-11-23T12:22:42.768643Z","shell.execute_reply.started":"2025-11-23T12:22:42.764223Z","shell.execute_reply":"2025-11-23T12:22:42.767787Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 21 – Generate submission CSV\n# ============================================================\n# Load best model\nmodel.load_state_dict(torch.load(best_model_path, map_location=DEVICE))\nmodel.to(DEVICE)\nmodel.eval()\n\nsub_df = pd.read_csv(SAMPLE_SUB_PATH)\nprint(sub_df.head())\n\ndef test_image_path(case_id):\n    return TEST_IMG_DIR / f\"{case_id}.png\"\n\nannotations = []\n\nfor idx, row in tqdm(sub_df.iterrows(), total=len(sub_df)):\n    case_id = str(row[\"case_id\"])\n    img_path = test_image_path(case_id)\n    prob_map = predict_mask(model, img_path)\n    annotation = mask_to_annotation(prob_map)\n    annotations.append(annotation)\n\nsub_df[\"annotation\"] = annotations\nsub_df.to_csv(\"/kaggle/working/submission_unet_baseline.csv\", index=False)\nsub_df.head()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-23T12:23:49.012449Z","iopub.execute_input":"2025-11-23T12:23:49.013091Z","iopub.status.idle":"2025-11-23T12:23:49.231416Z","shell.execute_reply.started":"2025-11-23T12:23:49.013067Z","shell.execute_reply":"2025-11-23T12:23:49.230789Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}