{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":117682,"databundleVersionId":14443416,"sourceType":"competition"},{"sourceId":14044945,"sourceType":"datasetVersion","datasetId":8941352},{"sourceId":285286358,"sourceType":"kernelVersion"}],"dockerImageVersionId":31193,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"#IMPORTS\nimport sys\nimport subprocess\nimport os\nimport shutil\nimport re\nimport importlib\nimport numpy as np\nimport pandas as pd\nimport zipfile\nimport tifffile\nimport torch\nimport torch.nn as nn\nimport scipy.ndimage as ndi","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T20:03:11.576970Z","iopub.execute_input":"2025-12-11T20:03:11.577625Z","iopub.status.idle":"2025-12-11T20:03:15.578191Z","shell.execute_reply.started":"2025-12-11T20:03:11.577601Z","shell.execute_reply":"2025-12-11T20:03:15.577605Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#  CONFIGURATION AND PATHS \n\nDATASET_PATH = \"/kaggle/input/vesuvius-segresmamba-libs-2025\" \nWEIGHTS_PATH = \"/kaggle/input/segresmamba-training/checkpoint_epoch_120.pth\" \n\n# Environment Paths\nWRITABLE_PATH = \"/kaggle/working/fixed_model\"\nroot_dir = \"/kaggle/input/vesuvius-challenge-surface-detection\" # Standard competition data path\ncustom_test_path = [\"/kaggle/input/vesuvius-challenge-surface-detection/train_images/1004283650.tif\",\n                  \"/kaggle/input/vesuvius-challenge-surface-detection/train_images/1006462223.tif\"]\ncustom_output_path = \"/kaggle/working/custom_masks\"\ncustom_zip_dir = \"/kaggle/working/custom.zip\"\nos.makedirs(custom_output_path, exist_ok=True)\n\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nprint(f\"Using device: {DEVICE}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T20:03:15.579149Z","iopub.execute_input":"2025-12-11T20:03:15.579452Z","iopub.status.idle":"2025-12-11T20:03:15.608943Z","shell.execute_reply.started":"2025-12-11T20:03:15.579433Z","shell.execute_reply":"2025-12-11T20:03:15.607798Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#  ENVIRONMENT SETUP AND CODE PATCHING \n# This block ensures Mamba libraries are installed and model code is compatible.\nprint(\"\\n--- Starting Environment Setup (Installing & Patching) ---\")\n\n# --- PART 1: INSTALL LIBRARIES ---\nwheels_path = DATASET_PATH \nfor root, dirs, files in os.walk(DATASET_PATH):\n    if any(f.startswith(\"mamba_ssm\") and f.endswith(\".whl\") for f in files):\n        wheels_path = root\n        break\n\ntry:\n    import mamba_ssm\n    import monai\n    print(\"Mamba and MONAI seem already installed.\")\nexcept ImportError:\n    print(f\"...Installing libraries from {wheels_path}...\")\n    try:\n        subprocess.check_call([\n            sys.executable, \"-m\", \"pip\", \"install\", \n            \"--no-index\", \n            \"--no-deps\", \n            f\"--find-links={wheels_path}\", \n            \"causal_conv1d\", \"mamba_ssm\", \"monai\"\n        ])\n        print(\" Libraries installed successfully.\")\n    except Exception as e:\n        print(f\"Installation failed: {e}\")\n\n# --- PART 2: FIND & COPY MODEL CODE ---\nprint(\"...Locating Model Code...\")\nsource_code_path = None\nfor root, dirs, files in os.walk(DATASET_PATH):\n    # Find the directory containing the model source code\n    if \"segresmamba.py\" in files:\n        source_code_path = os.path.dirname(root) \n        break\n    # Add common alternative path in case of nested structure\n    for dir_name in dirs:\n        if dir_name == \"model\":\n             if \"segresmamba.py\" in os.listdir(os.path.join(root, dir_name)):\n                source_code_path = os.path.dirname(os.path.join(root, dir_name))\n                break\n    if source_code_path:\n        break\n\n\nif not source_code_path:\n    print(\"CRITICAL: Could not find segresmamba.py\")\nelse:\n    # Copy code\n    if os.path.exists(WRITABLE_PATH):\n        shutil.rmtree(WRITABLE_PATH)\n    try:\n        shutil.copytree(source_code_path, WRITABLE_PATH)\n    except Exception as e:\n        # Fallback copy method\n        os.makedirs(WRITABLE_PATH, exist_ok=True)\n        subprocess.call([\"cp\", \"-r\", f\"{source_code_path}/.\", WRITABLE_PATH])\n    print(f\"Copied code to: {WRITABLE_PATH}\")\n\n    # --- PART 3: APPLY THE FIXES (Patch 'bimamba_type', 'nslices', 'use_fast_path') ---\n    print(\"...Applying patches to SegResMamba source...\")\n    patched_count = 0\n    for root, dirs, files in os.walk(WRITABLE_PATH):\n        for filename in files:\n            if filename.endswith(\".py\"):\n                file_path = os.path.join(root, filename)\n                with open(file_path, 'r') as f:\n                    content = f.read()\n                \n                original_content = content\n                \n                # FIX 1: Remove 'bimamba_type'\n                content = re.sub(r',?\\s*bimamba_type\\s*=\\s*[^,)]+', '', content)\n                # FIX 2: Remove 'nslices'\n                content = re.sub(r',?\\s*nslices\\s*=\\s*[^,)]+', '', content)\n                # FIX 3: Remove 'use_fast_path'\n                content = re.sub(r',?\\s*use_fast_path\\s*=\\s*[^,)]+', '', content)\n\n                if content != original_content:\n                    with open(file_path, 'w') as f:\n                        f.write(content)\n                    patched_count += 1\n    print(f\"Patched {patched_count} files.\")\n\n    # --- PART 4: IMPORT MODEL CLASS ---\n    if WRITABLE_PATH not in sys.path:\n        sys.path.insert(0, WRITABLE_PATH)\n\n# Import SegResMamba ONLY after environment setup is complete\nfrom model.segresmamba import SegResMamba\n\nprint(\"--- Setup Complete ---\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T20:03:15.609698Z","iopub.status.idle":"2025-12-11T20:03:15.609999Z","shell.execute_reply.started":"2025-12-11T20:03:15.609851Z","shell.execute_reply":"2025-12-11T20:03:15.609863Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from monai.inferers import sliding_window_inference ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T20:03:15.610964Z","iopub.status.idle":"2025-12-11T20:03:15.611418Z","shell.execute_reply.started":"2025-12-11T20:03:15.611269Z","shell.execute_reply":"2025-12-11T20:03:15.611283Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# DATA UTILITIES \n\ndef load_volume(path):\n    \"\"\"Loads the 3D TIF file and converts to float32 NumPy array (D, H, W).\"\"\"\n    # Teammate's loading logic\n    vol = tifffile.imread(path)          \n    return vol.astype(np.float32)\n\ndef val_transformation(image_np):\n    \"\"\"\n    Normalizes image (0-255 -> 0.0-1.0) and formats for PyTorch: (B, C, D, H, W).\n    \"\"\"\n    # Normalization (matches logic of ScaleIntensityRange)\n    normalized_volume = image_np / 255.0\n    \n    # Add Batch (0) and Channel (1) dimensions for PyTorch\n    # NumPy (D, H, W) -> Torch (1, 1, D, H, W)\n    volume_tensor = torch.from_numpy(normalized_volume).unsqueeze(0).unsqueeze(0).float()\n    \n    return volume_tensor\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T20:03:15.612379Z","iopub.status.idle":"2025-12-11T20:03:15.612631Z","shell.execute_reply.started":"2025-12-11T20:03:15.612507Z","shell.execute_reply":"2025-12-11T20:03:15.612517Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# MODEL UTILITY \n\ndef get_model(weights_path):\n    \"\"\"Initializes SegResMamba model and loads weights and handles the 'module.' prefix removal.\"\"\"\n    try:\n        model = SegResMamba(\n            spatial_dims=3, \n            in_chans=1,     # Grayscale Input\n            out_chans=1,    # Binary Output (Logit for Foreground)\n        )\n    except Exception as e:\n        print(f\"MODEL INSTANTIATION FAILED: {e}\")\n        raise Exception(\"Model instantiation failed. Check Mamba installation.\")\n    \n    if os.path.exists(weights_path):\n        print(f\"Loading weights from: {weights_path}\")\n        # Load the state dict\n        state_dict = torch.load(weights_path, map_location=DEVICE)\n        \n        # --- FIX: REMOVING 'module.' PREFIX ---\n        new_state_dict = {}\n        for k, v in state_dict.items():\n            # If the key starts with 'module.', strip it off\n            if k.startswith('module.'):\n                name = k[7:]  # remove 'module.' prefix\n            else:\n                name = k\n            new_state_dict[name] = v\n        # -----------------------------------\n        \n        # Load the cleaned state dict\n        model.load_state_dict(new_state_dict)\n\n    else:\n        # Crash if weights file is not found\n        print(\"\\n!!!!!!!!!!!!!!! CRITICAL ERROR !!!!!!!!!!!!!!!\")\n        print(f\"ERROR: Weights file not found at: {weights_path}\")\n        print(\"Please check the full file path in your Kaggle Inputs.\")\n        print(\"!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!\\n\")\n        raise FileNotFoundError(f\"Weights file not found at: {weights_path}\")\n        \n    model.to(DEVICE)\n    model.eval() \n    return model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T20:03:15.613568Z","iopub.status.idle":"2025-12-11T20:03:15.613997Z","shell.execute_reply.started":"2025-12-11T20:03:15.613845Z","shell.execute_reply":"2025-12-11T20:03:15.613858Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# TOPOLOGY-AWARE POST-PROCESSING\n\ndef topo_postprocess(\n    probs,          # (D, H, W) foreground probabilities\n    T_low=0.40,\n    T_high=0.80,\n    z_radius=2,\n    xy_radius=0,\n    dust_min_size=100,\n):\n    \"\"\"\n    Hysteresis + 3D propagation + dust removal using SciPy.\n    \"\"\"\n    # 1) Hysteresis Thresholds\n    strong = probs >= T_high\n    weak   = probs >= T_low\n\n    if not strong.any():\n        return np.zeros_like(probs, dtype=np.uint8)\n\n    # 2) Build connectivity structure for 3D propagation\n    z_sz  = 2 * z_radius + 1\n    xy_sz = 2 * xy_radius + 1\n    structure = np.ones((z_sz, xy_sz, xy_sz), dtype=bool)\n\n    # 3) Grow strong into weak (3D Hysteresis)\n    grown = ndi.binary_propagation(strong, mask=weak, structure=structure)\n\n    if not grown.any():\n        return np.zeros_like(probs, dtype=np.uint8)\n\n    # 4) Remove 3D dust (Connected Component Labeling)\n    labels, num = ndi.label(grown, structure=np.ones((3, 3, 3), bool))\n    \n    if num == 0:\n        return np.zeros_like(probs, dtype=np.uint8)\n\n    sizes = np.bincount(labels.ravel())\n    keep = sizes >= dust_min_size\n    keep[0] = False  # Exclude background (Label 0)\n    cleaned = keep[labels]\n\n    return cleaned.astype(np.uint8)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T20:03:15.614815Z","iopub.status.idle":"2025-12-11T20:03:15.615261Z","shell.execute_reply.started":"2025-12-11T20:03:15.615122Z","shell.execute_reply":"2025-12-11T20:03:15.615135Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# PREDICT FUNCTION \n\n# Initialize the model once\nmodel = get_model(WEIGHTS_PATH)\n\ndef predict(sample_tensor):\n    \"\"\"\n    Runs MONAI Sliding Window Inference and applies post-processing.\n    sample_tensor: (1, 1, D, H, W) torch.Tensor on GPU.\n    \"\"\"\n    # SWI parameters matching teammate's Keras config\n    ROI_SIZE = (128, 128, 128)\n    SW_BATCH_SIZE = 4 \n    OVERLAP = 0.5\n    \n    with torch.no_grad():\n        # Using mixed precision for faster inference\n        with torch.amp.autocast(DEVICE): \n            # MONAI Sliding Window Inference (handles patching and stitching)\n            val_outputs = sliding_window_inference(\n                inputs=sample_tensor, \n                roi_size=ROI_SIZE, \n                sw_batch_size=SW_BATCH_SIZE, \n                predictor=model, \n                overlap=OVERLAP\n            )\n    \n    # Post-processing steps:\n    # 1. Sigmoid to get probabilities (0.0 to 1.0)\n    # 2. Squeeze to remove Batch (0) and Channel (0) dimensions\n    # 3. .cpu().numpy() to convert to NumPy for SciPy topo_postprocess\n    probs_fg_np = torch.sigmoid(val_outputs).squeeze().cpu().numpy()\n\n    # Apply topology-aware postprocessing\n    #final_mask = topo_postprocess(probs_fg_np) \n\n    return probs_fg_np # (D, H, W) uint8 {0,1}\n\n# --- 7. MAIN EXECUTION LOOP (Submission Generation) ---\n\n# 1. Load test fragments list\n\n\n# 2. Run inference and create submission ZIP\nwith zipfile.ZipFile(\n    custom_zip_dir, \"w\", compression=zipfile.ZIP_DEFLATED\n) as z:\n    print(f\"\\nStarting inference on {len(custom_test_path)} test fragments...\")\n    \n    for tif_path in custom_test_path:\n        image_id = os.path.basename(tif_path).replace(\".tif\", \"\")\n        \n        # Load and preprocess\n        volume_np = load_volume(tif_path) \n        volume_tensor = val_transformation(volume_np).to(DEVICE)\n        \n        # Predict\n        output_probs = predict(volume_tensor) # (D, H, W) uint8\n\n        # Save and zip mask\n        out_path = f\"{custom_output_path}/{image_id}.npy\" \n        np.save(out_path, output_probs)\n\n        # Add the .npy file to the zip and clean up\n        z.write(out_path, arcname=f\"{image_id}.npy\")\n        os.remove(out_path)\n        \n        print(f\"  Processed {image_id}\", end='\\r')\n\nprint(\"\\n\\n Inference Complete.\")\nprint(f\"Submission ZIP ready: {custom_zip_dir}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T20:03:15.616060Z","iopub.status.idle":"2025-12-11T20:03:15.616324Z","shell.execute_reply.started":"2025-12-11T20:03:15.616195Z","shell.execute_reply":"2025-12-11T20:03:15.616207Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}