{"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":"none","dataSources":[{"sourceId":117682,"databundleVersionId":14443416,"sourceType":"competition"},{"sourceId":10038,"sourceType":"datasetVersion","datasetId":6978},{"sourceId":11855918,"sourceType":"datasetVersion","datasetId":7449695},{"sourceId":11885533,"sourceType":"datasetVersion","datasetId":7470225},{"sourceId":11939723,"sourceType":"datasetVersion","datasetId":7506362}],"dockerImageVersionId":31193,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"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\n\n# --- CONFIGURATION ---\nPATH = Path('/kaggle/input/vesuvius-challenge-surface-detection')\nBS = 16            # Batch Size (Lower if GPU runs out of memory)\nIMG_SIZE = 320     # Input size (Resizing slices to this for the model)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-11-29T17:10:07.486512Z","iopub.execute_input":"2025-11-29T17:10:07.486796Z","iopub.status.idle":"2025-11-29T17:10:29.114090Z","shell.execute_reply.started":"2025-11-29T17:10:07.486768Z","shell.execute_reply":"2025-11-29T17:10:29.112986Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 1. ROBUST DATASET (Lazy Loading + 2.5D Stack) ---\nclass VesuviusDataset(torch.utils.data.Dataset):\n    def __init__(self, df, mode='train'):\n        self.df = df\n        self.mode = mode\n        self.samples = []\n        \n        # Pre-calculate valid indices\n        print(f\"Preparing {mode} dataset...\")\n        for _, row in self.df.iterrows():\n            img_id = row['id']\n            # Determine folder\n            folder = 'train_images' if mode != 'test' else 'test_images'\n            path_to_img = PATH/folder/f'{img_id}.tif'\n            \n            # Get volume depth without loading the whole file\n            with Image.open(path_to_img) as img:\n                n_frames = img.n_frames\n            \n            # We skip the very first and very last slice to allow for neighbors\n            # (start from 1, end at n-1)\n            for i in range(1, n_frames - 1):\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        folder = 'train_images' if self.mode != 'test' else 'test_images'\n        \n        # --- A. LOAD 3 SLICES (2.5D) ---\n        img_path = PATH/folder/f'{img_id}.tif'\n        slices = []\n        \n        with Image.open(img_path) as img:\n            # We load: [Previous, Current, Next]\n            for offset in [-1, 0, 1]:\n                img.seek(slice_idx + offset)\n                arr = np.array(img)\n                slices.append(arr)\n        \n        # Stack into [3, Height, Width]\n        stack = np.stack(slices, axis=0)\n        \n        # Normalize: Check if 16-bit or 8-bit\n        if stack.dtype == np.uint16:\n            stack = stack.astype(np.float32) / 65535.0\n        else:\n            stack = stack.astype(np.float32) / 255.0\n            \n        tensor_img = torch.tensor(stack) # Shape: [3, H, W]\n\n        # Resize Image (Bilinear for continuous values)\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        # --- B. HANDLE TEST MODE ---\n        if self.mode == 'test':\n            return tensor_img\n\n        # --- C. LOAD MASK (TRAINING ONLY) ---\n        mask_path = PATH/'train_labels'/f'{img_id}.tif'\n        with Image.open(mask_path) as m:\n            m.seek(slice_idx)\n            mask_arr = np.array(m) # Values: 0, 1, 2\n            \n        mask_tensor = torch.tensor(mask_arr).long() # Keep as Integers!\n\n        # Resize Mask (Nearest Neighbor to preserve 0,1,2 classes)\n        mask_tensor = F.interpolate(\n            mask_tensor.unsqueeze(0).unsqueeze(0).float(), \n            size=(IMG_SIZE, IMG_SIZE), \n            mode='nearest'\n        ).long().squeeze()\n        \n        return tensor_img, mask_tensor","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-29T17:12:17.926891Z","iopub.execute_input":"2025-11-29T17:12:17.927861Z","iopub.status.idle":"2025-11-29T17:12:17.940405Z","shell.execute_reply.started":"2025-11-29T17:12:17.927829Z","shell.execute_reply":"2025-11-29T17:12:17.939264Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 2. SETUP DATALOADERS ---\ntrain_df = pd.read_csv(PATH/'train.csv')[:3]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-29T17:12:21.818661Z","iopub.execute_input":"2025-11-29T17:12:21.819069Z","iopub.status.idle":"2025-11-29T17:12:21.840322Z","shell.execute_reply.started":"2025-11-29T17:12:21.819037Z","shell.execute_reply":"2025-11-29T17:12:21.839219Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Validation Split: Hold out one fragment (e.g., the last one) to prevent overfitting\nvalid_id = train_df['id'].unique()[-1]\nt_df = train_df[train_df['id'] != valid_id]\nv_df = train_df[train_df['id'] == valid_id]\n\ntrain_ds = VesuviusDataset(t_df, mode='train')\nvalid_ds = VesuviusDataset(v_df, mode='valid')\n\ndls = DataLoaders.from_dsets(train_ds, valid_ds, bs=BS, num_workers=2).cuda()\nprint(f\"Data ready. Train: {len(train_ds)}, Valid: {len(valid_ds)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-29T17:12:24.416003Z","iopub.execute_input":"2025-11-29T17:12:24.416340Z","iopub.status.idle":"2025-11-29T17:12:26.747883Z","shell.execute_reply.started":"2025-11-29T17:12:24.416312Z","shell.execute_reply":"2025-11-29T17:12:26.746472Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 3. METRICS & LOSS ---\n# Accuracy that ignores the '2' (Unlabeled) class\ndef acc_surface(inp, targ):\n    targ = targ.squeeze(1)\n    mask = targ != 2  # Create filter for valid pixels\n    if mask.sum() == 0: return 0.0 # Prevent division by zero\n    return (inp.argmax(dim=1)[mask] == targ[mask]).float().mean()\n\n# Loss: Ignore index 2\nloss_func = CrossEntropyLossFlat(axis=1, ignore_index=2)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-29T17:12:30.009049Z","iopub.execute_input":"2025-11-29T17:12:30.009620Z","iopub.status.idle":"2025-11-29T17:12:30.017140Z","shell.execute_reply.started":"2025-11-29T17:12:30.009509Z","shell.execute_reply":"2025-11-29T17:12:30.015836Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from fastai.vision.all import *\nimport torch\nimport timm\n\n# --- 1. Config ---\nweight_path = '/kaggle/input/convnext-weights/convnext_small.pth'\n# Use the updated name provided in your warning\narch = 'convnext_small.fb_in22k' \n\n# Safety check for memory\nif dls.bs > 32: dls = dls.new(bs=16)\n\n# --- 2. Create the Base Model (Timm) ---\n# We create the full model first so we can load weights easily\n# num_classes=0 removes the head, but keeps the structure\nbase_model = timm.create_model(arch, pretrained=False, num_classes=0, global_pool='')\n\n# --- 3. Robust Weight Loading (Before splitting) ---\nprint(f\"Loading weights from {weight_path}...\")\ncheckpoint = torch.load(weight_path, map_location='cpu', weights_only=False)\n\nif 'model' in checkpoint:\n    state_dict = checkpoint['model']\nelif 'state_dict' in checkpoint:\n    state_dict = checkpoint['state_dict']\nelse:\n    state_dict = checkpoint\n\n# Clean keys\nclean_state_dict = {}\nfor k, v in state_dict.items():\n    name = k.replace('module.', '')\n    clean_state_dict[name] = v\n\n# Load weights into the base model\n# strict=False allows ignoring the classifier head keys that might be in the file\nmissing, unexpected = base_model.load_state_dict(clean_state_dict, strict=False)\n\n# Validation\ntotal_params = len(list(base_model.parameters()))\nif len(missing) > (total_params * 0.5):\n    print(f\"⚠️ WARNING: High mismatch. Top missing: {missing[:3]}\")\nelse:\n    print(f\"✅ ConvNeXt Weights loaded! (Ignored {len(missing)} keys)\")\n\n# --- 4. The Bridge: Convert to FastAI Format ---\n# FastAI needs an iterable nn.Sequential to build a U-Net.\n# ConvNeXt structure is: .stem -> .stages[0] -> .stages[1] ...\nlayers = [base_model.stem]\nfor stage in base_model.stages:\n    layers.append(stage)\n\n# Wrap it up. Now FastAI can iterate through it.\nbody = nn.Sequential(*layers)\n\n# --- 5. Create U-Net and Learner ---\nimg_size = dls.one_batch()[0].shape[-2:]\n\n# DynamicUnet will now scan 'body', see the size changes in the stages,\n# and automatically create the skip connections.\nmodel = models.unet.DynamicUnet(body, n_out=2, img_size=img_size)\n\nlearn = Learner(dls, \n                model, \n                loss_func=loss_func, \n                metrics=acc_surface)\n\n# --- 6. Train ---\nprint(\"Freezing encoder...\")\n# Freeze the body (which we manually created)\nfor m in body:\n    for param in m.parameters():\n        param.requires_grad = False\n\nprint('Start training')\nlearn.fit_one_cycle(6, lr_max=1e-3)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-29T17:12:32.684548Z","iopub.execute_input":"2025-11-29T17:12:32.684865Z","iopub.status.idle":"2025-11-29T17:14:39.207702Z","shell.execute_reply.started":"2025-11-29T17:12:32.684841Z","shell.execute_reply":"2025-11-29T17:14:39.206467Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 5. INFERENCE & SUBMISSION ---\nprint(\"Starting inference...\")\ntest_df = pd.read_csv(PATH/'test.csv')\nos.makedirs(\"/kaggle/working/pred_masks\", exist_ok=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-29T17:15:05.232745Z","iopub.execute_input":"2025-11-29T17:15:05.233901Z","iopub.status.idle":"2025-11-29T17:15:05.254859Z","shell.execute_reply.started":"2025-11-29T17:15:05.233841Z","shell.execute_reply":"2025-11-29T17:15:05.253533Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Helper to upscale prediction back to original size\ndef predict_and_restore(tensor_chunk, original_shape):\n    # It is good practice to ensure the model is in eval mode\n    model.eval() \n    \n    with torch.no_grad():\n        # CORRECTED LINE: Call 'model' directly, not 'model.model'\n        preds = model(tensor_chunk.to(dls.device))\n        \n        # Upscale back to original resolution [1, 2, H, W]\n        preds_resized = F.interpolate(\n            preds, \n            size=original_shape, \n            mode='bilinear', \n            align_corners=False\n        )\n        # Argmax to get class 0 or 1\n        return preds_resized.argmax(dim=1).cpu().numpy()[0]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-29T17:22:52.716770Z","iopub.execute_input":"2025-11-29T17:22:52.717141Z","iopub.status.idle":"2025-11-29T17:22:52.724155Z","shell.execute_reply.started":"2025-11-29T17:22:52.717116Z","shell.execute_reply":"2025-11-29T17:22:52.723006Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for _, row in test_df.iterrows():\n    vid = row['id']\n    print(f\"Processing Volume {vid}...\")\n    \n    # Init Dataset for this specific volume\n    test_ds = VesuviusDataset(pd.DataFrame([row]), mode='test')\n    \n    # We need to know the original size to resize back\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_frames = img.n_frames\n    \n    predictions = []\n    \n    # Add dummy prediction for first slice (since we skipped it in dataset)\n    predictions.append(np.zeros((H, W), dtype=np.uint8))\n    \n    # Loop through dataset predictions\n    for i in range(len(test_ds)):\n        # Load single tensor\n        t_img = test_ds[i].unsqueeze(0) # Add batch dim -> [1, 3, 224, 224]\n        \n        # Predict and Resize back to (H, W)\n        mask = predict_and_restore(t_img, (H, W))\n        predictions.append(mask.astype(np.uint8))\n        \n    # Add dummy prediction for last slice\n    predictions.append(np.zeros((H, W), dtype=np.uint8))\n    \n    # Stack and Save\n    full_mask = np.stack(predictions)\n    out_path = f\"/kaggle/working/pred_masks/{vid}.tif\"\n    \n    # Save as multipage 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\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-29T17:22:56.141564Z","iopub.execute_input":"2025-11-29T17:22:56.142002Z","iopub.status.idle":"2025-11-29T17:23:50.691477Z","shell.execute_reply.started":"2025-11-29T17:22:56.141962Z","shell.execute_reply":"2025-11-29T17:23:50.690390Z"}},"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    test_id = test_df['id'].iloc[index]\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    # 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 Ink 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\"Ink Pixels (1): {np.sum(pixels == 1)}\")\n\n# Run the visualization\nif len(test_df) > 0:\n    visualize_prediction(test_df)\nelse:\n    print(\"Test dataframe is empty (common in interactive sessions). Submit to see results.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-29T17:23:55.147757Z","iopub.execute_input":"2025-11-29T17:23:55.148118Z","iopub.status.idle":"2025-11-29T17:23:56.134412Z","shell.execute_reply.started":"2025-11-29T17:23:55.148086Z","shell.execute_reply":"2025-11-29T17:23:56.133189Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Create Submission ZIP\nprint(\"Zipping...\")\nzip_path = \"/kaggle/working/submission.zip\"\nwith zipfile.ZipFile(zip_path, '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\")\n\nprint(\"Submission Ready!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-29T17:24:06.566393Z","iopub.execute_input":"2025-11-29T17:24:06.567211Z","iopub.status.idle":"2025-11-29T17:24:06.574945Z","shell.execute_reply.started":"2025-11-29T17:24:06.567176Z","shell.execute_reply":"2025-11-29T17:24:06.573967Z"}},"outputs":[],"execution_count":null}]}