{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.11.13"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":117682,"databundleVersionId":14443416,"sourceType":"competition"},{"sourceId":11855918,"sourceType":"datasetVersion","datasetId":7449695},{"sourceId":284413864,"sourceType":"kernelVersion"}],"dockerImageVersionId":31192,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Vesuvius Challenge - FastAI U-Net 2.5D","metadata":{"_uuid":"f23c99f4-afa1-449a-afee-4062c85d3c5e","_cell_guid":"108448e2-9f73-4f73-b0da-5603b4b8ba9b","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nfrom PIL import Image\nimport zipfile\nfrom pathlib import Path\nfrom fastai.vision.all import *\nimport torch.nn.functional as F\nimport timm\n\n# --- CONFIGURATION ---\nPATH = Path('/kaggle/input/vesuvius-challenge-surface-detection')\n\nBS = 16             \nIMG_SIZE = 128       \nARCH = 'convnext_small.fb_in22k'","metadata":{"_uuid":"804a11d5-7bbb-43e7-96cb-0c191edb9aee","_cell_guid":"86f721ba-52c1-40a4-b751-2db6ae1257cd","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-12-07T09:25:43.143951Z","iopub.execute_input":"2025-12-07T09:25:43.144112Z","iopub.status.idle":"2025-12-07T09:25:56.248857Z","shell.execute_reply.started":"2025-12-07T09:25:43.144096Z","shell.execute_reply":"2025-12-07T09:25:56.248117Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load the data\ndf = pd.read_csv('/kaggle/input/vesuvius-challenge-surface-detection/train.csv')\n\n# Group by 'scroll_id' and count the number of 'id's (fragments)\nstats_df = df.groupby('scroll_id')['id'].count().reset_index()\n\n# Rename columns for clarity\nstats_df.rename(columns={'id': 'fragment_count'}, inplace=True)\n\n# Display the stats\nprint(stats_df)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T09:25:56.250247Z","iopub.execute_input":"2025-12-07T09:25:56.250456Z","iopub.status.idle":"2025-12-07T09:25:56.277376Z","shell.execute_reply.started":"2025-12-07T09:25:56.250439Z","shell.execute_reply":"2025-12-07T09:25:56.276687Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import glob\nfrom pathlib import Path\nimport random\n\n# CONFIG\nPROCESSED_DATA = Path('/kaggle/input/fastai-u-net-2-5d-pre-processing/processed_dataset_v1')\nVALID_ID = \"26002\"  # strictly hold this out\n\n# 1. Collect All Files\ntrain_path = PROCESSED_DATA / 'train'\nall_files = []\n\nif train_path.exists():\n    for folder in train_path.iterdir():\n        if folder.is_dir():\n            # Grab images in this scroll folder\n            scroll_files = list((folder / 'images').glob('*.npy'))\n            all_files.extend(scroll_files)\n            print(f\"Scroll {folder.name}: Found {len(scroll_files)} images\")\n\nall_files = sorted(all_files)\n\n# 2. Strict Split\ntrain_files = []\nvalid_files = []\n\nfor f in all_files:\n    # Get scroll ID from parent folder name\n    scroll_id = f.parent.parent.name\n    \n    if scroll_id == VALID_ID:\n        valid_files.append(f)\n    else:\n        train_files.append(f)\n\n# 3. SAFETY CHECK\nprint(f\"Final Split -> Train: {len(train_files)} | Valid: {len(valid_files)}\")\n\nif len(train_files) == 0:\n    print(\"\\n⚠️ WARNING: Train set is empty! logic skipped 26002, but found no other scrolls.\")\n    print(\"If you ONLY have 26002, you MUST split it or training will fail.\")\n    print(\"Uncomment the lines below to force a split on 26002 if it's your only data:\")\n    \n    # UNCOMMENT THIS ONLY IF YOU HAVE NO OTHER DATA\n    # random.shuffle(valid_files)\n    # split_idx = int(0.8 * len(valid_files))\n    # train_files = valid_files[:split_idx]\n    # valid_files = valid_files[split_idx:]\n    # print(f\"Fixed -> Train: {len(train_files)} | Valid: {len(valid_files)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T07:46:14.421314Z","iopub.execute_input":"2025-12-07T07:46:14.421602Z","iopub.status.idle":"2025-12-07T07:46:14.776679Z","shell.execute_reply.started":"2025-12-07T07:46:14.421584Z","shell.execute_reply":"2025-12-07T07:46:14.775844Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class FastVesuviusDataset(torch.utils.data.Dataset):\n    def __init__(self, file_paths):\n        self.files = file_paths\n\n    def __len__(self):\n        return len(self.files)\n\n    def __getitem__(self, idx):\n        img_path = str(self.files[idx])\n        \n        # 1. Load Image\n        image = np.load(img_path)\n        \n        # 2. Transpose (C, H, W)\n        # Fix shape: if channels are last, move them to front\n        if image.ndim == 3 and image.shape[-1] <= 65:\n            image = image.transpose(2, 0, 1)\n        \n        # Capture dimensions (C, H, W) -> H, W\n        h, w = image.shape[1], image.shape[2]\n        \n        # 3. Load Mask (Only looking for .npy in 'labels')\n        # Path: .../train/26002/images/abc.npy -> .../train/26002/labels/abc.npy\n        mask_path = img_path.replace('images', 'masks')\n\n        mask = np.load(mask_path)\n            \n        return torch.from_numpy(image).float(), torch.from_numpy(mask).long()\n\n# --- 2. TEST DATASET (For Inference) ---\n# Keeps original logic to read raw TIFFs for the hidden test set\nclass VesuviusTestDataset(torch.utils.data.Dataset):\n    def __init__(self, df):\n        self.df = df\n        self.samples = []\n        for _, row in self.df.iterrows():\n            img_id = row['id']\n            # Only checking size to init range, actual load is in getitem\n            path_to_img = PATH/'test_images'/f'{img_id}.tif'\n            with Image.open(path_to_img) as img:\n                n_frames = img.n_frames\n            for i in range(n_frames):\n                self.samples.append((img_id, i))\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        img_id, slice_idx = self.samples[idx]\n        img_path = PATH/'test_images'/f'{img_id}.tif'\n        \n        # Load 3 neighbors\n        slices = []\n        with Image.open(img_path) as img:\n            n_frames = img.n_frames\n            for offset in [-1, 0, 1]:\n                nid = np.clip(slice_idx + offset, 0, n_frames - 1)\n                img.seek(nid)\n                slices.append(np.array(img))\n        \n        stack = np.stack(slices, axis=0).astype(np.float32)\n        stack = (stack - stack.mean()) / (stack.std() + 1e-6) # Z-Score\n        \n        tensor_img = torch.tensor(stack)\n        \n        # Resize on the fly\n        tensor_img = F.interpolate(\n            tensor_img.unsqueeze(0), \n            size=(IMG_SIZE, IMG_SIZE), \n            mode='bilinear', \n            align_corners=False\n        ).squeeze(0)\n\n        return tensor_img","metadata":{"_uuid":"1b596401-3962-442a-95bf-361ce16fb320","_cell_guid":"f752ff8e-6a67-41a4-aba0-7e2fb1854020","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-12-07T07:46:24.477039Z","iopub.execute_input":"2025-12-07T07:46:24.477331Z","iopub.status.idle":"2025-12-07T07:46:24.488466Z","shell.execute_reply.started":"2025-12-07T07:46:24.477279Z","shell.execute_reply":"2025-12-07T07:46:24.487249Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==========================================\n# 3. CREATE DATASETS & DATALOADERS\n# ==========================================\n\n# Instantiate the Dataset class we defined earlier\ntrain_ds = FastVesuviusDataset(train_files)\nvalid_ds = FastVesuviusDataset(valid_files)\n\n# Define GPU Augmentations (FastAI style)\naug_tfms = [\n    Rotate(p=0.5, max_deg=180),\n    Flip(p=0.5),                \n    Dihedral(p=0.5),            \n    Zoom(max_zoom=1.1, p=0.5),  \n    Warp(magnitude=0.2, p=0.5)  \n]\n\n# Create DataLoaders\ndls = DataLoaders.from_dsets(\n    train_ds, \n    valid_ds, \n    bs=BS, \n    num_workers=4,\n    after_batch=aug_tfms, # Apply augs on GPU\n    device=torch.device('cuda')\n)\n\n# ==========================================\n# 4. VERIFICATION\n# ==========================================\n# Always check a batch to ensure shapes match model expectations\nif len(dls.train) > 0:\n    b = dls.one_batch()\n    print(f\"\\n✅ DataLoaders Ready!\")\n    print(f\"Input Batch Shape: {b[0].shape} (BS, Channels, H, W)\")\n    print(f\"Mask Batch Shape:  {b[1].shape} (BS, H, W)\")\nelse:\n    print(\"\\n❌ Error: DataLoaders are empty. Check file paths.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T07:46:28.661572Z","iopub.execute_input":"2025-12-07T07:46:28.662531Z","iopub.status.idle":"2025-12-07T07:46:28.852548Z","shell.execute_reply.started":"2025-12-07T07:46:28.662507Z","shell.execute_reply":"2025-12-07T07:46:28.851698Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport numpy as np\n\ndef visualize_train_sample(dataset, idx=0):\n    \"\"\"\n    Visualizes a single sample from the VesuviusDataset.\n    Shows: Input (Middle Slice), Ground Truth Mask, and Overlay.\n    \"\"\"\n    # 1. Load Data\n    tensor_img, mask_tensor = dataset[idx]\n    \n    # 2. Process Image (Middle Slice)\n    # tensor_img shape is [3, H, W] (from the 3-slice stack). \n    # We grab index 1 (the middle slice) for visualization.\n    img_np = tensor_img[1].cpu().numpy()\n    \n    # Normalize image to 0-1 range for better display (since it is Z-scored)\n    img_disp = (img_np - img_np.min()) / (img_np.max() - img_np.min())\n    \n    # 3. Process Mask\n    mask_np = mask_tensor.cpu().numpy()\n    \n    # 4. Create Plots\n    fig, axes = plt.subplots(1, 3, figsize=(15, 5))\n    \n    # --- Plot A: Input Image ---\n    axes[0].imshow(img_disp, cmap='gray')\n    axes[0].set_title(f\"Input Image (Sample {idx})\")\n    axes[0].axis('off')\n    \n    # --- Plot B: Ground Truth Mask ---\n    # 0 = Background, 1 = Papyrus, 2 = Ignore\n    axes[1].imshow(mask_np, cmap='viridis', interpolation='nearest')\n    axes[1].set_title(f\"Ground Truth (Unique: {np.unique(mask_np)})\")\n    axes[1].axis('off')\n    \n    # --- Plot C: Overlay ---\n    axes[2].imshow(img_disp, cmap='gray')\n    \n    # Create a masked array so we ONLY display the papyrus sheet (Value 1)\n    # We hide (mask) everything that is NOT 1\n    papyrus_overlay = np.ma.masked_where(mask_np != 1, mask_np)\n    \n    # Overlay in Red ('autumn' cmap goes nicely over gray)\n    axes[2].imshow(papyrus_overlay, cmap='autumn', alpha=0.6)\n    axes[2].set_title(\"Overlay (Papyrus in Red)\")\n    axes[2].axis('off')\n    \n    plt.tight_layout()\n    plt.show()\n\n# --- DEMO ---\n# Loop to find a sample that actually contains papyrus (value 1)\nprint(\"Looking for a sample with papyrus...\")\nfound_papyrus = False\nfor i in range(20): # Check first 20 samples\n    _, m = train_ds[i]\n    if 1 in m:\n        visualize_train_sample(train_ds, idx=i)\n        found_papyrus = True\n        break\n\nif not found_papyrus:\n    print(\"No papyrus found in the first 20 samples. Showing sample 0 anyway:\")\n    visualize_train_sample(train_ds, idx=0)","metadata":{"_uuid":"92e670af-8a31-4ed0-94f0-1fb5ef9f8409","_cell_guid":"f15f7cac-9805-4006-8579-878c61aa1c1b","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-12-07T07:46:33.115563Z","iopub.execute_input":"2025-12-07T07:46:33.115839Z","iopub.status.idle":"2025-12-07T07:46:33.466686Z","shell.execute_reply.started":"2025-12-07T07:46:33.115816Z","shell.execute_reply":"2025-12-07T07:46:33.463694Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- METRICS & LOSS ---\ndef acc_surface(inp, targ):\n    targ = targ.squeeze(1)\n    mask = targ != 2\n    if mask.sum() == 0: return 0.0\n    return (inp.argmax(dim=1)[mask] == targ[mask]).float().mean()\n\n# --- METRICS & LOSS ---\nclass DiceLoss(nn.Module):\n    def __init__(self, smooth=1, ignore_index=2):\n        super(DiceLoss, self).__init__()\n        self.smooth = smooth\n        self.ignore_index = ignore_index\n\n    def forward(self, inputs, targets):\n        # inputs: [BS, 2, H, W]\n        inputs = F.softmax(inputs, dim=1)\n        \n        # Flatten tensors\n        # We focus on channel 1 (papyrus)\n        input_flat = inputs[:, 1].contiguous().view(-1)\n        target_flat = targets.contiguous().view(-1)\n        \n        # --- IMPROVEMENT 1: Handle Ignore Index ---\n        # Create a mask where targets are NOT the ignore index (2)\n        valid_mask = target_flat != self.ignore_index\n        \n        # Filter both inputs and targets using this mask\n        input_flat = input_flat[valid_mask]\n        target_flat = target_flat[valid_mask]\n        \n        # Now calculate Dice on valid pixels only\n        intersection = (input_flat * target_flat).sum()\n        dice = (2. * intersection + self.smooth) / (input_flat.sum() + target_flat.sum() + self.smooth)\n        \n        return 1 - dice\n\nclass ComboLoss(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.ce = CrossEntropyLossFlat(axis=1, ignore_index=2)\n        self.dice = DiceLoss(ignore_index=2)\n        \n    def forward(self, inputs, targets):\n        return self.ce(inputs, targets) + self.dice(inputs, targets)\n\n# UPDATE YOUR LEARNER\nloss_func = ComboLoss()","metadata":{"_uuid":"191d96ae-0a99-4a7c-afb3-2e21510b6d8e","_cell_guid":"316a2196-5884-4bba-8737-b15e240f466b","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-12-07T07:46:53.558677Z","iopub.execute_input":"2025-12-07T07:46:53.559239Z","iopub.status.idle":"2025-12-07T07:46:53.568177Z","shell.execute_reply.started":"2025-12-07T07:46:53.559216Z","shell.execute_reply":"2025-12-07T07:46:53.567345Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- MODEL & TRAINING ---\n\n# 1. Load Pretrained Encoder\nweight_path = '/kaggle/input/convnext-weights/convnext_small.pth'\nbase_model = timm.create_model(ARCH, pretrained=False, num_classes=0, global_pool='')\n\n# Load weights safely\nprint(f\"Loading weights from {weight_path}...\")\ncheckpoint = torch.load(weight_path, map_location='cpu')\nstate_dict = checkpoint['model'] if 'model' in checkpoint else checkpoint\nclean_state_dict = {k.replace('module.', ''): v for k, v in state_dict.items()}\nbase_model.load_state_dict(clean_state_dict, strict=False)\n\n# 2. Build U-Net\nlayers = [base_model.stem]\nfor stage in base_model.stages:\n    layers.append(stage)\nbody = nn.Sequential(*layers)\n\nimg_size = dls.one_batch()[0].shape[-2:]\nmodel = models.unet.DynamicUnet(body, n_out=2, img_size=img_size)\n\nlearn = Learner(dls, model, loss_func=loss_func, metrics=acc_surface, wd=0.2)\n\n# 3. Training Loop\n# Step 1: Freeze Encoder\nprint(\"Step 1: Training Head (Frozen)...\")\nfor m in body:\n    for param in m.parameters():\n        param.requires_grad = False\nlearn.fit_one_cycle(1, lr_max=1e-3)","metadata":{"_uuid":"41a3f43e-c6cd-469b-997d-a55761222552","_cell_guid":"a0bc4cdf-a5e8-44a7-ac2c-215cc6f3f32b","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-12-07T07:46:56.923256Z","iopub.execute_input":"2025-12-07T07:46:56.923763Z","iopub.status.idle":"2025-12-07T07:48:56.504704Z","shell.execute_reply.started":"2025-12-07T07:46:56.923742Z","shell.execute_reply":"2025-12-07T07:48:56.503492Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Run LR Finder and capture the result\nres = learn.lr_find(suggest_funcs=(minimum, steep, valley, slide))\n\n# Print the numerical suggestions\nprint(f\"Minimum: {res.minimum:.2e}\")\nprint(f\"Steepest: {res.steep:.2e}\")\nprint(f\"Valley:   {res.valley:.2e}\") # <--- Usually the best choice\nprint(f\"Slide:    {res.slide:.2e}\")","metadata":{"_uuid":"99074538-22f2-4494-931f-ef32f0c6dcad","_cell_guid":"7af9bdf3-38ff-4cb1-b3e6-3c8999b831ac","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-12-07T00:24:24.007217Z","iopub.execute_input":"2025-12-07T00:24:24.008096Z","iopub.status.idle":"2025-12-07T00:26:15.381209Z","shell.execute_reply.started":"2025-12-07T00:24:24.008064Z","shell.execute_reply":"2025-12-07T00:26:15.380448Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Step 2: Unfreeze and Fine-tune\nprint(\"Step 2: Fine-tuning Full Model (Unfrozen)...\")\n\nfor m in body:\n    for param in m.parameters():\n        param.requires_grad = True\n\nlr = res.valley\nlearn.fit_one_cycle(7, lr_max=slice(lr/10, lr))","metadata":{"_uuid":"ee36c281-d37e-4423-a54b-d8947025c652","_cell_guid":"3797bb46-4b74-4fce-8a57-ec8172afdf30","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-12-07T00:26:40.960183Z","iopub.execute_input":"2025-12-07T00:26:40.960914Z","iopub.status.idle":"2025-12-07T00:34:20.088797Z","shell.execute_reply.started":"2025-12-07T00:26:40.960890Z","shell.execute_reply":"2025-12-07T00:34:20.087631Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- INFERENCE & SUBMISSION ---\nfrom skimage.morphology import remove_small_objects # Standard in Kaggle\n\n# Helper: Predict single slice\ndef predict_and_restore(tensor_chunk, original_shape):\n    model.eval() \n    with torch.no_grad():\n        preds = model(tensor_chunk.to(dls.device))\n        preds_resized = F.interpolate(\n            preds, \n            size=original_shape, \n            mode='bilinear', \n            align_corners=False\n        )\n        return preds_resized.argmax(dim=1).cpu().numpy()[0]\n\n# --- INFERENCE LOOP ---\nprint(\"Starting inference...\")\ntest_df = pd.read_csv(PATH/'test.csv')\nos.makedirs(\"/kaggle/working/pred_masks\", exist_ok=True)\n\nfor _, row in test_df.iterrows():\n    vid = row['id']\n    print(f\"Processing Volume {vid}...\")\n    \n    # We use the separate Test class we defined earlier\n    test_ds = VesuviusTestDataset(pd.DataFrame([row]))\n    \n    path_to_img = PATH/'test_images'/f'{vid}.tif'\n    with Image.open(path_to_img) as img:\n        W, H = img.size\n    \n    predictions = []\n    \n    # 1. Predict all slices (Raw)\n    for i in range(len(test_ds)):\n        t_img = test_ds[i].unsqueeze(0)\n        raw_mask = predict_and_restore(t_img, (H, W))\n        predictions.append(raw_mask.astype(np.uint8))\n        \n    # 2. Stack into Full 3D Volume\n    full_volume = np.stack(predictions)\n    \n    # --- IMPROVEMENT 3: 3D Post-Processing ---\n    # Instead of cleaning slice-by-slice, we clean the 3D volume.\n    # This prevents cutting vertical strands of papyrus.\n    print(f\"Applying 3D cleaning to volume shape {full_volume.shape}...\")\n    \n    # Convert to boolean for morphology\n    bool_vol = full_volume.astype(bool)\n    \n    # Remove objects smaller than ~1000 voxels (adjust based on resolution)\n    # This removes dust but keeps papyrus strands connected\n    clean_vol = remove_small_objects(bool_vol, min_size=1000)\n    \n    full_mask = clean_vol.astype(np.uint8)\n    \n    # 3. Save\n    out_path = f\"/kaggle/working/pred_masks/{vid}.tif\"\n    frames = [Image.fromarray(full_mask[i]) for i in range(len(full_mask))]\n    frames[0].save(out_path, save_all=True, append_images=frames[1:], compression=\"tiff_deflate\")\n    print(f\"Saved {vid}.tif\")\n\n# Create Submission ZIP\nprint(\"Zipping...\")\nwith zipfile.ZipFile(\"/kaggle/working/submission.zip\", 'w') as z:\n    for _, row in test_df.iterrows():\n        vid = row['id']\n        z.write(f\"/kaggle/working/pred_masks/{vid}.tif\", f\"{vid}.tif\")\nprint(\"Submission Ready!\")","metadata":{"_uuid":"1f84f51d-66e6-46a8-a211-df47f2924b09","_cell_guid":"68759527-bc10-45fe-bab7-766bd96bac96","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nfrom PIL import Image\nimport numpy as np\n\ndef visualize_prediction(test_df, index=0):\n    # 1. Get ID and Paths\n    try:\n        test_id = test_df['id'].iloc[index]\n    except IndexError:\n        print(f\"Index {index} out of bounds for test_df with length {len(test_df)}\")\n        return\n\n    print(f\"Visualizing Test Volume: {test_id}\")\n    \n    # Paths\n    vol_path = PATH/'test_images'/f'{test_id}.tif'\n    pred_path = f\"/kaggle/working/pred_masks/{test_id}.tif\"\n    \n    # Check if prediction exists (it should if inference ran)\n    if not os.path.exists(pred_path):\n        print(f\"Prediction file not found: {pred_path}\")\n        return\n\n    # 2. Determine Middle Slice Index (without loading data)\n    with Image.open(vol_path) as img:\n        n_frames = img.n_frames\n        mid_idx = n_frames // 2\n        \n    # 3. Load ONLY the Middle Slice (Memory Efficient)\n    # Load Original\n    with Image.open(vol_path) as img:\n        img.seek(mid_idx)\n        orig_slice = np.array(img)\n        \n    # Load Prediction\n    with Image.open(pred_path) as img:\n        img.seek(mid_idx)\n        pred_slice = np.array(img)\n\n    # 4. Visualization\n    fig, axes = plt.subplots(1, 3, figsize=(15, 5))\n    \n    # Original Input\n    axes[0].imshow(orig_slice, cmap='gray')\n    axes[0].set_title(f'Original Volume (Slice {mid_idx})')\n    axes[0].axis('off')\n\n    # Prediction\n    axes[1].imshow(pred_slice, cmap='Reds')\n    axes[1].set_title('Predicted Papyrus Mask')\n    axes[1].axis('off')\n\n    # Overlay\n    axes[2].imshow(orig_slice, cmap='gray')\n    axes[2].imshow(pred_slice, cmap='Reds', alpha=0.5) # Alpha blends them\n    axes[2].set_title('Overlay')\n    axes[2].axis('off')\n\n    plt.tight_layout()\n    plt.show()\n    \n    # Stats\n    pixels = pred_slice.flatten()\n    print(f\"Slice {mid_idx} Prediction Stats:\")\n    print(f\"Background Pixels (0): {np.sum(pixels == 0)}\")\n    print(f\"Papyrus Pixels (1): {np.sum(pixels == 1)}\")\n\n# Run the visualization\n# Check if test_df exists from previous cell\nif 'test_df' in locals() and len(test_df) > 0:\n    visualize_prediction(test_df)\nelse:\n    print(\"Test dataframe is empty or undefined. (This is normal in 'Edit' mode if the test set is hidden).\")","metadata":{"_uuid":"9fd62b83-5fe6-4b14-bc8a-b45531359434","_cell_guid":"2f19ed8b-3b43-4e91-b090-024ff59a6852","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null}]}