{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":113558,"databundleVersionId":14878066,"sourceType":"competition"},{"sourceId":14420385,"sourceType":"datasetVersion","datasetId":9210245},{"sourceId":290470675,"sourceType":"kernelVersion"}],"dockerImageVersionId":31236,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Cell 0 — Setup + Config\n","metadata":{}},{"cell_type":"code","source":"import os, random, math, gc, csv, re, warnings\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\n\nimport cv2\ncv2.setNumThreads(0)\nfrom PIL import Image\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\n\nfrom sklearn.model_selection import GroupShuffleSplit, GroupKFold\n\nwarnings.filterwarnings(\"ignore\", category=UserWarning)\nwarnings.filterwarnings(\"ignore\", category=FutureWarning)\n\nCOMP_ROOT = Path(\"/kaggle/input/recodai-luc-scientific-image-forgery-detection\")\nTRAIN_AUTH_DIR = COMP_ROOT / \"train_images\" / \"authentic\"\nTRAIN_FORG_DIR = COMP_ROOT / \"train_images\" / \"forged\"\nTRAIN_MASK_DIR = COMP_ROOT / \"train_masks\"\nSUPP_IMG_DIR   = COMP_ROOT / \"supplemental_images\"\nSUPP_MASK_DIR  = COMP_ROOT / \"supplemental_masks\"\nTEST_IMG_DIR   = COMP_ROOT / \"test_images\"\nSAMPLE_SUB_PATH= COMP_ROOT / \"sample_submission.csv\"\n\nWORKDIR = Path(\"/kaggle/working\")\nOUTDIR  = WORKDIR / \"outputs\"\nOUTDIR.mkdir(parents=True, exist_ok=True)\n\nRESNET34_WEIGHTS_PATH = \"/kaggle/input/resnet-train/resnet34-b627a593.pth\"\n\nFOLD0_PATH = WORKDIR / \"best_unet_fold0.pt\"\nFOLD1_PATH = WORKDIR / \"best_unet_fold1.pt\"\nOUT_FOLD0_PATH = OUTDIR / \"best_unet_fold0.pt\"\nOUT_FOLD1_PATH = OUTDIR / \"best_unet_fold1.pt\"\n\nSEED = 42\ndef seed_everything(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.benchmark = True\nseed_everything(SEED)\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nprint(\"CUDA available:\", torch.cuda.is_available(), \"| n_gpus:\", torch.cuda.device_count())\nprint(\"COMP_ROOT:\", COMP_ROOT)\nprint(\"RESNET34_WEIGHTS_PATH exists:\", os.path.exists(RESNET34_WEIGHTS_PATH))\n\nDO_TRAIN = False                                               \nEPOCHS = 8                                                  \nLR = 1e-4                                                   \nWEIGHT_DECAY = 1e-4\n\nCALIB_FRAC = 0.10                                                                                \n\nCROP_SIZE = 256\nBATCH_SIZE = 16 if torch.cuda.is_available() else 4\nNUM_WORKERS = 2\nSAVE_THR = 0.30\nPATIENCE = 4\n\nTILE = 768\nSTRIDE = 384                                              \nUSE_TTA = True                                         \n\nMIN_PIXELS = 12\nTHR_SWEEP = [0.15,0.20,0.25,0.30,0.35,0.40]\nAREA_SWEEP = [10,20,30,40,60]\n\ndef find_in_inputs(filename: str) -> str | None:\n    base = Path(\"/kaggle/input\")\n    if not base.exists():\n        return None\n    hits = list(base.rglob(filename))\n    hits = sorted(hits, key=lambda p: (len(str(p)), str(p)))\n    return str(hits[0]) if hits else None\n\nIN_FOLD0 = find_in_inputs(\"best_unet_fold0.pt\")\nIN_FOLD1 = find_in_inputs(\"best_unet_fold1.pt\")\nprint(\"Discovered in inputs:\", IN_FOLD0, IN_FOLD1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-19T22:02:19.834142Z","iopub.execute_input":"2026-01-19T22:02:19.834586Z","iopub.status.idle":"2026-01-19T22:02:22.714798Z","shell.execute_reply.started":"2026-01-19T22:02:19.834554Z","shell.execute_reply":"2026-01-19T22:02:22.714135Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Cell 1 — Helpers: listing + robust required case_ids + RLE formatting\n","metadata":{}},{"cell_type":"code","source":"\ndef list_pngs(folder: Path):\n    if not folder.exists():\n        return []\n    return sorted([str(p) for p in folder.glob(\"*.png\")])\n\ndef list_npys(folder: Path):\n    if not folder.exists():\n        return []\n    return sorted([str(p) for p in folder.glob(\"*.npy\")])\n\ndef parse_case_id(path_str: str) -> int:\n    return int(Path(path_str).stem)\n\ndef read_required_case_ids(sample_csv_path: Path) -> list[int]:\n    try:\n        ss = pd.read_csv(sample_csv_path)\n        if \"case_id\" in ss.columns and len(ss) > 0:\n            return ss[\"case_id\"].astype(int).tolist()\n    except Exception:\n        pass\n    try:\n        out = []\n        with open(sample_csv_path, newline=\"\", encoding=\"utf-8-sig\") as f:\n            reader = csv.DictReader(f)\n            for row in reader:\n                if \"case_id\" in row and row[\"case_id\"] != \"\":\n                    out.append(int(row[\"case_id\"]))\n        if len(out) > 0:\n            return out\n    except Exception:\n        pass\n    txt = sample_csv_path.read_text(encoding=\"utf-8\", errors=\"ignore\")\n    ids = [int(x) for x in re.findall(r\"\\b\\d+\\b\", txt)]\n    seen, out = set(), []\n    for x in ids:\n        if x not in seen:\n            seen.add(x)\n            out.append(x)\n    return out\n\nrequired_case_ids = read_required_case_ids(SAMPLE_SUB_PATH)\nprint(\"Required case_ids:\", len(required_case_ids), \"| first 10:\", required_case_ids[:10])\n\nRLE_WITH_COMMAS = True\ntry:\n    with open(SAMPLE_SUB_PATH, \"r\", encoding=\"utf-8-sig\") as f:\n        next(f)\n        for _ in range(200):\n            line = f.readline()\n            if not line:\n                break\n            if \"[\" in line and \"]\" in line:\n                RLE_WITH_COMMAS = (\",\" in line)\n                break\nexcept Exception:\n    pass\nprint(\"RLE_WITH_COMMAS inferred:\", RLE_WITH_COMMAS)\n\nprint(\"Train authentic images:\", len(list_pngs(TRAIN_AUTH_DIR)))\nprint(\"Train forged images   :\", len(list_pngs(TRAIN_FORG_DIR)))\nprint(\"Train masks           :\", len(list_npys(TRAIN_MASK_DIR)))\nprint(\"Supplemental images   :\", len(list_pngs(SUPP_IMG_DIR)))\nprint(\"Supplemental masks    :\", len(list_npys(SUPP_MASK_DIR)))\nprint(\"Visible test images here:\", len(list_pngs(TEST_IMG_DIR)))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-19T22:02:22.716029Z","iopub.execute_input":"2026-01-19T22:02:22.716274Z","iopub.status.idle":"2026-01-19T22:02:22.757778Z","shell.execute_reply.started":"2026-01-19T22:02:22.716251Z","shell.execute_reply":"2026-01-19T22:02:22.757251Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Cell 2 — Build dataframe + calibration split\n","metadata":{}},{"cell_type":"code","source":"\ndef build_mask_map(mask_dir: Path) -> dict[int, list[str]]:\n    m = {}\n    if not mask_dir.exists():\n        return m\n    for p in mask_dir.glob(\"*.npy\"):\n        mm = re.match(r\"(\\d+)\", p.stem)\n        if not mm:\n            continue\n        cid = int(mm.group(1))\n        m.setdefault(cid, []).append(str(p))\n    for k in m:\n        m[k] = sorted(m[k])\n    return m\n\nmask_map = {}\nfor d in [TRAIN_MASK_DIR, SUPP_MASK_DIR]:\n    mm = build_mask_map(d)\n    for k, v in mm.items():\n        mask_map.setdefault(k, []).extend(v)\nfor k in mask_map:\n    mask_map[k] = sorted(mask_map[k])\n\ndef make_rows(img_paths: list[str], source: str, forged_default: bool) -> list[dict]:\n    rows = []\n    for p in img_paths:\n        cid = parse_case_id(p)\n        mps = mask_map.get(cid, [])\n        is_forged = bool(forged_default and (len(mps) > 0))\n        rows.append({\"case_id\": cid, \"img_path\": p, \"mask_paths\": mps, \"is_forged\": is_forged, \"source\": source})\n    return rows\n\nauth_paths = list_pngs(TRAIN_AUTH_DIR)\nforg_paths = list_pngs(TRAIN_FORG_DIR)\nsupp_paths = list_pngs(SUPP_IMG_DIR)\n\nrows = []\nrows += make_rows(auth_paths, \"train_authentic\", forged_default=False)\nrows += make_rows(forg_paths, \"train_forged\", forged_default=True)\nrows += make_rows(supp_paths, \"supplemental\", forged_default=True)\n\ndf = pd.DataFrame(rows)\ndf.loc[(df[\"source\"]==\"train_forged\") & (df[\"mask_paths\"].apply(len)==0), \"is_forged\"] = False\n\nprint(\"Total image rows:\", len(df))\nprint(\"Forged rows:\", int(df.is_forged.sum()), \"| Authentic rows:\", int((~df.is_forged).sum()))\n\ngss = GroupShuffleSplit(n_splits=1, test_size=CALIB_FRAC, random_state=SEED)\ntrain_idx, calib_idx = next(gss.split(df, groups=df[\"case_id\"]))\ntrain_base_df = df.iloc[train_idx].reset_index(drop=True)\ncalib_df      = df.iloc[calib_idx].reset_index(drop=True)\nprint(\"Train-base rows:\", len(train_base_df), \"| Calib rows:\", len(calib_df),\n      \"| Calib forged:\", int(calib_df.is_forged.sum()), \"| Calib authentic:\", int((~calib_df.is_forged).sum()))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-19T22:02:22.758661Z","iopub.execute_input":"2026-01-19T22:02:22.758911Z","iopub.status.idle":"2026-01-19T22:02:22.847079Z","shell.execute_reply.started":"2026-01-19T22:02:22.758878Z","shell.execute_reply":"2026-01-19T22:02:22.846401Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Cell 3 — Image/mask utils + augment + postprocess + RLE\n","metadata":{}},{"cell_type":"code","source":"\ndef read_image(path: str) -> np.ndarray:\n    try:\n        img = cv2.imread(path, cv2.IMREAD_COLOR)\n        if img is not None and img.size > 0:\n            return cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    except Exception:\n        pass\n    try:\n        return np.array(Image.open(path).convert(\"RGB\"))\n    except Exception:\n        return np.zeros((256,256,3), dtype=np.uint8)\n\ndef read_and_merge_masks(mask_paths: list[str], shape_hw: tuple[int,int]) -> np.ndarray:\n    H, W = int(shape_hw[0]), int(shape_hw[1])\n    if H<=0 or W<=0:\n        return np.zeros((1,1), dtype=np.uint8)\n    out = np.zeros((H,W), dtype=np.uint8)\n    for mp in (mask_paths or []):\n        try:\n            m = np.load(mp)\n        except Exception:\n            continue\n        m = np.asarray(m)\n        if m.ndim == 3: m = m[...,0]\n        if m.size == 0: continue\n        m = (m>0).astype(np.uint8)\n        if m.shape != out.shape:\n            if m.shape[0] <= 0 or m.shape[1] <= 0: \n                continue\n            m = cv2.resize(m, (W,H), interpolation=cv2.INTER_NEAREST)\n        out = np.maximum(out, m)\n    return out.astype(np.uint8)\n\ndef pad_to(img: np.ndarray, mask: np.ndarray, min_size: int):\n    H,W = img.shape[:2]\n    pad_h = max(0, min_size-H)\n    pad_w = max(0, min_size-W)\n    if pad_h==0 and pad_w==0:\n        return img, mask\n    top = pad_h//2; bottom = pad_h-top\n    left = pad_w//2; right = pad_w-left\n    img2 = cv2.copyMakeBorder(img, top,bottom,left,right, borderType=cv2.BORDER_REFLECT_101)\n    mask2= cv2.copyMakeBorder(mask, top,bottom,left,right, borderType=cv2.BORDER_CONSTANT, value=0)\n    return img2, mask2\n\ndef random_crop(img: np.ndarray, mask: np.ndarray, crop: int):\n    H,W = img.shape[:2]\n    y0 = 0 if H<=crop else random.randint(0, H-crop)\n    x0 = 0 if W<=crop else random.randint(0, W-crop)\n    return img[y0:y0+crop, x0:x0+crop], mask[y0:y0+crop, x0:x0+crop]\n\ndef crop_around_mask(img: np.ndarray, mask: np.ndarray, crop: int):\n    ys,xs = np.where(mask>0)\n    if len(ys)==0:\n        return random_crop(img, mask, crop)\n    H,W = img.shape[:2]\n    i = random.randrange(len(ys))\n    cy,cx = int(ys[i]), int(xs[i])\n    jitter = crop//8\n    cy = int(np.clip(cy + random.randint(-jitter,jitter), 0, H-1))\n    cx = int(np.clip(cx + random.randint(-jitter,jitter), 0, W-1))\n    y0 = int(np.clip(cy - crop//2, 0, max(0,H-crop)))\n    x0 = int(np.clip(cx - crop//2, 0, max(0,W-crop)))\n    return img[y0:y0+crop, x0:x0+crop], mask[y0:y0+crop, x0:x0+crop]\n\ndef augment(img: np.ndarray, mask: np.ndarray):\n    if random.random()<0.5:\n        img = np.fliplr(img).copy(); mask = np.fliplr(mask).copy()\n    if random.random()<0.5:\n        img = np.flipud(img).copy(); mask = np.flipud(mask).copy()\n    k = random.randint(0,3)\n    if k:\n        img = np.rot90(img,k).copy(); mask = np.rot90(mask,k).copy()\n    if random.random()<0.30:\n        alpha = 1.0 + random.uniform(-0.15,0.15)\n        beta  = random.uniform(-20,20)\n        img = np.clip(alpha*img + beta, 0,255).astype(np.uint8)\n    return img, mask\n\ndef postprocess_mask(mask01: np.ndarray, min_area: int = 30, do_close: bool = True):\n    mask01 = (mask01>0).astype(np.uint8)\n    if mask01.sum()==0:\n        return mask01\n    if do_close:\n        k = cv2.getStructuringElement(cv2.MORPH_ELLIPSE,(3,3))\n        mask01 = cv2.morphologyEx(mask01, cv2.MORPH_CLOSE, k, iterations=1)\n    num, lab = cv2.connectedComponents(mask01, connectivity=8)\n    out = np.zeros_like(mask01)\n    for i in range(1,num):\n        area = int((lab==i).sum())\n        if area >= int(min_area):\n            out[lab==i] = 1\n    return out.astype(np.uint8)\n\ndef rle_encode(mask: np.ndarray) -> str:\n    m = (mask>0).astype(np.uint8)\n    pixels = m.T.flatten()\n    pixels = np.concatenate([[0], pixels, [0]])\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 1\n    runs[1::2] -= runs[::2]\n    if RLE_WITH_COMMAS:\n        return \"[\" + \", \".join(str(x) for x in runs) + \"]\"\n    else:\n        return \"[\" + \" \".join(str(x) for x in runs) + \"]\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-19T22:02:22.848572Z","iopub.execute_input":"2026-01-19T22:02:22.848828Z","iopub.status.idle":"2026-01-19T22:02:22.867923Z","shell.execute_reply.started":"2026-01-19T22:02:22.848808Z","shell.execute_reply":"2026-01-19T22:02:22.867416Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Cell 4 — Dataset + loader factory\n","metadata":{}},{"cell_type":"code","source":"\nclass ForgeryDataset(Dataset):\n    def __init__(self, df: pd.DataFrame, train: bool, crop: int):\n        self.df = df.reset_index(drop=True)\n        self.train = train\n        self.crop = int(crop)\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        for _ in range(10):\n            row = self.df.iloc[idx]\n            img = read_image(row.img_path)\n            H,W = img.shape[:2]\n            if H<=0 or W<=0:\n                idx = random.randint(0, len(self.df)-1)\n                continue\n\n            if bool(row.is_forged):\n                mask = read_and_merge_masks(row.mask_paths, (H,W))\n            else:\n                mask = np.zeros((H,W), dtype=np.uint8)\n\n            img, mask = pad_to(img, mask, self.crop)\n\n            if self.train:\n                if bool(row.is_forged) and mask.sum()>0 and random.random()<0.80:\n                    img, mask = crop_around_mask(img, mask, self.crop)\n                else:\n                    img, mask = random_crop(img, mask, self.crop)\n                img, mask = augment(img, mask)\n            else:\n                H2,W2 = img.shape[:2]\n                y0 = max(0, (H2-self.crop)//2)\n                x0 = max(0, (W2-self.crop)//2)\n                img = img[y0:y0+self.crop, x0:x0+self.crop]\n                mask= mask[y0:y0+self.crop, x0:x0+self.crop]\n\n            x = torch.from_numpy(img).float().permute(2,0,1)/255.0\n            y = torch.from_numpy((mask>0).astype(np.float32)).unsqueeze(0)\n            return x,y,int(row.case_id)\n\n        x = torch.zeros((3,self.crop,self.crop), dtype=torch.float32)\n        y = torch.zeros((1,self.crop,self.crop), dtype=torch.float32)\n        return x,y,-1\n\ndef make_loader(df_part: pd.DataFrame, train: bool):\n    ds = ForgeryDataset(df_part, train=train, crop=CROP_SIZE)\n    return DataLoader(ds, batch_size=BATCH_SIZE, shuffle=train,\n                      num_workers=NUM_WORKERS, pin_memory=True, drop_last=train)\n\nsm = make_loader(train_base_df.sample(min(128, len(train_base_df)), random_state=SEED), train=True)\nxb, yb, _ = next(iter(sm))\nprint(\"Smoke batch:\", xb.shape, yb.shape, \"x range\", (float(xb.min()), float(xb.max())))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-19T22:02:22.868588Z","iopub.execute_input":"2026-01-19T22:02:22.868783Z","iopub.status.idle":"2026-01-19T22:02:25.268192Z","shell.execute_reply.started":"2026-01-19T22:02:22.868764Z","shell.execute_reply":"2026-01-19T22:02:25.267445Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Cell 5 — Model definition (ResNet34 U-Net offline)\n","metadata":{}},{"cell_type":"code","source":"\nfrom torchvision.models import resnet34\n\ndef load_resnet34_offline(enc: nn.Module, weights_path: str) -> bool:\n    if weights_path and os.path.exists(weights_path):\n        sd = torch.load(weights_path, map_location=\"cpu\")\n        if isinstance(sd, dict) and \"state_dict\" in sd:\n            sd = sd[\"state_dict\"]\n        if isinstance(sd, dict):\n            new_sd = {}\n            for k,v in sd.items():\n                kk = k\n                for pref in [\"module.\", \"model.\", \"encoder.\"]:\n                    if kk.startswith(pref):\n                        kk = kk[len(pref):]\n                new_sd[kk] = v\n            sd = new_sd\n        try:\n            enc.load_state_dict(sd, strict=False)\n            return True\n        except Exception as e:\n            print(\"Failed loading offline resnet34:\", e)\n    return False\n\nclass UpBlock(nn.Module):\n    def __init__(self, in_ch, skip_ch, out_ch):\n        super().__init__()\n        self.conv1 = nn.Conv2d(in_ch + skip_ch, out_ch, 3, padding=1)\n        self.bn1 = nn.BatchNorm2d(out_ch)\n        self.conv2 = nn.Conv2d(out_ch, out_ch, 3, padding=1)\n        self.bn2 = nn.BatchNorm2d(out_ch)\n    def forward(self, x, skip):\n        x = F.interpolate(x, size=skip.shape[-2:], mode=\"bilinear\", align_corners=False)\n        x = torch.cat([x, skip], dim=1)\n        x = F.relu(self.bn1(self.conv1(x)))\n        x = F.relu(self.bn2(self.conv2(x)))\n        return x\n\nclass ResNet34UNet(nn.Module):\n    def __init__(self, resnet34_weights_path: str = \"\"):\n        super().__init__()\n        enc = resnet34(weights=None)\n        ok = load_resnet34_offline(enc, resnet34_weights_path)\n        print(\"Encoder pretrained loaded:\", ok)\n\n        self.enc0 = nn.Sequential(enc.conv1, enc.bn1, enc.relu)\n        self.pool = enc.maxpool\n        self.enc1 = enc.layer1\n        self.enc2 = enc.layer2\n        self.enc3 = enc.layer3\n        self.enc4 = enc.layer4\n\n        self.bridge = nn.Sequential(\n            nn.Conv2d(512, 512, 3, padding=1),\n            nn.BatchNorm2d(512),\n            nn.ReLU(inplace=True),\n        )\n\n        self.up4 = UpBlock(512, 256, 256)\n        self.up3 = UpBlock(256, 128, 128)\n        self.up2 = UpBlock(128,  64,  64)\n        self.up1 = UpBlock( 64,  64,  64)\n\n        self.head = nn.Conv2d(64, 1, kernel_size=1)\n\n    def forward(self, x):\n        x0 = self.enc0(x)\n        x1 = self.enc1(self.pool(x0))\n        x2 = self.enc2(x1)\n        x3 = self.enc3(x2)\n        x4 = self.enc4(x3)\n\n        b = self.bridge(x4)\n\n        d4 = self.up4(b, x3)\n        d3 = self.up3(d4, x2)\n        d2 = self.up2(d3, x1)\n        d1 = self.up1(d2, x0)\n\n        out = self.head(d1)\n        out = F.interpolate(out, size=x.shape[-2:], mode=\"bilinear\", align_corners=False)\n        return out\n\nprint(\"Model ready.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-19T22:02:25.269303Z","iopub.execute_input":"2026-01-19T22:02:25.269591Z","iopub.status.idle":"2026-01-19T22:02:28.753701Z","shell.execute_reply.started":"2026-01-19T22:02:25.269563Z","shell.execute_reply":"2026-01-19T22:02:28.753076Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Cell 6 — Train/load folds + build ensemble\n","metadata":{}},{"cell_type":"code","source":"\nuse_amp = torch.cuda.is_available()\nscaler = torch.amp.GradScaler(\"cuda\", enabled=use_amp)\n\nPOS_WEIGHT = torch.tensor([5.0], device=device)\nbce = nn.BCEWithLogitsLoss(pos_weight=POS_WEIGHT)\n\ndef dice_loss(logits, targets, eps=1e-6):\n    probs = torch.sigmoid(logits)\n    num = 2*(probs*targets).sum(dim=(2,3))\n    den = (probs+targets).sum(dim=(2,3)) + eps\n    return (1-(num/den)).mean()\n\ndef dice_score(pred, target, eps=1e-6):\n    inter = (pred*target).sum(dim=(2,3))\n    union = pred.sum(dim=(2,3)) + target.sum(dim=(2,3))\n    score = torch.where(union==0, torch.ones_like(union), (2*inter+eps)/(union+eps))\n    return score.mean().item()\n\n@torch.no_grad()\ndef eval_dice_single(model, loader, thr=0.3):\n    model.eval()\n    scores=[]\n    for x,y,_ in loader:\n        x = x.to(device, non_blocking=True)\n        y = y.to(device, non_blocking=True)\n        with torch.amp.autocast(\"cuda\", enabled=use_amp):\n            logits = model(x)\n        probs = torch.sigmoid(logits)\n        pred = (probs>thr).float()\n        scores.append(dice_score(pred,y))\n    return float(np.mean(scores)) if scores else 0.0\n\ndef save_state(model, path: Path):\n    st = model.module.state_dict() if hasattr(model,\"module\") else model.state_dict()\n    torch.save(st, str(path))\n\ndef load_state(model, path: str):\n    st = torch.load(path, map_location=device)\n    mm = model.module if hasattr(model,\"module\") else model\n    mm.load_state_dict(st, strict=True)\n    mm.eval()\n    return mm\n\ndef train_one_fold(fold_id: int, tr_df: pd.DataFrame, va_df: pd.DataFrame, out_path: Path):\n    train_loader = make_loader(tr_df, train=True)\n    val_loader   = make_loader(va_df, train=False)\n\n    model = ResNet34UNet(resnet34_weights_path=RESNET34_WEIGHTS_PATH).to(device)\n    if torch.cuda.device_count()>1:\n        model = nn.DataParallel(model)\n\n    opt = torch.optim.AdamW(model.parameters(), lr=LR, weight_decay=WEIGHT_DECAY)\n    sched = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=max(1,EPOCHS))\n\n    best=-1.0\n    bad=0\n\n    for epoch in range(1, EPOCHS+1):\n        model.train()\n        losses=[]\n        for x,y,_ in train_loader:\n            x = x.to(device, non_blocking=True)\n            y = y.to(device, non_blocking=True)\n            opt.zero_grad(set_to_none=True)\n\n            with torch.amp.autocast(\"cuda\", enabled=use_amp):\n                logits = model(x)\n                probs = torch.sigmoid(logits)\n                pt = probs*y + (1-probs)*(1-y)\n                focal = ((1-pt)**2).mean()\n                loss = bce(logits,y) + dice_loss(logits,y) + 0.05*focal\n\n            scaler.scale(loss).backward()\n            scaler.step(opt)\n            scaler.update()\n            losses.append(float(loss.item()))\n\n        vd = eval_dice_single(model, val_loader, thr=SAVE_THR)\n        lr_now = opt.param_groups[0][\"lr\"]\n        print(f\"[Fold {fold_id}] Epoch {epoch}/{EPOCHS} | lr={lr_now:.2e} | train_loss={np.mean(losses):.4f} | val_dice@{SAVE_THR}={vd:.4f}\")\n\n        if vd > best + 1e-6:\n            best = vd\n            bad = 0\n            save_state(model, out_path)\n            print(f\"✅ [Fold {fold_id}] Saved -> {out_path}\")\n        else:\n            bad += 1\n            if bad >= PATIENCE:\n                print(f\"[Fold {fold_id}] Early stopping.\")\n                break\n\n        sched.step()\n        gc.collect()\n        if torch.cuda.is_available():\n            torch.cuda.empty_cache()\n\n    print(f\"[Fold {fold_id}] Best val_dice@{SAVE_THR}: {best:.4f}\")\n    return str(out_path)\n\nfold0_ckpt = str(FOLD0_PATH) if FOLD0_PATH.exists() else (IN_FOLD0 if IN_FOLD0 else None)\nfold1_ckpt = str(FOLD1_PATH) if FOLD1_PATH.exists() else (IN_FOLD1 if IN_FOLD1 else None)\n\ngkf = GroupKFold(n_splits=2)\nsplits = list(gkf.split(train_base_df, groups=train_base_df[\"case_id\"]))\n(tr0, va0) = splits[0]\n(tr1, va1) = splits[1]\n\nif DO_TRAIN or (not fold0_ckpt):\n    fold0_ckpt = train_one_fold(0, train_base_df.iloc[tr0], train_base_df.iloc[va0], FOLD0_PATH)\n    if FOLD0_PATH.exists():\n        import shutil; shutil.copy2(str(FOLD0_PATH), str(OUT_FOLD0_PATH))\n        print(\"✅ Copied fold0 to outputs:\", OUT_FOLD0_PATH)\n\nif DO_TRAIN or (not fold1_ckpt):\n    fold1_ckpt = train_one_fold(1, train_base_df.iloc[tr1], train_base_df.iloc[va1], FOLD1_PATH)\n    if FOLD1_PATH.exists():\n        import shutil; shutil.copy2(str(FOLD1_PATH), str(OUT_FOLD1_PATH))\n        print(\"✅ Copied fold1 to outputs:\", OUT_FOLD1_PATH)\n\nprint(\"Fold checkpoints:\", fold0_ckpt, fold1_ckpt)\n\nm0 = ResNet34UNet(resnet34_weights_path=RESNET34_WEIGHTS_PATH).to(device)\nm1 = ResNet34UNet(resnet34_weights_path=RESNET34_WEIGHTS_PATH).to(device)\nm0 = load_state(m0, fold0_ckpt)\nm1 = load_state(m1, fold1_ckpt)\nmodels_ens = [m0, m1]\nprint(\"Ensemble ready.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-19T22:02:28.754484Z","iopub.execute_input":"2026-01-19T22:02:28.754881Z","iopub.status.idle":"2026-01-19T22:37:26.734337Z","shell.execute_reply.started":"2026-01-19T22:02:28.754857Z","shell.execute_reply":"2026-01-19T22:37:26.733451Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Cell 7 — Calibrate THR + MIN_AREA on calibration set using ensemble (crop-level proxy)\n","metadata":{}},{"cell_type":"code","source":"\ncalib_small = calib_df.sample(min(len(calib_df), 800), random_state=SEED).reset_index(drop=True)\ncalib_loader = make_loader(calib_small, train=False)\n\n@torch.no_grad()\ndef eval_dice_post_ensemble(models, loader, thr=0.3, min_area=30):\n    scores=[]\n    for x,y,_ in loader:\n        x = x.to(device, non_blocking=True)\n        y_np = y.numpy()\n        with torch.amp.autocast(\"cuda\", enabled=use_amp):\n            ps = None\n            for m in models:\n                logits = m(x)\n                pr = torch.sigmoid(logits).detach().float()\n                ps = pr if ps is None else (ps+pr)\n            probs = (ps/len(models)).cpu().numpy()             \n\n        for i in range(probs.shape[0]):\n            pm = (probs[i,0] > thr).astype(np.uint8)\n            pm = postprocess_mask(pm, min_area=min_area, do_close=True)\n            ym = (y_np[i,0] > 0.5).astype(np.uint8)\n            inter = int((pm*ym).sum())\n            union = int(pm.sum()+ym.sum())\n            scores.append(1.0 if union==0 else (2*inter+1e-6)/(union+1e-6))\n    return float(np.mean(scores)) if scores else 0.0\n\nbest=(-1.0,None,None)\nfor a in AREA_SWEEP:\n    for t in THR_SWEEP:\n        sc = eval_dice_post_ensemble(models_ens, calib_loader, thr=t, min_area=a)\n        print(f\"thr {t:.2f} | min_area {a:>2} | calib_dice_post {sc:.4f}\")\n        if sc > best[0]:\n            best=(sc,t,a)\n\nBEST_CALIB=float(best[0])\nBEST_THR=float(best[1]) if best[1] is not None else 0.30\nBEST_MIN_AREA=int(best[2]) if best[2] is not None else 30\nprint(\"BEST_THR:\", BEST_THR, \"| BEST_MIN_AREA:\", BEST_MIN_AREA, \"| BEST_CALIB:\", BEST_CALIB)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-19T22:37:26.735924Z","iopub.execute_input":"2026-01-19T22:37:26.736259Z","iopub.status.idle":"2026-01-19T22:42:42.595187Z","shell.execute_reply.started":"2026-01-19T22:37:26.736216Z","shell.execute_reply":"2026-01-19T22:42:42.594262Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Cell 8 — Ensemble tiled inference + robust submission.csv (row-count safe)\n","metadata":{}},{"cell_type":"code","source":"\n@torch.no_grad()\ndef predict_prob_full_ens(img: np.ndarray, models: list[nn.Module],\n                          crop: int = 512, stride: int = 384, batch: int = 8, tta: bool = False) -> np.ndarray:\n    H0, W0 = img.shape[:2]\n    img_pad, _ = pad_to(img, np.zeros((H0,W0), dtype=np.uint8), crop)\n    Hp, Wp = img_pad.shape[:2]\n\n    ys = list(range(0, max(1, Hp - crop + 1), stride))\n    xs = list(range(0, max(1, Wp - crop + 1), stride))\n    if ys[-1] != Hp - crop: ys.append(Hp - crop)\n    if xs[-1] != Wp - crop: xs.append(Wp - crop)\n\n    acc = np.zeros((Hp, Wp), dtype=np.float32)\n    cnt = np.zeros((Hp, Wp), dtype=np.float32)\n\n    tiles, coords = [], []\n\n    def forward_models(x_tensor):\n        with torch.amp.autocast(\"cuda\", enabled=torch.cuda.is_available()):\n            ps = None\n            for m in models:\n                logits = m(x_tensor)\n                pr = torch.sigmoid(logits).detach().float()\n                ps = pr if ps is None else (ps + pr)\n            pr_mean = (ps/len(models))[:,0]                 \n        return pr_mean\n\n    def run_batch(tiles, coords):\n        if not tiles:\n            return\n        x = torch.stack(tiles, dim=0).to(device)\n        probs = forward_models(x).cpu().numpy()\n\n        if tta:\n            x_h = torch.flip(x, dims=[3])\n            p_h = torch.flip(forward_models(x_h), dims=[2]).cpu().numpy()\n\n            x_v = torch.flip(x, dims=[2])\n            p_v = torch.flip(forward_models(x_v), dims=[1]).cpu().numpy()\n\n            x_hv = torch.flip(x, dims=[2,3])\n            p_hv = torch.flip(forward_models(x_hv), dims=[1,2]).cpu().numpy()\n\n            probs = (probs + p_h + p_v + p_hv) / 4.0\n\n        for (y0, x0), pr in zip(coords, probs):\n            acc[y0:y0+crop, x0:x0+crop] += pr\n            cnt[y0:y0+crop, x0:x0+crop] += 1.0\n\n    for y0 in ys:\n        for x0 in xs:\n            tile = img_pad[y0:y0+crop, x0:x0+crop]\n            t = torch.from_numpy(tile).float().permute(2,0,1)/255.0\n            tiles.append(t)\n            coords.append((y0,x0))\n            if len(tiles) >= batch:\n                run_batch(tiles, coords)\n                tiles, coords = [], []\n\n    run_batch(tiles, coords)\n\n    prob = acc / np.maximum(cnt, 1e-6)\n    y0 = (Hp - H0)//2\n    x0 = (Wp - W0)//2\n    prob = prob[y0:y0+H0, x0:x0+W0]\n    return prob\n\n@torch.no_grad()\ndef predict_mask_for_path_ens(img_path: str, thr: float, min_area: int) -> np.ndarray:\n    img = read_image(img_path)\n    prob = predict_prob_full_ens(img, models_ens, crop=TILE, stride=STRIDE, batch=8, tta=USE_TTA)\n    mask = (prob > thr).astype(np.uint8)\n    mask = postprocess_mask(mask, min_area=min_area, do_close=True)\n    return mask\n\nvisible_test_paths = sorted([str(p) for p in TEST_IMG_DIR.glob(\"*.png\")])\npath_map = {parse_case_id(p): p for p in visible_test_paths}\nprint(\"Visible test images here:\", len(path_map), \"| first 5:\", list(path_map.items())[:5])\n\ncase_ids = required_case_ids\nprint(\"Submitting rows:\", len(case_ids))\n\nTHR = float(BEST_THR) if \"BEST_THR\" in globals() else 0.30\nMIN_AREA = int(BEST_MIN_AREA) if \"BEST_MIN_AREA\" in globals() else 30\n\nrows=[]\nfailed=0\nmissing=0\n\nfor cid in case_ids:\n    cid = int(cid)\n    p = path_map.get(cid)\n    if p is None:\n        missing += 1\n        rows.append((cid, \"authentic\"))\n        continue\n    try:\n        m = predict_mask_for_path_ens(p, thr=THR, min_area=MIN_AREA)\n        ann = \"authentic\" if int(m.sum()) < int(MIN_PIXELS) else rle_encode(m)\n    except Exception:\n        failed += 1\n        ann = \"authentic\"\n    rows.append((cid, ann))\n\nsub = pd.DataFrame(rows, columns=[\"case_id\",\"annotation\"])\nassert list(sub.columns)==[\"case_id\",\"annotation\"]\nassert len(sub)==len(case_ids), f\"Row count mismatch: {len(sub)} vs {len(case_ids)}\"\nassert sub[\"case_id\"].isna().sum()==0\nassert sub[\"annotation\"].isna().sum()==0\n\nout_path = str(WORKDIR / \"submission.csv\")\nsub.to_csv(out_path, index=False)\n\nprint(\"Saved:\", out_path)\nprint(\"Rows:\", len(sub), \"| failed:\", failed, \"| missing_images_here:\", missing)\nprint(sub.head(10))\nprint(\"authentic:\", int((sub['annotation']=='authentic').sum()), \"| rle:\", int((sub['annotation']!='authentic').sum()))\nprint(\"Using THR:\", THR, \"| MIN_PIXELS:\", MIN_PIXELS, \"| MIN_AREA:\", MIN_AREA, \"| TILE:\", TILE, \"| STRIDE:\", STRIDE, \"| USE_TTA:\", USE_TTA)\n\nchk = pd.read_csv(out_path)\nprint(\"Sanity check:\", chk.shape, \"| cols:\", list(chk.columns))\nassert chk.shape[1]==2 and chk.shape[0]==len(sub)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-19T22:42:42.596733Z","iopub.execute_input":"2026-01-19T22:42:42.597062Z","iopub.status.idle":"2026-01-19T22:42:48.176281Z","shell.execute_reply.started":"2026-01-19T22:42:42.597032Z","shell.execute_reply":"2026-01-19T22:42:48.175539Z"}},"outputs":[],"execution_count":null}]}