{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"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"}],"dockerImageVersionId":31236,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install -q segmentation-models-pytorch pytorch-lightning albumentations timm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-22T10:09:57.199509Z","iopub.execute_input":"2025-12-22T10:09:57.199797Z","iopub.status.idle":"2025-12-22T10:10:03.001129Z","shell.execute_reply.started":"2025-12-22T10:09:57.199772Z","shell.execute_reply":"2025-12-22T10:10:03.000367Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, gc, cv2, glob, numpy as np, pandas as pd, torch, torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nimport pytorch_lightning as pl\nfrom pytorch_lightning.callbacks import ModelCheckpoint\nimport segmentation_models_pytorch as smp\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom tqdm.auto import tqdm\nfrom sklearn.model_selection import StratifiedKFold\nimport matplotlib.pyplot as plt","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-22T10:10:03.002949Z","iopub.execute_input":"2025-12-22T10:10:03.003199Z","iopub.status.idle":"2025-12-22T10:10:22.724314Z","shell.execute_reply.started":"2025-12-22T10:10:03.003151Z","shell.execute_reply":"2025-12-22T10:10:22.723657Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CFG:\n    seed = 42\n    img_size = 512\n    base_path = \"/kaggle/input/recodai-luc-scientific-image-forgery-detection/\"\n    encoder = \"efficientnet-b4\"\n    batch_size = 8\n    lr = 3e-4\n    epochs = 12\n    precision = \"16-mixed\"\n    num_workers = 3 \n\npl.seed_everything(CFG.seed)\nos.makedirs('/kaggle/working/weights', exist_ok=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-22T10:10:22.725216Z","iopub.execute_input":"2025-12-22T10:10:22.725784Z","iopub.status.idle":"2025-12-22T10:10:22.736413Z","shell.execute_reply.started":"2025-12-22T10:10:22.725744Z","shell.execute_reply":"2025-12-22T10:10:22.735878Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_train_df(base_path):\n    auth_paths = glob.glob(os.path.join(base_path, \"train_images/authentic/*.png\"))\n    forged_paths = glob.glob(os.path.join(base_path, \"train_images/forged/*.png\"))\n    \n    data = []\n    for p in auth_paths:\n        data.append({'case_id': os.path.splitext(os.path.basename(p))[0], 'file_path': p, 'is_forged': 0})\n    for p in forged_paths:\n        data.append({'case_id': os.path.splitext(os.path.basename(p))[0], 'file_path': p, 'is_forged': 1})\n        \n    df = pd.DataFrame(data).reset_index(drop=True)\n    \n    skf = StratifiedKFold(n_splits=5, shuffle=True, random_state=CFG.seed)\n    df['fold'] = -1\n    for fold, (_, val_idx) in enumerate(skf.split(df, df['is_forged'])):\n        df.loc[val_idx, 'fold'] = fold\n        \n    counts = df['is_forged'].value_counts()\n    weights = {0: 1.0/counts[0], 1: 1.0/counts[1]}\n    df['sample_weight'] = df['is_forged'].map(weights)\n    return df\n\ndf = get_train_df(CFG.base_path)\nprint(f\"Dataset Balanced: {df['is_forged'].value_counts().to_dict()}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-22T10:10:22.738078Z","iopub.execute_input":"2025-12-22T10:10:22.738433Z","iopub.status.idle":"2025-12-22T10:10:22.871263Z","shell.execute_reply.started":"2025-12-22T10:10:22.738406Z","shell.execute_reply":"2025-12-22T10:10:22.870367Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ForgeryDataset(Dataset):\n    def __init__(self, df, transform=None, mode='train'):\n        self.df, self.transform, self.mode = df, transform, mode\n        self.mask_dir = os.path.join(CFG.base_path, \"train_masks/\")\n\n    def __len__(self): return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        image = cv2.cvtColor(cv2.imread(row['file_path']), cv2.COLOR_BGR2RGB)\n        if self.mode != 'test':\n            if row['is_forged'] == 0:\n                mask = np.zeros(image.shape[:2], dtype=np.float32)\n            else:\n                m_path = os.path.join(self.mask_dir, f\"{row['case_id']}.npy\")\n                mask = np.load(m_path).astype(np.float32) if os.path.exists(m_path) else np.zeros(image.shape[:2], dtype=np.float32)\n                if mask.ndim == 3: mask = np.max(mask, axis=2)\n                if image.shape[:2] != mask.shape[:2]:\n                    mask = cv2.resize(mask, (image.shape[1], image.shape[0]), interpolation=cv2.INTER_NEAREST)\n            mask = (mask > 0).astype(np.float32)\n            if self.transform:\n                aug = self.transform(image=image, mask=mask)\n                image, mask = aug['image'], aug['mask']\n            return image, mask.unsqueeze(0)\n        if self.transform: image = self.transform(image=image)['image']\n        return image, row['case_id']\n\nclass ForgeryModel(pl.LightningModule):\n    def __init__(self):\n        super().__init__()\n        self.model = smp.UnetPlusPlus(encoder_name=CFG.encoder, encoder_weights=\"imagenet\", in_channels=3, classes=1)\n        self.loss = lambda y_hat, y: 0.5 * smp.losses.DiceLoss(mode='binary', from_logits=True)(y_hat, y) + 0.5 * nn.BCEWithLogitsLoss()(y_hat, y)\n        \n    def forward(self, x): return self.model(x)\n    def training_step(self, batch, batch_idx):\n        return self.loss(self(batch[0]), batch[1])\n    def validation_step(self, batch, batch_idx):\n        preds = (self(batch[0]).sigmoid() > 0.5).float()\n        tp, fp, fn, tn = smp.metrics.get_stats(preds.long(), batch[1].long(), mode='binary')\n        self.log(\"val_f1\", smp.metrics.f1_score(tp, fp, fn, tn, reduction=\"micro\"), prog_bar=True)\n    def configure_optimizers(self): return torch.optim.AdamW(self.parameters(), lr=CFG.lr)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-22T10:10:22.872279Z","iopub.execute_input":"2025-12-22T10:10:22.873199Z","iopub.status.idle":"2025-12-22T10:10:22.884288Z","shell.execute_reply.started":"2025-12-22T10:10:22.873134Z","shell.execute_reply":"2025-12-22T10:10:22.883494Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_trans = A.Compose([A.Resize(CFG.img_size, CFG.img_size), A.HorizontalFlip(p=0.5), A.VerticalFlip(p=0.5), A.Normalize(), ToTensorV2()])\nval_trans = A.Compose([A.Resize(CFG.img_size, CFG.img_size), A.Normalize(), ToTensorV2()])\n\nfor fold in range(5):\n    print(f\"\\n--- Training Fold {fold} ---\")\n    t_df, v_df = df[df['fold']!=fold].reset_index(), df[df['fold']==fold].reset_index()\n    sampler = WeightedRandomSampler(t_df['sample_weight'], len(t_df))\n    \n    train_loader = DataLoader(ForgeryDataset(t_df, train_trans), batch_size=CFG.batch_size, sampler=sampler, num_workers=CFG.num_workers, pin_memory=True)\n    val_loader = DataLoader(ForgeryDataset(v_df, val_trans), batch_size=CFG.batch_size, num_workers=CFG.num_workers, pin_memory=True)\n    \n    model = ForgeryModel()\n    ckpt = ModelCheckpoint(monitor='val_f1', mode='max', filename=f'fold_{fold}', dirpath='/kaggle/working/weights/')\n    pl.Trainer(max_epochs=CFG.epochs, accelerator='gpu', precision=CFG.precision, callbacks=[ckpt]).fit(model, train_loader, val_loader)\n    del model; gc.collect(); torch.cuda.empty_cache()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-22T10:10:22.885257Z","iopub.execute_input":"2025-12-22T10:10:22.88556Z","iopub.status.idle":"2025-12-22T16:11:28.609234Z","shell.execute_reply.started":"2025-12-22T10:10:22.885533Z","shell.execute_reply":"2025-12-22T16:11:28.608504Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_supplemental(base_path):\n    models = [ForgeryModel.load_from_checkpoint(p).cuda().eval() for p in glob.glob('/kaggle/working/weights/*.ckpt')]\n    supp_paths = glob.glob(os.path.join(base_path, \"supplemental_images/*.png\"))[:3]\n    \n    plt.figure(figsize=(18, 6))\n    for i, p in enumerate(supp_paths):\n        img = cv2.cvtColor(cv2.imread(p), cv2.COLOR_BGR2RGB)\n        img_t = val_trans(image=img)['image'].unsqueeze(0).cuda()\n        with torch.no_grad():\n            preds = [torch.sigmoid(m(img_t)).cpu().numpy()[0,0] for m in models]\n        avg_p = cv2.resize(np.mean(preds, axis=0), (img.shape[1], img.shape[0]))\n        \n        plt.subplot(1, 3, i+1); plt.imshow(avg_p, cmap='jet'); plt.title(f\"Supp: {os.path.basename(p)}\")\n        plt.axis('off')\n    plt.tight_layout(); plt.show()\n\nvisualize_supplemental(CFG.base_path)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-22T16:13:14.846721Z","iopub.execute_input":"2025-12-22T16:13:14.847219Z","iopub.status.idle":"2025-12-22T16:13:21.664804Z","shell.execute_reply.started":"2025-12-22T16:13:14.847159Z","shell.execute_reply":"2025-12-22T16:13:21.664152Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def rle_encode(mask):\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[::2]\n    return ' '.join(str(x) for x in runs)\n\ntest_files = glob.glob(os.path.join(CFG.base_path, \"test_images/*.png\"))\nresults = []\nmodels = [ForgeryModel.load_from_checkpoint(p).cuda().eval() for p in glob.glob('/kaggle/working/weights/*.ckpt')]\n\nfor p in tqdm(test_files):\n    case_id = os.path.splitext(os.path.basename(p))[0]\n    img = cv2.imread(p)\n    h, w = img.shape[:2]\n    img_rgb = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    \n    # Simple TTA: Original + Horizontal Flip\n    probs = []\n    for view in [img_rgb, cv2.flip(img_rgb, 1)]:\n        img_t = val_trans(image=view)['image'].unsqueeze(0).cuda()\n        with torch.no_grad():\n            for m in models:\n                p_map = torch.sigmoid(m(img_t)).cpu().numpy()[0,0]\n                if len(probs) % 2 != 0: p_map = cv2.flip(p_map, 1) # Un-flip\n                probs.append(cv2.resize(p_map, (w, h)))\n    \n    mask = (np.mean(probs, axis=0) > 0.5).astype(np.uint8)\n    rle = rle_encode(mask)\n    results.append({'case_id': case_id, 'annotation': rle if rle != \"\" else \"authentic\"})\n\npd.DataFrame(results).to_csv('submission.csv', index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-22T16:13:37.338229Z","iopub.execute_input":"2025-12-22T16:13:37.338559Z","iopub.status.idle":"2025-12-22T16:13:42.799718Z","shell.execute_reply.started":"2025-12-22T16:13:37.33853Z","shell.execute_reply":"2025-12-22T16:13:42.799143Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def tta_inference_supplemental(base_path, models, output_csv=\"supplemental_submission.csv\"):\n    # Target only forged supplemental images for this test\n    supp_files = glob.glob(os.path.join(base_path, \"supplemental_images/*.png\"))\n    results = []\n    \n    print(f\"Running TTA Inference on {len(supp_files)} supplemental images...\")\n    \n    for p in tqdm(supp_files):\n        case_id = os.path.splitext(os.path.basename(p))[0]\n        img = cv2.imread(p)\n        h, w = img.shape[:2]\n        img_rgb = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        \n        # Define TTA Views: [Original, Horizontal Flip, Vertical Flip]\n        tta_views = [\n            img_rgb, \n            cv2.flip(img_rgb, 1), \n            cv2.flip(img_rgb, 0)\n        ]\n        \n        all_probs = []\n        \n        for idx, view in enumerate(tta_views):\n            img_t = val_trans(image=view)['image'].unsqueeze(0).cuda()\n            \n            with torch.no_grad():\n                for m in models:\n                    prob_map = torch.sigmoid(m(img_t)).cpu().numpy()[0, 0]\n                    \n                    # Inverse Transform (Un-flip) the TTA views\n                    if idx == 1: # Un-flip Horizontal\n                        prob_map = cv2.flip(prob_map, 1)\n                    elif idx == 2: # Un-flip Vertical\n                        prob_map = cv2.flip(prob_map, 0)\n                    \n                    all_probs.append(cv2.resize(prob_map, (w, h)))\n        \n        # Average all model predictions and TTA views\n        final_prob = np.mean(all_probs, axis=0)\n        \n        # Apply Threshold (0.5) and encode\n        mask = (final_prob > 0.5).astype(np.uint8)\n        rle = rle_encode(mask)\n        \n        results.append({\n            'case_id': case_id, \n            'annotation': rle if rle != \"\" else \"authentic\"\n        })\n\n    # Save to CSV\n    supp_df = pd.DataFrame(results)\n    supp_df.to_csv(output_csv, index=False)\n    print(f\"Saved supplemental results to {output_csv}\")\n    return supp_df\n\n# Execute TTA on Supplemental\nsupp_results_df = tta_inference_supplemental(CFG.base_path, models)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-22T16:13:48.160732Z","iopub.execute_input":"2025-12-22T16:13:48.161023Z","iopub.status.idle":"2025-12-22T16:14:33.383394Z","shell.execute_reply.started":"2025-12-22T16:13:48.160997Z","shell.execute_reply":"2025-12-22T16:14:33.382695Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def rle_decode(mask_rle, shape):\n    '''\n    mask_rle: run-length as string formated (start length)\n    shape: (height, width) of array to return \n    Returns numpy array, 1 - mask, 0 - background\n    '''\n    if mask_rle == \"authentic\" or not isinstance(mask_rle, str):\n        return np.zeros(shape, dtype=np.uint8)\n    \n    s = mask_rle.split()\n    starts, lengths = [np.asarray(x, dtype=int) for x in (s[0:][::2], s[1:][::2])]\n    starts -= 1\n    ends = starts + lengths\n    img = np.zeros(shape[0] * shape[1], dtype=np.uint8)\n    for lo, hi in zip(starts, ends):\n        img[lo:hi] = 1\n    return img.reshape(shape).T # Transpose matches the standard RLE encoding for this competition","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-22T16:19:52.902572Z","iopub.execute_input":"2025-12-22T16:19:52.903151Z","iopub.status.idle":"2025-12-22T16:19:52.909119Z","shell.execute_reply.started":"2025-12-22T16:19:52.903118Z","shell.execute_reply":"2025-12-22T16:19:52.908427Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_tta_check(df, base_path, num_samples=3):\n    samples = df.sample(num_samples)\n    plt.figure(figsize=(18, 6 * num_samples))\n    \n    for i, (_, row) in enumerate(samples.iterrows()):\n        # Find image path (assuming forged subfolder)\n        img_path = os.path.join(base_path, f\"supplemental_images/{row['case_id']}.png\")\n        img = cv2.cvtColor(cv2.imread(img_path), cv2.COLOR_BGR2RGB)\n        \n        # Decode the RLE we just saved to verify it's correct\n        mask = rle_decode(row['annotation'], img.shape[:2]) # Use the rle_decode function provided earlier\n        \n        plt.subplot(num_samples, 2, i*2+1)\n        plt.imshow(img)\n        plt.title(f\"Image: {row['case_id']}\")\n        plt.axis('off')\n        \n        plt.subplot(num_samples, 2, i*2+2)\n        plt.imshow(mask, cmap='gray')\n        plt.title(\"TTA + Ensemble Decoded Mask\")\n        plt.axis('off')\n        \n    plt.tight_layout()\n    plt.show()\n\nvisualize_tta_check(supp_results_df, CFG.base_path)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-22T16:19:56.632708Z","iopub.execute_input":"2025-12-22T16:19:56.633031Z","iopub.status.idle":"2025-12-22T16:20:01.810958Z","shell.execute_reply.started":"2025-12-22T16:19:56.633002Z","shell.execute_reply":"2025-12-22T16:20:01.809863Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def optimize_threshold(base_path, models, thresholds=[0.3, 0.35, 0.4, 0.45, 0.5, 0.55, 0.6]):\n    supp_imgs = glob.glob(os.path.join(base_path, \"supplemental_images/forged/*.png\"))\n    \n    # Pre-calculate probabilities for all images to save time\n    all_preds = []\n    all_gts = []\n    \n    print(\"Pre-calculating TTA probabilities for optimization...\")\n    for p in tqdm(supp_imgs):\n        case_id = os.path.splitext(os.path.basename(p))[0]\n        mask_search = glob.glob(os.path.join(base_path, f\"supplemental_masks/**/{case_id}.npy\"), recursive=True)\n        if not mask_search: continue\n            \n        img = cv2.imread(p)\n        h, w = img.shape[:2]\n        \n        # Ground Truth\n        gt = np.load(mask_search[0]).astype(np.uint8)\n        if gt.ndim == 3: gt = np.max(gt, axis=2)\n        all_gts.append((gt > 0).astype(np.uint8))\n        \n        # TTA Ensemble Probabilities\n        img_rgb = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        views = [img_rgb, cv2.flip(img_rgb, 1), cv2.flip(img_rgb, 0)]\n        probs = []\n        for idx, view in enumerate(views):\n            img_t = val_trans(image=view)['image'].unsqueeze(0).cuda()\n            with torch.no_grad():\n                for m in models:\n                    p_map = torch.sigmoid(m(img_t)).cpu().numpy()[0,0]\n                    if idx == 1: p_map = cv2.flip(p_map, 1)\n                    if idx == 2: p_map = cv2.flip(p_map, 0)\n                    probs.append(cv2.resize(p_map, (w, h)))\n        all_preds.append(np.mean(probs, axis=0))\n\n    # Evaluate each threshold\n    best_f1 = 0\n    best_threshold = 0.5\n    \n    print(\"\\nThreshold Search:\")\n    for t in thresholds:\n        f1_list = []\n        for pred_prob, gt_mask in zip(all_preds, all_gts):\n            pred_mask = (pred_prob > t).astype(np.uint8)\n            inter = np.logical_and(pred_mask, gt_mask).sum()\n            f1 = (2. * inter) / (pred_mask.sum() + gt_mask.sum() + 1e-7)\n            f1_list.append(f1)\n        \n        mean_f1 = np.mean(f1_list)\n        print(f\"Threshold {t:.2f} -> Mean F1: {mean_f1:.4f}\")\n        \n        if mean_f1 > best_f1:\n            best_f1 = mean_f1\n            best_threshold = t\n            \n    print(f\"\\n✅ Optimal Threshold Found: {best_threshold} (F1: {best_f1:.4f})\")\n    return best_threshold\n\n# Run the optimizer\nBEST_THRESHOLD = optimize_threshold(CFG.base_path, models)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-22T16:20:41.356775Z","iopub.execute_input":"2025-12-22T16:20:41.357072Z","iopub.status.idle":"2025-12-22T16:20:41.385681Z","shell.execute_reply.started":"2025-12-22T16:20:41.357047Z","shell.execute_reply":"2025-12-22T16:20:41.385091Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def score_supplemental(base_path, models):\n    supp_imgs = glob.glob(os.path.join(base_path, \"supplemental_images/forged/*.png\"))\n    f1_scores = []\n    \n    for p in tqdm(supp_imgs):\n        case_id = os.path.splitext(os.path.basename(p))[0]\n        mask_search = glob.glob(os.path.join(base_path, f\"supplemental_masks/**/{case_id}.npy\"), recursive=True)\n        if not mask_search: continue\n            \n        img = cv2.imread(p)\n        h, w = img.shape[:2]\n        \n        # Load and fix GT Mask shape\n        gt_mask = np.squeeze(np.load(mask_search[0]).astype(np.float32))\n        if gt_mask.ndim == 3: gt_mask = np.max(gt_mask, axis=2)\n        gt_mask = (gt_mask > 0).astype(np.uint8)\n        \n        # TTA Inference\n        img_rgb = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        views = [img_rgb, cv2.flip(img_rgb, 1), cv2.flip(img_rgb, 0)]\n        probs = []\n        for idx, view in enumerate(views):\n            img_t = val_trans(image=view)['image'].unsqueeze(0).cuda()\n            with torch.no_grad():\n                for m in models:\n                    p_map = torch.sigmoid(m(img_t)).cpu().numpy()[0,0]\n                    if idx == 1: p_map = cv2.flip(p_map, 1)\n                    if idx == 2: p_map = cv2.flip(p_map, 0)\n                    probs.append(cv2.resize(p_map, (w, h)))\n        \n        pred_prob = np.mean(probs, axis=0)\n        pred_mask = (pred_prob > 0.5).astype(np.uint8)\n\n        # FINAL SHAPE ALIGNMENT\n        if pred_mask.shape != gt_mask.shape:\n            pred_mask = cv2.resize(pred_mask, (gt_mask.shape[1], gt_mask.shape[0]), interpolation=cv2.INTER_NEAREST)\n        \n        intersection = np.logical_and(pred_mask, gt_mask).sum()\n        dice = (2. * intersection) / (pred_mask.sum() + gt_mask.sum() + 1e-7)\n        f1_scores.append(dice)\n\n    print(f\"\\n FIXED F1 SCORE: {np.mean(f1_scores):.4f}\")\n    return np.mean(f1_scores)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-22T16:23:09.257488Z","iopub.execute_input":"2025-12-22T16:23:09.257803Z","iopub.status.idle":"2025-12-22T16:23:09.267828Z","shell.execute_reply.started":"2025-12-22T16:23:09.257774Z","shell.execute_reply":"2025-12-22T16:23:09.267104Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_comparison(base_path, models, num_samples=2):\n    supp_imgs = glob.glob(os.path.join(base_path, \"supplemental_images/*.png\"))[:num_samples]\n    \n    plt.figure(figsize=(18, 5 * num_samples))\n    for i, p in enumerate(supp_imgs):\n        case_id = os.path.splitext(os.path.basename(p))[0]\n        img = cv2.cvtColor(cv2.imread(p), cv2.COLOR_BGR2RGB)\n        \n        # Ground Truth\n        mask_path = glob.glob(os.path.join(base_path, f\"**/supplemental_masks/**/{case_id}.npy\"), recursive=True)[0]\n        gt = np.load(mask_path)\n        if gt.ndim == 3: gt = np.max(gt, axis=2)\n        \n        # TTA Prediction (Simplified for viz)\n        img_t = val_trans(image=img)['image'].unsqueeze(0).cuda()\n        with torch.no_grad():\n            preds = [torch.sigmoid(m(img_t)).cpu().numpy()[0,0] for m in models]\n        pred_avg = cv2.resize(np.mean(preds, axis=0), (img.shape[1], img.shape[0]))\n\n        plt.subplot(num_samples, 3, i*3+1); plt.imshow(img); plt.title(\"Original Image\"); plt.axis('off')\n        plt.subplot(num_samples, 3, i*3+2); plt.imshow(gt, cmap='gray'); plt.title(\"Ground Truth Mask\"); plt.axis('off')\n        plt.subplot(num_samples, 3, i*3+3); plt.imshow(pred_avg, cmap='jet'); plt.title(\"TTA Ensemble Prediction\"); plt.axis('off')\n    plt.tight_layout(); plt.show()\n\nvisualize_comparison(CFG.base_path, models)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-22T16:23:19.779305Z","iopub.execute_input":"2025-12-22T16:23:19.779604Z","iopub.status.idle":"2025-12-22T16:23:28.596841Z","shell.execute_reply.started":"2025-12-22T16:23:19.779575Z","shell.execute_reply":"2025-12-22T16:23:28.595887Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}