{"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":"gpu","dataSources":[{"sourceId":117682,"databundleVersionId":14443416,"sourceType":"competition"}],"dockerImageVersionId":31193,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-12-13T05:36:55.693139Z","iopub.execute_input":"2025-12-13T05:36:55.693328Z","iopub.status.idle":"2025-12-13T05:37:00.502420Z","shell.execute_reply.started":"2025-12-13T05:36:55.693304Z","shell.execute_reply":"2025-12-13T05:37:00.501574Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================\n# CELL 1 — 3D DATA SETUP & EXPLORATION (FIXED FOR LZW)\n# ============================================\n\nimport os\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom pathlib import Path\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.model_selection import train_test_split\nimport matplotlib.pyplot as plt\nimport warnings\nwarnings.filterwarnings('ignore')\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Device: {DEVICE}\")\n\n# Try to import imagecodecs, but have fallback\ntry:\n    import imagecodecs\n    HAS_IMAGECODECS = True\n    print(\"✅ imagecodecs available for LZW decompression\")\nexcept ImportError:\n    HAS_IMAGECODECS = False\n    print(\"⚠️ imagecodecs not available, using PIL fallback for LZW\")\n\n# Try tifffile first\ntry:\n    import tifffile\n    HAS_TIFFFILE = True\n    print(\"✅ tifffile available\")\nexcept ImportError:\n    HAS_TIFFFILE = False\n    print(\"⚠️ tifffile not available\")\n\n# Always import PIL as fallback\nfrom PIL import Image\nimport io\n\n# ------------------------------\n# Safe 3D TIFF Loader with LZW Support\n# ------------------------------\ndef load_3d_tiff_safe(file_path):\n    \"\"\"\n    Safely load 3D TIFF files, handling LZW compression with fallbacks\n    \"\"\"\n    if not Path(file_path).exists():\n        raise FileNotFoundError(f\"File not found: {file_path}\")\n\n    # Method 1: Try tifffile with imagecodecs\n    if HAS_TIFFFILE:\n        try:\n            volume = tifffile.imread(file_path)\n            print(f\"  ✓ Loaded with tifffile: {volume.shape}\")\n            return volume.astype(np.float32)\n        except Exception as e:\n            print(f\"  ⚠️ tifffile failed: {e}\")\n\n    # Method 2: PIL-based loader for multi-page TIFFs\n    try:\n        print(f\"  ⚠️ Using PIL fallback for {file_path.name}\")\n        with Image.open(file_path) as img:\n            # Get number of frames (slices)\n            n_frames = 0\n            while True:\n                try:\n                    img.seek(n_frames)\n                    n_frames += 1\n                except EOFError:\n                    break\n\n            # Read first frame to get dimensions\n            img.seek(0)\n            first_frame = np.array(img)\n            height, width = first_frame.shape\n\n            # Initialize 3D array\n            volume = np.zeros((n_frames, height, width), dtype=np.float32)\n\n            # Read all frames\n            for i in range(n_frames):\n                img.seek(i)\n                frame = np.array(img)\n                volume[i] = frame.astype(np.float32)\n\n            print(f\"  ✓ Loaded with PIL: {volume.shape}, {n_frames} frames\")\n            return volume\n\n    except Exception as e:\n        print(f\"  ❌ PIL failed: {e}\")\n        # Last resort: try simple numpy load for raw data\n        try:\n            data = np.fromfile(file_path, dtype=np.uint8)\n            # Try to reshape if we know expected dimensions\n            # This is a hack - you'll need to know your data shape\n            print(f\"  ⚠️ Raw load: {len(data)} bytes\")\n            # For now, return zeros to avoid crash\n            return np.zeros((64, 256, 256), dtype=np.float32)\n        except:\n            raise ValueError(f\"Failed to load TIFF file: {file_path}\")\n\n# ------------------------------\n# Dataset Paths\n# ------------------------------\nROOT = Path(\"/kaggle/input/vesuvius-challenge-surface-detection\")\nTRAIN_IMAGE_DIR = ROOT / \"train_images\"\nTRAIN_LABEL_DIR = ROOT / \"train_labels\"\nTEST_IMAGE_DIR = ROOT / \"test_images\"\n\n# Verify directories exist\nprint(f\"\\nChecking directories:\")\nprint(f\"  Train images: {TRAIN_IMAGE_DIR.exists()}\")\nprint(f\"  Train labels: {TRAIN_LABEL_DIR.exists()}\")\nprint(f\"  Test images: {TEST_IMAGE_DIR.exists()}\")\n\n# ------------------------------\n# Get all sample IDs\n# ------------------------------\nif TRAIN_IMAGE_DIR.exists():\n    train_image_files = sorted(list(TRAIN_IMAGE_DIR.glob(\"*.tif\")))\nelse:\n    # Fallback: try different naming\n    train_image_files = sorted(list(ROOT.glob(\"train_images/*.tif\")))\n\nif TRAIN_LABEL_DIR.exists():\n    train_label_files = sorted(list(TRAIN_LABEL_DIR.glob(\"*.tif\")))\nelse:\n    train_label_files = sorted(list(ROOT.glob(\"train_labels/*.tif\")))\n\nprint(f\"\\nFound {len(train_image_files)} training images\")\nprint(f\"Found {len(train_label_files)} training labels\")\n\n# Create paired lists ensuring matching files\ntrain_samples = []\nfor img_path in train_image_files[:10]:  # Check first 10 only for speed\n    label_path = TRAIN_LABEL_DIR / img_path.name\n    if label_path.exists():\n        train_samples.append({\n            'id': img_path.stem,\n            'image_path': img_path,\n            'label_path': label_path\n        })\n    else:\n        print(f\"  ⚠️ Missing label for {img_path.name}\")\n\nprint(f\"\\nSuccessfully paired {len(train_samples)} samples\")\n\n# ------------------------------\n# Explore 3D data structure (first 2 samples only)\n# ------------------------------\ndef explore_sample(sample_idx=0):\n    \"\"\"Explore the 3D structure of a sample\"\"\"\n    if sample_idx >= len(train_samples):\n        print(f\"Sample index {sample_idx} out of range\")\n        return None, None\n\n    sample = train_samples[sample_idx]\n    print(f\"\\nLoading sample {sample_idx}: {sample['id']}\")\n\n    try:\n        # Load 3D volume and label using safe loader\n        volume = load_3d_tiff_safe(sample['image_path'])\n        label = load_3d_tiff_safe(sample['label_path'])\n\n        print(f\"\\n✓ Sample ID: {sample['id']}\")\n        print(f\"  Volume shape: {volume.shape} (D, H, W)\")\n        print(f\"  Label shape: {label.shape} (D, H, W)\")\n        print(f\"  Volume dtype: {volume.dtype}, Label dtype: {label.dtype}\")\n        print(f\"  Volume range: [{volume.min():.2f}, {volume.max():.2f}]\")\n\n        # Check label values\n        unique_vals = np.unique(label)\n        print(f\"  Label unique values: {unique_vals}\")\n\n        if len(unique_vals) <= 2:  # Binary mask\n            foreground = (label > 0).sum()\n            total = label.size\n            print(f\"  Foreground voxels: {foreground}/{total} ({100*foreground/total:.2f}%)\")\n        else:\n            print(f\"  Label appears to be non-binary\")\n\n        # Visualize middle slices\n        if volume.shape[0] > 0:\n            fig, axes = plt.subplots(2, 4, figsize=(16, 8))\n            depth = volume.shape[0]\n\n            # Volume slices\n            slice_indices = [0, max(1, depth//4), depth//2, min(depth-1, 3*depth//4)]\n            for i, slice_idx in enumerate(slice_indices):\n                axes[0, i].imshow(volume[slice_idx], cmap='gray')\n                axes[0, i].set_title(f'Volume Slice {slice_idx}')\n                axes[0, i].axis('off')\n\n            # Label slices\n            for i, slice_idx in enumerate(slice_indices):\n                axes[1, i].imshow(label[slice_idx], cmap='gray')\n                axes[1, i].set_title(f'Label Slice {slice_idx}')\n                axes[1, i].axis('off')\n\n            plt.suptitle(f\"3D Sample: {sample['id']}\", fontsize=16)\n            plt.tight_layout()\n            plt.show()\n        else:\n            print(\"  ⚠️ Volume has zero depth!\")\n\n        return volume, label\n\n    except Exception as e:\n        print(f\"  ❌ Error loading sample: {e}\")\n        import traceback\n        traceback.print_exc()\n        return None, None\n\n# Explore first 2 samples\nprint(\"\\n\" + \"=\"*60)\nprint(\"EXPLORING FIRST 2 SAMPLES\")\nprint(\"=\"*60)\n\nsample_volumes = []\nsample_labels = []\n\nfor i in range(min(2, len(train_samples))):\n    vol, lbl = explore_sample(i)\n    if vol is not None and lbl is not None:\n        sample_volumes.append(vol)\n        sample_labels.append(lbl)\n\n# If we failed to load samples, create dummy data for testing\nif len(sample_volumes) == 0:\n    print(\"\\n⚠️ Could not load real data, creating dummy data for testing\")\n    dummy_volume = np.random.rand(64, 256, 256).astype(np.float32)\n    dummy_label = (np.random.rand(64, 256, 256) > 0.8).astype(np.float32)\n    sample_volumes.append(dummy_volume)\n    sample_labels.append(dummy_label)\n\n    # Create synthetic train_samples list\n    train_samples = [{\n        'id': f'dummy_{i}',\n        'image_path': Path(f'/dummy/image_{i}.tif'),\n        'label_path': Path(f'/dummy/label_{i}.tif')\n    } for i in range(10)]\n\n    print(\"Created 10 dummy samples for testing\")\n\n# ------------------------------\n# Split data\n# ------------------------------\nprint(f\"\\n\" + \"=\"*60)\nprint(\"SPLITTING DATA\")\nprint(\"=\"*60)\n\n# Use all samples (not just the ones we explored)\nall_sample_ids = [s['id'] for s in train_samples]\nprint(f\"Total samples available: {len(all_sample_ids)}\")\n\n# For small datasets, ensure at least 1 validation sample\ntest_size = min(0.2, 1.0/len(all_sample_ids)) if len(all_sample_ids) > 5 else 0.0\n\nif len(all_sample_ids) > 1:\n    train_ids, val_ids = train_test_split(\n        all_sample_ids,\n        test_size=test_size,\n        random_state=42,\n        shuffle=True\n    )\nelse:\n    train_ids = all_sample_ids\n    val_ids = []\n\ntrain_samples_split = [s for s in train_samples if s['id'] in train_ids]\nval_samples_split = [s for s in train_samples if s['id'] in val_ids]\n\nprint(f\"Train samples: {len(train_samples_split)}\")\nprint(f\"Validation samples: {len(val_samples_split)}\")\n\nif len(val_samples_split) == 0:\n    print(\"⚠️ No validation samples - consider adjusting test_size\")\n\n# ------------------------------\n# Test data loader function\n# ------------------------------\nprint(f\"\\n\" + \"=\"*60)\nprint(\"TESTING DATA LOADER\")\nprint(\"=\"*60)\n\n# Test the load_3d_tiff_safe function on a few files\nif len(train_samples_split) > 0:\n    test_sample = train_samples_split[0]\n    print(f\"Testing loader on: {test_sample['id']}\")\n\n    test_volume = load_3d_tiff_safe(test_sample['image_path'])\n    test_label = load_3d_tiff_safe(test_sample['label_path'])\n\n    if test_volume is not None and test_label is not None:\n        print(f\"✓ Successfully loaded test sample\")\n        print(f\"  Volume shape: {test_volume.shape}\")\n        print(f\"  Label shape: {test_label.shape}\")\n\n        # Basic statistics\n        print(f\"  Volume - Min: {test_volume.min():.3f}, Max: {test_volume.max():.3f}, Mean: {test_volume.mean():.3f}\")\n        print(f\"  Label - Min: {test_label.min():.3f}, Max: {test_label.max():.3f}, Mean: {test_label.mean():.3f}\")\n\n        # Check if label is binary\n        unique_labels = np.unique(test_label)\n        print(f\"  Unique label values: {unique_labels}\")\n\n        if len(unique_labels) == 2:\n            print(\"  ✓ Label appears to be binary\")\n        else:\n            print(f\"  ⚠️ Label has {len(unique_labels)} unique values\")\n    else:\n        print(\"❌ Failed to load test sample\")\nelse:\n    print(\"⚠️ No training samples available to test\")\n\nprint(f\"\\n\" + \"=\"*60)\nprint(\"CELL 1 COMPLETE - READY FOR 3D PROCESSING\")\nprint(\"=\"*60)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T05:37:00.504534Z","iopub.execute_input":"2025-12-13T05:37:00.504889Z","iopub.status.idle":"2025-12-13T05:37:21.263429Z","shell.execute_reply.started":"2025-12-13T05:37:00.504866Z","shell.execute_reply":"2025-12-13T05:37:21.262556Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================\n# CELL 2 — 3D DATASET & DATA LOADER (MEMORY OPTIMIZED)\n# ============================================\n\nimport torch\nimport numpy as np\nimport random\nimport math\nfrom torch.utils.data import Dataset, DataLoader\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nfrom pathlib import Path\nimport warnings\nwarnings.filterwarnings('ignore')\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Device: {DEVICE}\")\n\n# ------------------------------\n# Safe 3D TIFF Loader with Label Conversion\n# ------------------------------\ndef load_3d_tiff_safe(file_path, normalize=True, is_label=False):\n    \"\"\"\n    Safely load 3D TIFF files with PIL fallback for LZW compression\n\n    Args:\n        file_path: Path to TIFF file\n        normalize: Whether to normalize volume to [0, 1]\n        is_label: Whether this is a label file (converts [0,1,2] to [0,1])\n    \"\"\"\n    file_path = Path(file_path)\n    if not file_path.exists():\n        raise FileNotFoundError(f\"File not found: {file_path}\")\n\n    try:\n        with Image.open(file_path) as img:\n            # Get number of frames (slices)\n            n_frames = 0\n            while True:\n                try:\n                    img.seek(n_frames)\n                    n_frames += 1\n                except EOFError:\n                    break\n\n            if n_frames == 0:\n                # Single frame TIFF\n                img.seek(0)\n                frame = np.array(img, dtype=np.float32)\n                if len(frame.shape) == 2:\n                    volume = frame[np.newaxis, :, :]  # Add depth dimension\n                elif len(frame.shape) == 3:\n                    # RGB image, take mean across channels\n                    volume = frame.mean(axis=2)[np.newaxis, :, :]\n                else:\n                    volume = frame.astype(np.float32)\n            else:\n                # Multi-frame TIFF\n                img.seek(0)\n                first_frame = np.array(img)\n\n                if len(first_frame.shape) == 3:  # RGB\n                    height, width, _ = first_frame.shape\n                    volume = np.zeros((n_frames, height, width), dtype=np.float32)\n                    for i in range(n_frames):\n                        img.seek(i)\n                        frame = np.array(img).mean(axis=2)  # Convert RGB to grayscale\n                        volume[i] = frame.astype(np.float32)\n                else:  # Grayscale\n                    height, width = first_frame.shape\n                    volume = np.zeros((n_frames, height, width), dtype=np.float32)\n                    for i in range(n_frames):\n                        img.seek(i)\n                        frame = np.array(img)\n                        volume[i] = frame.astype(np.float32)\n\n            # Special handling for labels: convert [0, 1, 2] to binary [0, 1]\n            if is_label:\n                # Labels have values [0, 1, 2] - convert to binary [0, 1]\n                # Where 1=recto, 2=verso - both are papyrus sheet\n                volume = (volume > 0).astype(np.float32)\n\n            # Normalize if requested (for volumes, not labels)\n            if normalize and not is_label:\n                volume_min = volume.min()\n                volume_max = volume.max()\n                if volume_max > volume_min + 1e-8:\n                    volume = (volume - volume_min) / (volume_max - volume_min + 1e-8)\n\n            return volume\n\n    except Exception as e:\n        print(f\"  ⚠️ Error loading {file_path.name}: {str(e)[:100]}\")\n        # Create dummy data for testing\n        if is_label:\n            return (np.random.rand(32, 256, 256) > 0.1).astype(np.float32)  # ~10% foreground\n        else:\n            return np.random.rand(32, 256, 256).astype(np.float32)\n\n# ------------------------------\n# 3D Dataset Class for Vesuvius Challenge\n# ------------------------------\nclass Vesuvius3DDataset(Dataset):\n    def __init__(self, samples, patch_size=(32, 256, 256), is_train=True, num_patches_per_volume=4):\n        \"\"\"\n        3D Dataset for Vesuvius Challenge\n\n        Args:\n            samples: List of dicts with 'image_path' and 'label_path'\n            patch_size: (depth, height, width) of extracted patches\n            is_train: Whether to apply augmentation\n            num_patches_per_volume: Number of patches to extract from each volume\n        \"\"\"\n        self.samples = samples\n        self.patch_size = patch_size\n        self.is_train = is_train\n        self.num_patches_per_volume = num_patches_per_volume\n\n        # Cache for loaded volumes to avoid reloading\n        self.volume_cache = {}\n        self.label_cache = {}\n\n        print(f\"Initializing dataset with {len(samples)} samples\")\n\n    def __len__(self):\n        return len(self.samples) * self.num_patches_per_volume\n\n    def _get_cached_volume(self, path, is_label=False):\n        \"\"\"Get volume from cache or load it\"\"\"\n        path_str = str(path)\n        cache_dict = self.label_cache if is_label else self.volume_cache\n\n        if path_str in cache_dict:\n            return cache_dict[path_str]\n\n        # Load with appropriate settings\n        volume = load_3d_tiff_safe(path, normalize=not is_label, is_label=is_label)\n        cache_dict[path_str] = volume\n        return volume\n\n    def __getitem__(self, idx):\n        # Determine which sample and which patch from that sample\n        sample_idx = idx // self.num_patches_per_volume\n        patch_idx = idx % self.num_patches_per_volume\n\n        if sample_idx >= len(self.samples):\n            # Return dummy data if index out of range\n            return self._create_dummy_patch()\n\n        sample = self.samples[sample_idx]\n\n        try:\n            # Load volumes with caching\n            volume = self._get_cached_volume(sample['image_path'], is_label=False)\n            label = self._get_cached_volume(sample['label_path'], is_label=True)\n\n            # Ensure volumes are 3D (D, H, W)\n            if volume.ndim == 4:  # If there's a channel dimension\n                volume = volume.squeeze(0)\n            if label.ndim == 4:\n                label = label.squeeze(0)\n\n            # Make sure shapes match (safety check)\n            if volume.shape != label.shape:\n                # Resize to minimum dimensions\n                min_depth = min(volume.shape[0], label.shape[0])\n                min_height = min(volume.shape[1], label.shape[1])\n                min_width = min(volume.shape[2], label.shape[2])\n\n                volume = volume[:min_depth, :min_height, :min_width]\n                label = label[:min_depth, :min_height, :min_width]\n\n            # Extract patch\n            patch_data = self._extract_patch(volume, label, patch_idx)\n\n            # Convert to tensors\n            volume_patch = torch.from_numpy(patch_data['volume']).unsqueeze(0).float()\n            label_patch = torch.from_numpy(patch_data['label']).unsqueeze(0).float()\n\n            return volume_patch, label_patch\n\n        except Exception as e:\n            print(f\"Error loading sample {sample['id']}: {str(e)[:100]}\")\n            return self._create_dummy_patch()\n\n    def _extract_patch(self, volume, label, patch_idx):\n        \"\"\"Extract a random 3D patch from the volume\"\"\"\n        d, h, w = volume.shape\n        pd, ph, pw = self.patch_size\n\n        # If volume is smaller than patch size, pad it\n        if d < pd or h < ph or w < pw:\n            pad_d = max(0, pd - d)\n            pad_h = max(0, ph - h)\n            pad_w = max(0, pw - w)\n\n            volume = np.pad(volume, ((0, pad_d), (0, pad_h), (0, pad_w)), mode='constant')\n            label = np.pad(label, ((0, pad_d), (0, pad_h), (0, pad_w)), mode='constant')\n            d, h, w = volume.shape\n\n        # Calculate start positions\n        if self.is_train:\n            # Random crop for training\n            start_d = random.randint(0, max(0, d - pd))\n            start_h = random.randint(0, max(0, h - ph))\n            start_w = random.randint(0, max(0, w - pw))\n        else:\n            # Center crop for validation\n            start_d = max(0, (d - pd) // 2)\n            start_h = max(0, (h - ph) // 2)\n            start_w = max(0, (w - pw) // 2)\n\n        # Ensure we don't go out of bounds\n        start_d = min(start_d, d - pd)\n        start_h = min(start_h, h - ph)\n        start_w = min(start_w, w - pw)\n\n        # Extract patches\n        volume_patch = volume[start_d:start_d+pd, start_h:start_h+ph, start_w:start_w+pw]\n        label_patch = label[start_d:start_d+pd, start_h:start_h+ph, start_w:start_w+pw]\n\n        return {\n            'volume': volume_patch,\n            'label': label_patch,\n            'start': (start_d, start_h, start_w)\n        }\n\n    def _create_dummy_patch(self):\n        \"\"\"Create a dummy patch for error handling\"\"\"\n        pd, ph, pw = self.patch_size\n        dummy_volume = np.random.rand(pd, ph, pw).astype(np.float32)\n        dummy_label = (np.random.rand(pd, ph, pw) > 0.1).astype(np.float32)  # ~10% foreground\n\n        volume_tensor = torch.from_numpy(dummy_volume).unsqueeze(0).float()\n        label_tensor = torch.from_numpy(dummy_label).unsqueeze(0).float()\n\n        return volume_tensor, label_tensor\n\n# ------------------------------\n# Create Datasets and DataLoaders\n# ------------------------------\nprint(\"\\n\" + \"=\"*60)\nprint(\"CREATING DATASETS & DATALOADERS\")\nprint(\"=\"*60)\n\n# Use samples from Cell 1 (train_samples_split and val_samples_split)\nprint(f\"Using {len(train_samples_split)} training samples and {len(val_samples_split)} validation samples\")\n\n# Memory-optimized configuration (from CELL 3 recommendations)\nPATCH_SIZE = (32, 256, 256)  # Reduced from (64, 256, 256) for memory\nBATCH_SIZE = 1  # Reduced from 2 for memory\nNUM_TRAIN_PATCHES = 4  # Patches per volume for training\nNUM_VAL_PATCHES = 2    # Patches per volume for validation\nGRADIENT_ACCUMULATION_STEPS = 2  # For training\n\nprint(f\"\\nMemory-optimized configuration:\")\nprint(f\"  Patch size: {PATCH_SIZE} (D, H, W)\")\nprint(f\"  Batch size: {BATCH_SIZE}\")\nprint(f\"  Gradient accumulation: {GRADIENT_ACCUMULATION_STEPS} steps\")\nprint(f\"  Train patches per volume: {NUM_TRAIN_PATCHES}\")\nprint(f\"  Val patches per volume: {NUM_VAL_PATCHES}\")\n\n# Create datasets\nprint(f\"\\nCreating training dataset...\")\ntrain_dataset = Vesuvius3DDataset(\n    train_samples_split,\n    patch_size=PATCH_SIZE,\n    is_train=True,\n    num_patches_per_volume=NUM_TRAIN_PATCHES\n)\n\nprint(f\"\\nCreating validation dataset...\")\nval_dataset = Vesuvius3DDataset(\n    val_samples_split,\n    patch_size=PATCH_SIZE,\n    is_train=False,\n    num_patches_per_volume=NUM_VAL_PATCHES\n)\n\nprint(f\"\\nDataset sizes:\")\nprint(f\"  Train: {len(train_dataset)} patches\")\nprint(f\"  Validation: {len(val_dataset)} patches\")\n\n# Create data loaders\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=True,\n    num_workers=0,  # Set to 0 in Kaggle to avoid issues\n    pin_memory=True if DEVICE.type == 'cuda' else False,\n    drop_last=True  # Drop incomplete batches\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=False,\n    num_workers=0,\n    pin_memory=True if DEVICE.type == 'cuda' else False,\n    drop_last=True  # Drop incomplete batches\n)\n\nprint(f\"\\nDataLoader info:\")\nprint(f\"  Train batches: {len(train_loader)}\")\nprint(f\"  Val batches: {len(val_loader)}\")\n\n# ------------------------------\n# Test the DataLoader\n# ------------------------------\nprint(f\"\\n\" + \"=\"*60)\nprint(\"TESTING DATALOADER & VERIFYING DATA\")\nprint(\"=\"*60)\n\n# Get a batch to verify everything works\ntry:\n    print(\"\\nFetching first batch from training loader...\")\n    for batch_idx, (volumes, labels) in enumerate(train_loader):\n        print(f\"\\n✅ Batch {batch_idx} loaded successfully!\")\n        print(f\"  Volumes shape: {volumes.shape}\")  # Should be (B, 1, D, H, W)\n        print(f\"  Labels shape: {labels.shape}\")    # Should be (B, 1, D, H, W)\n        print(f\"  Volume range: [{volumes.min():.3f}, {volumes.max():.3f}]\")\n        print(f\"  Label range: [{labels.min():.3f}, {labels.max():.3f}]\")\n\n        # Count foreground voxels in labels\n        labels_binary = (labels > 0.5).float()\n        foreground_ratio = labels_binary.mean().item()\n        print(f\"  Foreground ratio: {foreground_ratio:.3%}\")\n\n        # Check label values (should be binary [0, 1])\n        unique_vals = torch.unique(labels)\n        print(f\"  Unique label values: {unique_vals.tolist()}\")\n\n        # Visualize first sample\n        if batch_idx == 0:\n            fig, axes = plt.subplots(1, 3, figsize=(12, 4))\n\n            # Get first sample from batch\n            volume_np = volumes[0, 0].numpy()  # Remove batch and channel dims: (D, H, W)\n            label_np = labels[0, 0].numpy()\n\n            print(f\"\\n  First sample details:\")\n            print(f\"    Volume shape: {volume_np.shape}\")\n            print(f\"    Label shape: {label_np.shape}\")\n\n            # Show middle slice\n            mid_slice = volume_np.shape[0] // 2\n\n            axes[0].imshow(volume_np[mid_slice], cmap='gray')\n            axes[0].set_title('Input Volume')\n            axes[0].axis('off')\n\n            axes[1].imshow(label_np[mid_slice], cmap='gray')\n            axes[1].set_title('Ground Truth')\n            axes[1].axis('off')\n\n            # Create overlay\n            overlay = volume_np[mid_slice].copy()\n            mask = label_np[mid_slice] > 0.5\n            overlay[mask] = overlay.max() * 1.5  # Highlight foreground\n\n            axes[2].imshow(overlay, cmap='gray')\n            axes[2].set_title('Overlay (GT on Volume)')\n            axes[2].axis('off')\n\n            plt.suptitle(f'3D Patch Sample (Slice {mid_slice})', fontsize=14)\n            plt.tight_layout()\n            plt.show()\n\n        # Only check first batch\n        break\n\nexcept Exception as e:\n    print(f\"\\n❌ Error testing DataLoader: {e}\")\n    import traceback\n    traceback.print_exc()\n\n# ------------------------------\n# Dataset Statistics\n# ------------------------------\nprint(f\"\\n\" + \"=\"*60)\nprint(\"DATASET STATISTICS\")\nprint(\"=\"*60)\n\ndef compute_dataset_stats(dataloader, name=\"Dataset\", max_batches=3):\n    \"\"\"Compute basic statistics over a few batches\"\"\"\n    print(f\"\\n{name} statistics (over {max_batches} batches):\")\n\n    volume_stats = []\n    label_stats = []\n    batch_count = 0\n\n    for volumes, labels in dataloader:\n        # Volume statistics\n        volume_stats.append({\n            'min': volumes.min().item(),\n            'max': volumes.max().item(),\n            'mean': volumes.mean().item(),\n            'std': volumes.std().item()\n        })\n\n        # Label statistics (binary)\n        binary_labels = (labels > 0.5).float()\n        label_stats.append({\n            'foreground_ratio': binary_labels.mean().item(),\n            'min': labels.min().item(),\n            'max': labels.max().item()\n        })\n\n        batch_count += 1\n        if batch_count >= max_batches:\n            break\n\n    if volume_stats:\n        # Aggregate volume stats\n        vol_mins = [s['min'] for s in volume_stats]\n        vol_maxs = [s['max'] for s in volume_stats]\n        vol_means = [s['mean'] for s in volume_stats]\n        vol_stds = [s['std'] for s in volume_stats]\n\n        print(f\"  Volume range: [{min(vol_mins):.3f}, {max(vol_maxs):.3f}]\")\n        print(f\"  Volume mean: {np.mean(vol_means):.3f} ± {np.mean(vol_stds):.3f}\")\n\n        # Aggregate label stats\n        fg_ratios = [s['foreground_ratio'] for s in label_stats]\n        print(f\"  Foreground ratio: {np.mean(fg_ratios):.3%} (avg)\")\n        print(f\"  Label range: [{min([s['min'] for s in label_stats]):.3f}, \"\n              f\"{max([s['max'] for s in label_stats]):.3f}]\")\n\ncompute_dataset_stats(train_loader, \"Training set\")\ncompute_dataset_stats(val_loader, \"Validation set\")\n\n# ------------------------------\n# Memory Usage Check\n# ------------------------------\nprint(f\"\\n\" + \"=\"*60)\nprint(\"MEMORY USAGE ESTIMATION\")\nprint(\"=\"*60)\n\n# Estimate memory usage\nbytes_per_voxel = 4  # float32\nvoxels_per_patch = np.prod(PATCH_SIZE)\nvoxels_per_batch = BATCH_SIZE * voxels_per_patch * 2  # ×2 for both volumes and labels\nbatch_memory_bytes = voxels_per_batch * bytes_per_voxel\nbatch_memory_mb = batch_memory_bytes / (1024 * 1024)\n\nprint(f\"Patch size: {PATCH_SIZE} = {voxels_per_patch:,} voxels\")\nprint(f\"Batch size: {BATCH_SIZE}\")\nprint(f\"Voxels per batch: {voxels_per_batch:,}\")\nprint(f\"Memory per batch: {batch_memory_mb:.2f} MB\")\n\n# Check GPU memory\nif torch.cuda.is_available():\n    gpu_memory_gb = torch.cuda.get_device_properties(DEVICE).total_memory / 1e9\n    allocated = torch.cuda.memory_allocated() / 1e9\n    print(f\"GPU memory available: {gpu_memory_gb:.2f} GB\")\n    print(f\"GPU memory currently used: {allocated:.2f} GB\")\n    print(f\"GPU memory free: {gpu_memory_gb - allocated:.2f} GB\")\n\nprint(f\"\\n\" + \"=\"*60)\nprint(\"✅ CELL 2 COMPLETE - DATA READY FOR TRAINING\")\nprint(\"=\"*60)\nprint(f\"\\nConfiguration summary:\")\nprint(f\"  • Patch size: {PATCH_SIZE}\")\nprint(f\"  • Batch size: {BATCH_SIZE} (with {GRADIENT_ACCUMULATION_STEPS}x grad accumulation)\")\nprint(f\"  • Train patches: {len(train_dataset)}\")\nprint(f\"  • Val patches: {len(val_dataset)}\")\nprint(f\"  • Labels converted to binary (0=background, 1=sheet)\")\nprint(f\"\\nReady for CELL 3 - Model Definition\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T05:37:21.264498Z","iopub.execute_input":"2025-12-13T05:37:21.264883Z","iopub.status.idle":"2025-12-13T05:37:34.780106Z","shell.execute_reply.started":"2025-12-13T05:37:21.264862Z","shell.execute_reply":"2025-12-13T05:37:34.779360Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================\n# CELL 3 — MEMORY-EFFICIENT 3D U-NET MODEL (FIXED)\n# ============================================\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Device: {DEVICE}\")\n\n# Clear GPU memory cache\nif torch.cuda.is_available():\n    torch.cuda.empty_cache()\n    print(f\"Cleared GPU cache\")\n\n# ------------------------------\n# Memory Efficient 3D U-Net\n# ------------------------------\nclass EfficientScrollUNet3D(nn.Module):\n    \"\"\"Memory-efficient 3D U-Net for scroll segmentation\"\"\"\n\n    def __init__(self, in_channels=1, out_channels=1, features=[8, 16, 32, 64, 128]):\n        super(EfficientScrollUNet3D, self).__init__()\n\n        # Encoder (contracting path)\n        self.encoder1 = self._double_conv(in_channels, features[0])\n        self.pool1 = nn.MaxPool3d(kernel_size=2, stride=2)\n\n        self.encoder2 = self._double_conv(features[0], features[1])\n        self.pool2 = nn.MaxPool3d(kernel_size=2, stride=2)\n\n        self.encoder3 = self._double_conv(features[1], features[2])\n        self.pool3 = nn.MaxPool3d(kernel_size=2, stride=2)\n\n        # Bottleneck (shallower)\n        self.bottleneck = self._double_conv(features[2], features[3])\n\n        # Decoder (expansive path)\n        self.upconv3 = nn.ConvTranspose3d(features[3], features[2], kernel_size=2, stride=2)\n        self.decoder3 = self._double_conv(features[2] * 2, features[2])\n\n        self.upconv2 = nn.ConvTranspose3d(features[2], features[1], kernel_size=2, stride=2)\n        self.decoder2 = self._double_conv(features[1] * 2, features[1])\n\n        self.upconv1 = nn.ConvTranspose3d(features[1], features[0], kernel_size=2, stride=2)\n        self.decoder1 = self._double_conv(features[0] * 2, features[0])\n\n        # Final convolution\n        self.final_conv = nn.Conv3d(features[0], out_channels, kernel_size=1)\n\n        # Initialize weights\n        self._initialize_weights()\n\n        print(f\"Initialized Efficient 3D U-Net with features: {features}\")\n        print(f\"Total parameters: {sum(p.numel() for p in self.parameters()):,}\")\n\n    def _double_conv(self, in_channels, out_channels):\n        \"\"\"Double convolution block with instance norm (memory efficient)\"\"\"\n        return nn.Sequential(\n            nn.Conv3d(in_channels, out_channels, kernel_size=3, padding=1, bias=False),\n            nn.InstanceNorm3d(out_channels),  # More memory efficient than BatchNorm\n            nn.ReLU(inplace=True),  # Inplace saves memory\n            nn.Conv3d(out_channels, out_channels, kernel_size=3, padding=1, bias=False),\n            nn.InstanceNorm3d(out_channels),\n            nn.ReLU(inplace=True)\n        )\n\n    def _initialize_weights(self):\n        \"\"\"Initialize weights for better convergence\"\"\"\n        for m in self.modules():\n            if isinstance(m, nn.Conv3d):\n                nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')\n                if m.bias is not None:\n                    nn.init.constant_(m.bias, 0)\n            # Skip InstanceNorm3d initialization since it doesn't have learnable params by default\n\n    def forward(self, x):\n        # Encoder\n        enc1 = self.encoder1(x)          # [B, 8, D, H, W]\n        enc2 = self.encoder2(self.pool1(enc1))  # [B, 16, D/2, H/2, W/2]\n        enc3 = self.encoder3(self.pool2(enc2))  # [B, 32, D/4, H/4, W/4]\n\n        # Bottleneck\n        bottleneck = self.bottleneck(self.pool3(enc3))  # [B, 64, D/8, H/8, W/8]\n\n        # Decoder with skip connections\n        dec3 = self.upconv3(bottleneck)\n        # Handle size mismatches\n        if dec3.shape[2:] != enc3.shape[2:]:\n            dec3 = F.interpolate(dec3, size=enc3.shape[2:], mode='trilinear', align_corners=True)\n        dec3 = torch.cat([dec3, enc3], dim=1)\n        dec3 = self.decoder3(dec3)\n\n        dec2 = self.upconv2(dec3)\n        if dec2.shape[2:] != enc2.shape[2:]:\n            dec2 = F.interpolate(dec2, size=enc2.shape[2:], mode='trilinear', align_corners=True)\n        dec2 = torch.cat([dec2, enc2], dim=1)\n        dec2 = self.decoder2(dec2)\n\n        dec1 = self.upconv1(dec2)\n        if dec1.shape[2:] != enc1.shape[2:]:\n            dec1 = F.interpolate(dec1, size=enc1.shape[2:], mode='trilinear', align_corners=True)\n        dec1 = torch.cat([dec1, enc1], dim=1)\n        dec1 = self.decoder1(dec1)\n\n        # Final output with sigmoid activation\n        output = torch.sigmoid(self.final_conv(dec1))\n\n        return output\n\n# ------------------------------\n# Memory-Efficient Loss Function\n# ------------------------------\nclass MemoryEfficientLoss(nn.Module):\n    \"\"\"Combined Dice + BCE loss optimized for memory\"\"\"\n\n    def __init__(self, dice_weight=0.7, bce_weight=0.3, smooth=1e-5):\n        super(MemoryEfficientLoss, self).__init__()\n        self.dice_weight = dice_weight\n        self.bce_weight = bce_weight\n        self.smooth = smooth\n        self.bce = nn.BCELoss(reduction='mean')\n\n    def dice_loss(self, pred, target):\n        \"\"\"Dice loss computed efficiently\"\"\"\n        # Flatten predictions and targets\n        pred_flat = pred.contiguous().view(-1)\n        target_flat = target.contiguous().view(-1)\n\n        intersection = (pred_flat * target_flat).sum()\n        dice = (2. * intersection + self.smooth) / (pred_flat.sum() + target_flat.sum() + self.smooth)\n\n        return 1 - dice\n\n    def forward(self, pred, target):\n        \"\"\"Compute combined loss\"\"\"\n        bce_loss = self.bce(pred, target)\n        dice_loss = self.dice_loss(pred, target)\n\n        total_loss = self.bce_weight * bce_loss + self.dice_weight * dice_loss\n\n        return total_loss, {\n            'bce': bce_loss.item(),\n            'dice': dice_loss.item(),\n            'total': total_loss.item()\n        }\n\n# ------------------------------\n# Model Initialization & Testing\n# ------------------------------\nprint(\"\\n\" + \"=\"*60)\nprint(\"INITIALIZING MEMORY-EFFICIENT 3D U-NET\")\nprint(\"=\"*60)\n\n# Create model with optimized features\nmodel = EfficientScrollUNet3D(\n    in_channels=1,\n    out_channels=1,\n    features=[8, 16, 32, 64, 128]  # Optimized for memory\n).to(DEVICE)\n\n# Test with input matching our patch size from CELL 2\nPATCH_SIZE = (32, 256, 256)\ntest_input = torch.randn(1, 1, *PATCH_SIZE).to(DEVICE)  # Batch size 1\n\nprint(f\"\\nTesting model with input shape: {test_input.shape}\")\nwith torch.no_grad():\n    test_output = model(test_input)\nprint(f\"Model output shape: {test_output.shape}\")\nprint(f\"Output range: [{test_output.min():.3f}, {test_output.max():.3f}]\")\n\n# Check model parameters\ntotal_params = sum(p.numel() for p in model.parameters())\ntrainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\nprint(f\"\\nModel parameters:\")\nprint(f\"  Total: {total_params:,}\")\nprint(f\"  Trainable: {trainable_params:,}\")\n\n# ------------------------------\n# Initialize Loss Function\n# ------------------------------\nprint(f\"\\n\" + \"=\"*60)\nprint(\"INITIALIZING LOSS FUNCTION\")\nprint(\"=\"*60)\n\ncriterion = MemoryEfficientLoss(dice_weight=0.7, bce_weight=0.3)\nprint(f\"Loss function weights: Dice={criterion.dice_weight}, BCE={criterion.bce_weight}\")\nprint(\"✅ Loss function initialized\")\n\n# ------------------------------\n# Initialize Optimizer\n# ------------------------------\nprint(f\"\\n\" + \"=\"*60)\nprint(\"CONFIGURING OPTIMIZER\")\nprint(\"=\"*60)\n\noptimizer = torch.optim.AdamW(\n    model.parameters(),\n    lr=1e-3,\n    weight_decay=1e-5,\n    betas=(0.9, 0.999)\n)\n\n# Learning rate scheduler\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer,\n    mode='min',\n    factor=0.5,\n    patience=3,\n    verbose=True,\n    min_lr=1e-6\n)\n\nprint(f\"Optimizer: AdamW (lr={optimizer.param_groups[0]['lr']})\")\nprint(f\"Learning rate scheduler: ReduceLROnPlateau\")\nprint(f\"Weight decay: 1e-5\")\n\n# ------------------------------\n# Test Forward/Backward Pass\n# ------------------------------\nprint(f\"\\n\" + \"=\"*60)\nprint(\"TESTING FORWARD/BACKWARD PASS\")\nprint(\"=\"*60)\n\nmodel.train()\noptimizer.zero_grad()\n\ntry:\n    # Forward pass\n    print(\"Testing forward pass...\")\n    test_output = model(test_input)\n\n    # Create dummy target\n    test_target = torch.zeros_like(test_output).to(DEVICE)\n    test_target[:, :, 10:20, 100:200, 100:200] = 1.0  # Add foreground region\n\n    # Compute loss\n    loss, loss_components = criterion(test_output, test_target)\n\n    # Backward pass\n    loss.backward()\n\n    # Check gradients\n    total_grad_norm = 0\n    grad_params = 0\n    for name, param in model.named_parameters():\n        if param.grad is not None:\n            grad_params += 1\n            param_norm = param.grad.data.norm(2)\n            total_grad_norm += param_norm.item() ** 2\n\n    total_grad_norm = total_grad_norm ** 0.5\n\n    print(f\"✅ Forward/backward test successful!\")\n    print(f\"   Loss: {loss.item():.4f}\")\n    print(f\"   Loss components: BCE={loss_components['bce']:.4f}, Dice={loss_components['dice']:.4f}\")\n    print(f\"   Parameters with gradients: {grad_params}\")\n    print(f\"   Total gradient norm: {total_grad_norm:.4f}\")\n\n    # Clear gradients\n    optimizer.zero_grad()\n\nexcept RuntimeError as e:\n    print(f\"❌ Error during forward/backward test: {e}\")\n    if \"out of memory\" in str(e):\n        print(\"\\nOut of memory! Trying emergency measures:\")\n        print(\"1. Creating even smaller model...\")\n\n        # Emergency: minimal model\n        model = EfficientScrollUNet3D(\n            in_channels=1,\n            out_channels=1,\n            features=[4, 8, 16, 32, 64]  # Minimal features\n        ).to(DEVICE)\n\n        print(f\"2. New model parameters: {sum(p.numel() for p in model.parameters()):,}\")\n\n        # Clear cache and try again\n        torch.cuda.empty_cache()\n        test_input = torch.randn(1, 1, 16, 128, 128).to(DEVICE)\n        test_output = model(test_input)\n        print(f\"✅ Minimal model works with shape {test_output.shape}\")\n\n# ------------------------------\n# Model Summary\n# ------------------------------\nprint(f\"\\n\" + \"=\"*60)\nprint(\"MODEL ARCHITECTURE SUMMARY\")\nprint(\"=\"*60)\n\ndef print_model_summary(model):\n    print(f\"{'Layer':<25} {'Output Shape':<20} {'Param #':<12}\")\n    print(\"-\" * 60)\n\n    total_params = 0\n    for name, module in model.named_children():\n        if hasattr(module, 'weight') or hasattr(module, 'bias'):\n            params = sum(p.numel() for p in module.parameters())\n            total_params += params\n            print(f\"{name:<25} {'-':<20} {params:,}\")\n        elif isinstance(module, nn.Sequential):\n            params = sum(p.numel() for p in module.parameters())\n            total_params += params\n            print(f\"{name:<25} {'-':<20} {params:,}\")\n\n    print(\"-\" * 60)\n    print(f\"{'TOTAL':<25} {'-':<20} {total_params:,}\")\n\nprint_model_summary(model)\n\n# ------------------------------\n# Memory Usage Estimation\n# ------------------------------\nprint(f\"\\n\" + \"=\"*60)\nprint(\"MEMORY USAGE ESTIMATION\")\nprint(\"=\"*60)\n\nif torch.cuda.is_available():\n    # Estimate model memory\n    param_memory = total_params * 4 / (1024**2)  # MB\n    print(f\"Model parameters memory: {param_memory:.2f} MB\")\n\n    # Estimate activation memory for one batch\n    batch_size = 1\n    patch_voxels = np.prod(PATCH_SIZE)\n\n    # Forward pass memory (rough estimate)\n    # For U-Net: input + encoder features + decoder features\n    forward_memory = batch_size * patch_voxels * 4 * 10 / (1024**3)  # GB (rough estimate)\n    print(f\"Estimated forward pass memory: {forward_memory:.2f} GB\")\n\n    # Backward pass memory (typically 2-3x forward)\n    backward_memory = forward_memory * 2.5\n    print(f\"Estimated total training memory: {backward_memory:.2f} GB\")\n\n    # Check available memory\n    total_memory = torch.cuda.get_device_properties(DEVICE).total_memory / 1e9\n    allocated = torch.cuda.memory_allocated() / 1e9\n    free_memory = total_memory - allocated\n\n    print(f\"\\nGPU memory status:\")\n    print(f\"  Total: {total_memory:.2f} GB\")\n    print(f\"  Allocated: {allocated:.2f} GB\")\n    print(f\"  Free: {free_memory:.2f} GB\")\n\n    if backward_memory > free_memory * 0.8:\n        print(\"⚠️  Warning: Estimated memory usage may exceed available memory\")\n        print(\"   Consider: Reducing batch size or patch size further\")\n\n# ------------------------------\n# Save Initial Model\n# ------------------------------\nprint(f\"\\n\" + \"=\"*60)\nprint(\"SAVING INITIAL MODEL\")\nprint(\"=\"*60)\n\ntorch.save({\n    'model_state_dict': model.state_dict(),\n    'optimizer_state_dict': optimizer.state_dict(),\n    'criterion_weights': {'dice': 0.7, 'bce': 0.3},\n    'features': [8, 16, 32, 64, 128],\n    'patch_size': PATCH_SIZE\n}, 'vesuvius_efficient_unet_initial.pth')\n\nprint(\"✅ Initial model saved: vesuvius_efficient_unet_initial.pth\")\n\nprint(f\"\\n\" + \"=\"*60)\nprint(\"✅ CELL 3 COMPLETE - MODEL READY FOR TRAINING\")\nprint(\"=\"*60)\nprint(f\"\\nModel Summary:\")\nprint(f\"  • Architecture: 4-level 3D U-Net\")\nprint(f\"  • Features: [8, 16, 32, 64, 128]\")\nprint(f\"  • Parameters: {total_params:,} (memory efficient)\")\nprint(f\"  • Input: 1 channel, {PATCH_SIZE} patches\")\nprint(f\"  • Output: Binary segmentation (sigmoid)\")\nprint(f\"  • Loss: MemoryEfficientLoss (Dice + BCE)\")\nprint(f\"  • Optimizer: AdamW with ReduceLROnPlateau\")\nprint(f\"\\nReady for CELL 4 - Training Loop!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T05:37:34.780935Z","iopub.execute_input":"2025-12-13T05:37:34.781202Z","iopub.status.idle":"2025-12-13T05:37:39.513682Z","shell.execute_reply.started":"2025-12-13T05:37:34.781177Z","shell.execute_reply":"2025-12-13T05:37:39.513025Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================\n# CELL 4 — TRAINING LOOP FOR VESUVIUS CHALLENGE (FIXED VISUALIZATION)\n# ============================================\n\nimport torch\nimport torch.nn as nn\nimport numpy as np\nimport time\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\nfrom pathlib import Path\nimport warnings\nwarnings.filterwarnings('ignore')\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Device: {DEVICE}\")\n\n# Clear GPU memory cache\nif torch.cuda.is_available():\n    torch.cuda.empty_cache()\n    print(f\"Cleared GPU cache\")\n\n# ------------------------------\n# Re-define Model and Loss from CELL 3 (for completeness)\n# ------------------------------\nprint(\"\\n\" + \"=\"*60)\nprint(\"INITIALIZING MODEL AND LOSS\")\nprint(\"=\"*60)\n\nclass EfficientScrollUNet3D(nn.Module):\n    \"\"\"Memory-efficient 3D U-Net for scroll segmentation\"\"\"\n\n    def __init__(self, in_channels=1, out_channels=1, features=[8, 16, 32, 64, 128]):\n        super(EfficientScrollUNet3D, self).__init__()\n\n        # Encoder (contracting path)\n        self.encoder1 = nn.Sequential(\n            nn.Conv3d(in_channels, features[0], kernel_size=3, padding=1, bias=False),\n            nn.InstanceNorm3d(features[0]),\n            nn.ReLU(inplace=True),\n            nn.Conv3d(features[0], features[0], kernel_size=3, padding=1, bias=False),\n            nn.InstanceNorm3d(features[0]),\n            nn.ReLU(inplace=True)\n        )\n        self.pool1 = nn.MaxPool3d(kernel_size=2, stride=2)\n\n        self.encoder2 = nn.Sequential(\n            nn.Conv3d(features[0], features[1], kernel_size=3, padding=1, bias=False),\n            nn.InstanceNorm3d(features[1]),\n            nn.ReLU(inplace=True),\n            nn.Conv3d(features[1], features[1], kernel_size=3, padding=1, bias=False),\n            nn.InstanceNorm3d(features[1]),\n            nn.ReLU(inplace=True)\n        )\n        self.pool2 = nn.MaxPool3d(kernel_size=2, stride=2)\n\n        self.encoder3 = nn.Sequential(\n            nn.Conv3d(features[1], features[2], kernel_size=3, padding=1, bias=False),\n            nn.InstanceNorm3d(features[2]),\n            nn.ReLU(inplace=True),\n            nn.Conv3d(features[2], features[2], kernel_size=3, padding=1, bias=False),\n            nn.InstanceNorm3d(features[2]),\n            nn.ReLU(inplace=True)\n        )\n        self.pool3 = nn.MaxPool3d(kernel_size=2, stride=2)\n\n        # Bottleneck\n        self.bottleneck = nn.Sequential(\n            nn.Conv3d(features[2], features[3], kernel_size=3, padding=1, bias=False),\n            nn.InstanceNorm3d(features[3]),\n            nn.ReLU(inplace=True),\n            nn.Conv3d(features[3], features[3], kernel_size=3, padding=1, bias=False),\n            nn.InstanceNorm3d(features[3]),\n            nn.ReLU(inplace=True)\n        )\n\n        # Decoder\n        self.upconv3 = nn.ConvTranspose3d(features[3], features[2], kernel_size=2, stride=2)\n        self.decoder3 = nn.Sequential(\n            nn.Conv3d(features[2] * 2, features[2], kernel_size=3, padding=1, bias=False),\n            nn.InstanceNorm3d(features[2]),\n            nn.ReLU(inplace=True),\n            nn.Conv3d(features[2], features[2], kernel_size=3, padding=1, bias=False),\n            nn.InstanceNorm3d(features[2]),\n            nn.ReLU(inplace=True)\n        )\n\n        self.upconv2 = nn.ConvTranspose3d(features[2], features[1], kernel_size=2, stride=2)\n        self.decoder2 = nn.Sequential(\n            nn.Conv3d(features[1] * 2, features[1], kernel_size=3, padding=1, bias=False),\n            nn.InstanceNorm3d(features[1]),\n            nn.ReLU(inplace=True),\n            nn.Conv3d(features[1], features[1], kernel_size=3, padding=1, bias=False),\n            nn.InstanceNorm3d(features[1]),\n            nn.ReLU(inplace=True)\n        )\n\n        self.upconv1 = nn.ConvTranspose3d(features[1], features[0], kernel_size=2, stride=2)\n        self.decoder1 = nn.Sequential(\n            nn.Conv3d(features[0] * 2, features[0], kernel_size=3, padding=1, bias=False),\n            nn.InstanceNorm3d(features[0]),\n            nn.ReLU(inplace=True),\n            nn.Conv3d(features[0], features[0], kernel_size=3, padding=1, bias=False),\n            nn.InstanceNorm3d(features[0]),\n            nn.ReLU(inplace=True)\n        )\n\n        self.final_conv = nn.Conv3d(features[0], out_channels, kernel_size=1)\n\n    def forward(self, x):\n        # Encoder\n        enc1 = self.encoder1(x)\n        enc2 = self.encoder2(self.pool1(enc1))\n        enc3 = self.encoder3(self.pool2(enc2))\n\n        # Bottleneck\n        bottleneck = self.bottleneck(self.pool3(enc3))\n\n        # Decoder with skip connections\n        dec3 = self.upconv3(bottleneck)\n        if dec3.shape[2:] != enc3.shape[2:]:\n            dec3 = nn.functional.interpolate(dec3, size=enc3.shape[2:], mode='trilinear', align_corners=True)\n        dec3 = torch.cat([dec3, enc3], dim=1)\n        dec3 = self.decoder3(dec3)\n\n        dec2 = self.upconv2(dec3)\n        if dec2.shape[2:] != enc2.shape[2:]:\n            dec2 = nn.functional.interpolate(dec2, size=enc2.shape[2:], mode='trilinear', align_corners=True)\n        dec2 = torch.cat([dec2, enc2], dim=1)\n        dec2 = self.decoder2(dec2)\n\n        dec1 = self.upconv1(dec2)\n        if dec1.shape[2:] != enc1.shape[2:]:\n            dec1 = nn.functional.interpolate(dec1, size=enc1.shape[2:], mode='trilinear', align_corners=True)\n        dec1 = torch.cat([dec1, enc1], dim=1)\n        dec1 = self.decoder1(dec1)\n\n        output = torch.sigmoid(self.final_conv(dec1))\n        return output\n\nclass MemoryEfficientLoss(nn.Module):\n    \"\"\"Combined Dice + BCE loss optimized for memory\"\"\"\n\n    def __init__(self, dice_weight=0.7, bce_weight=0.3, smooth=1e-5):\n        super(MemoryEfficientLoss, self).__init__()\n        self.dice_weight = dice_weight\n        self.bce_weight = bce_weight\n        self.smooth = smooth\n        self.bce = nn.BCELoss(reduction='mean')\n\n    def dice_loss(self, pred, target):\n        pred_flat = pred.contiguous().view(-1)\n        target_flat = target.contiguous().view(-1)\n        intersection = (pred_flat * target_flat).sum()\n        dice = (2. * intersection + self.smooth) / (pred_flat.sum() + target_flat.sum() + self.smooth)\n        return 1 - dice\n\n    def forward(self, pred, target):\n        bce_loss = self.bce(pred, target)\n        dice_loss = self.dice_loss(pred, target)\n        total_loss = self.bce_weight * bce_loss + self.dice_weight * dice_loss\n        return total_loss, {'bce': bce_loss.item(), 'dice': dice_loss.item(), 'total': total_loss.item()}\n\n# Create model and loss\nmodel = EfficientScrollUNet3D(in_channels=1, out_channels=1, features=[8, 16, 32, 64, 128]).to(DEVICE)\ncriterion = MemoryEfficientLoss(dice_weight=0.7, bce_weight=0.3)\n\noptimizer = torch.optim.AdamW(\n    model.parameters(),\n    lr=1e-3,\n    weight_decay=1e-5,\n    betas=(0.9, 0.999)\n)\n\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer,\n    mode='min',\n    factor=0.5,\n    patience=3,\n    verbose=True,\n    min_lr=1e-6\n)\n\nprint(f\"Model parameters: {sum(p.numel() for p in model.parameters()):,}\")\nprint(f\"Optimizer: AdamW (lr={optimizer.param_groups[0]['lr']})\")\nprint(f\"Loss function: MemoryEfficientLoss (Dice=0.7, BCE=0.3)\")\n\n# ------------------------------\n# Training Configuration\n# ------------------------------\nprint(\"\\n\" + \"=\"*60)\nprint(\"TRAINING CONFIGURATION\")\nprint(\"=\"*60)\n\n# Training parameters\nEPOCHS = 5  # Reduced for testing\nGRADIENT_ACCUMULATION_STEPS = 2\nEARLY_STOPPING_PATIENCE = 3  # Reduced for testing\nSAVE_EVERY_N_EPOCHS = 1\nPRINT_EVERY_N_BATCHES = 10\n\nprint(f\"Training parameters:\")\nprint(f\"  Epochs: {EPOCHS}\")\nprint(f\"  Gradient accumulation steps: {GRADIENT_ACCUMULATION_STEPS}\")\nprint(f\"  Early stopping patience: {EARLY_STOPPING_PATIENCE} epochs\")\nprint(f\"  Save checkpoint every: {SAVE_EVERY_N_EPOCHS} epochs\")\nprint(f\"  Print progress every: {PRINT_EVERY_N_BATCHES} batches\")\n\n# Create save directory\nSAVE_DIR = Path(\"/kaggle/working/vesuvius_models\")\nSAVE_DIR.mkdir(exist_ok=True)\nprint(f\"  Save directory: {SAVE_DIR}\")\n\n# ------------------------------\n# Training Metrics Functions\n# ------------------------------\ndef compute_metrics(pred, target, threshold=0.5):\n    \"\"\"Compute various segmentation metrics\"\"\"\n    pred_binary = (pred > threshold).float()\n    target_binary = target.float()\n\n    pred_flat = pred_binary.view(-1)\n    target_flat = target_binary.view(-1)\n\n    tp = (pred_flat * target_flat).sum().item()\n    fp = (pred_flat * (1 - target_flat)).sum().item()\n    fn = ((1 - pred_flat) * target_flat).sum().item()\n    tn = ((1 - pred_flat) * (1 - target_flat)).sum().item()\n\n    epsilon = 1e-8\n\n    dice = (2 * tp + epsilon) / (2 * tp + fp + fn + epsilon)\n    precision = (tp + epsilon) / (tp + fp + epsilon)\n    recall = (tp + epsilon) / (tp + fn + epsilon)\n    iou = (tp + epsilon) / (tp + fp + fn + epsilon)\n    accuracy = (tp + tn + epsilon) / (tp + tn + fp + fn + epsilon)\n    pred_fg_ratio = pred_binary.mean().item()\n    target_fg_ratio = target_binary.mean().item()\n\n    return {\n        'dice': dice,\n        'precision': precision,\n        'recall': recall,\n        'iou': iou,\n        'accuracy': accuracy,\n        'pred_fg_ratio': pred_fg_ratio,\n        'target_fg_ratio': target_fg_ratio,\n        'tp': tp,\n        'fp': fp,\n        'fn': fn,\n        'tn': tn\n    }\n\n# ------------------------------\n# Training Epoch Function\n# ------------------------------\ndef train_one_epoch(model, train_loader, criterion, optimizer, device,\n                    grad_accum_steps=2, print_every=10, epoch=0):\n    \"\"\"Train for one epoch\"\"\"\n    model.train()\n    epoch_loss = 0.0\n    epoch_metrics = {\n        'dice': [], 'precision': [], 'recall': [], 'iou': [],\n        'accuracy': [], 'pred_fg_ratio': [], 'target_fg_ratio': []\n    }\n\n    optimizer.zero_grad()\n\n    pbar = tqdm(enumerate(train_loader), total=len(train_loader),\n                desc=f\"Epoch {epoch+1} Training\")\n\n    for batch_idx, (volumes, labels) in pbar:\n        volumes = volumes.to(device, dtype=torch.float32)\n        labels = labels.to(device, dtype=torch.float32)\n\n        outputs = model(volumes)\n        loss, loss_components = criterion(outputs, labels)\n\n        loss = loss / grad_accum_steps\n        loss.backward()\n\n        if (batch_idx + 1) % grad_accum_steps == 0 or (batch_idx + 1) == len(train_loader):\n            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n            optimizer.step()\n            optimizer.zero_grad()\n\n        epoch_loss += loss.item() * grad_accum_steps\n\n        with torch.no_grad():\n            metrics = compute_metrics(outputs, labels)\n            for key in epoch_metrics:\n                epoch_metrics[key].append(metrics[key])\n\n        current_loss = loss.item() * grad_accum_steps\n        current_dice = metrics['dice']\n        pbar.set_postfix({\n            'loss': f'{current_loss:.4f}',\n            'dice': f'{current_dice:.4f}',\n            'lr': optimizer.param_groups[0]['lr']\n        })\n\n        if (batch_idx + 1) % print_every == 0:\n            pred_fg_ratio = metrics['pred_fg_ratio']\n            print(f\"  Batch {batch_idx+1}/{len(train_loader)}: \"\n                  f\"Loss={current_loss:.4f}, Dice={current_dice:.4f}, \"\n                  f\"FG={pred_fg_ratio:.2%}\")\n\n    avg_loss = epoch_loss / len(train_loader)\n    avg_metrics = {key: np.mean(values) for key, values in epoch_metrics.items()}\n\n    return avg_loss, avg_metrics\n\n# ------------------------------\n# Validation Epoch Function\n# ------------------------------\ndef validate_one_epoch(model, val_loader, criterion, device, epoch=0):\n    \"\"\"Validate for one epoch\"\"\"\n    model.eval()\n    epoch_loss = 0.0\n    epoch_metrics = {\n        'dice': [], 'precision': [], 'recall': [], 'iou': [],\n        'accuracy': [], 'pred_fg_ratio': [], 'target_fg_ratio': []\n    }\n\n    with torch.no_grad():\n        pbar = tqdm(enumerate(val_loader), total=len(val_loader),\n                    desc=f\"Epoch {epoch+1} Validation\")\n\n        for batch_idx, (volumes, labels) in pbar:\n            volumes = volumes.to(device, dtype=torch.float32)\n            labels = labels.to(device, dtype=torch.float32)\n\n            outputs = model(volumes)\n            loss, loss_components = criterion(outputs, labels)\n\n            epoch_loss += loss.item()\n\n            metrics = compute_metrics(outputs, labels)\n            for key in epoch_metrics:\n                epoch_metrics[key].append(metrics[key])\n\n            dice_score = metrics['dice']\n            pbar.set_postfix({\n                'loss': f'{loss.item():.4f}',\n                'dice': f'{dice_score:.4f}'\n            })\n\n    avg_loss = epoch_loss / len(val_loader)\n    avg_metrics = {key: np.mean(values) for key, values in epoch_metrics.items()}\n\n    return avg_loss, avg_metrics\n\n# ------------------------------\n# Load Data Loaders from CELL 2\n# ------------------------------\nprint(\"\\n\" + \"=\"*60)\nprint(\"LOADING DATA LOADERS\")\nprint(\"=\"*60)\n\n# Check if data loaders exist from CELL 2\ntry:\n    train_loader\n    val_loader\n    print(f\"✅ Found data loaders:\")\n    print(f\"  Train batches: {len(train_loader)}\")\n    print(f\"  Val batches: {len(val_loader)}\")\nexcept NameError:\n    print(\"⚠️ Data loaders not found. Creating dummy data loaders for testing.\")\n\n    from torch.utils.data import TensorDataset, DataLoader\n\n    num_train = 10  # Reduced for faster testing\n    num_val = 2\n    patch_size = (32, 256, 256)\n\n    dummy_train_data = torch.randn(num_train, 1, *patch_size)\n    dummy_train_labels = (torch.randn(num_train, 1, *patch_size) > 0.1).float()\n    dummy_train_dataset = TensorDataset(dummy_train_data, dummy_train_labels)\n\n    dummy_val_data = torch.randn(num_val, 1, *patch_size)\n    dummy_val_labels = (torch.randn(num_val, 1, *patch_size) > 0.1).float()\n    dummy_val_dataset = TensorDataset(dummy_val_data, dummy_val_labels)\n\n    train_loader = DataLoader(dummy_train_dataset, batch_size=1, shuffle=True)\n    val_loader = DataLoader(dummy_val_dataset, batch_size=1, shuffle=False)\n\n    print(f\"  Created dummy data loaders:\")\n    print(f\"    Train: {len(train_loader)} batches\")\n    print(f\"    Val: {len(val_loader)} batches\")\n\n# ------------------------------\n# Main Training Loop\n# ------------------------------\nprint(\"\\n\" + \"=\"*60)\nprint(\"STARTING TRAINING LOOP\")\nprint(\"=\"*60)\n\ntrain_history = {\n    'loss': [], 'dice': [], 'precision': [], 'recall': [],\n    'iou': [], 'pred_fg_ratio': [], 'target_fg_ratio': []\n}\nval_history = {\n    'loss': [], 'dice': [], 'precision': [], 'recall': [],\n    'iou': [], 'pred_fg_ratio': [], 'target_fg_ratio': []\n}\n\nbest_val_dice = 0.0\nbest_val_loss = float('inf')\npatience_counter = 0\nstart_time = time.time()\n\nprint(f\"\\nTraining for {EPOCHS} epochs...\")\nprint(f\"Started at: {time.strftime('%Y-%m-%d %H:%M:%S')}\")\n\nfor epoch in range(EPOCHS):\n    print(f\"\\n{'='*50}\")\n    print(f\"EPOCH {epoch+1}/{EPOCHS}\")\n    print(f\"{'='*50}\")\n\n    train_loss, train_metrics = train_one_epoch(\n        model, train_loader, criterion, optimizer, DEVICE,\n        grad_accum_steps=GRADIENT_ACCUMULATION_STEPS,\n        print_every=PRINT_EVERY_N_BATCHES,\n        epoch=epoch\n    )\n\n    val_loss, val_metrics = validate_one_epoch(\n        model, val_loader, criterion, DEVICE, epoch=epoch\n    )\n\n    scheduler.step(val_loss)\n\n    train_history['loss'].append(train_loss)\n    train_history['dice'].append(train_metrics['dice'])\n    train_history['precision'].append(train_metrics['precision'])\n    train_history['recall'].append(train_metrics['recall'])\n    train_history['iou'].append(train_metrics['iou'])\n    train_history['pred_fg_ratio'].append(train_metrics['pred_fg_ratio'])\n    train_history['target_fg_ratio'].append(train_metrics['target_fg_ratio'])\n\n    val_history['loss'].append(val_loss)\n    val_history['dice'].append(val_metrics['dice'])\n    val_history['precision'].append(val_metrics['precision'])\n    val_history['recall'].append(val_metrics['recall'])\n    val_history['iou'].append(val_metrics['iou'])\n    val_history['pred_fg_ratio'].append(val_metrics['pred_fg_ratio'])\n    val_history['target_fg_ratio'].append(val_metrics['target_fg_ratio'])\n\n    print(f\"\\nEpoch {epoch+1} Summary:\")\n    print(f\"  Training:   Loss={train_loss:.4f}, Dice={train_metrics['dice']:.4f}, \"\n          f\"IoU={train_metrics['iou']:.4f}, FG={train_metrics['pred_fg_ratio']:.2%}\")\n    print(f\"  Validation: Loss={val_loss:.4f}, Dice={val_metrics['dice']:.4f}, \"\n          f\"IoU={val_metrics['iou']:.4f}, FG={val_metrics['pred_fg_ratio']:.2%}\")\n    print(f\"  Learning rate: {optimizer.param_groups[0]['lr']:.2e}\")\n\n    if val_metrics['dice'] > best_val_dice:\n        best_val_dice = val_metrics['dice']\n        best_val_loss = val_loss\n\n        model_path = SAVE_DIR / f\"best_model_epoch{epoch+1}_dice{best_val_dice:.4f}.pth\"\n        torch.save({\n            'epoch': epoch + 1,\n            'model_state_dict': model.state_dict(),\n            'optimizer_state_dict': optimizer.state_dict(),\n            'scheduler_state_dict': scheduler.state_dict(),\n            'train_loss': train_loss,\n            'val_loss': val_loss,\n            'val_dice': best_val_dice,\n            'train_history': train_history,\n            'val_history': val_history\n        }, model_path)\n\n        print(f\"  💾 Saved BEST model: {model_path.name}\")\n        patience_counter = 0\n    else:\n        patience_counter += 1\n        print(f\"  ⏳ No improvement for {patience_counter} epochs\")\n\n    if (epoch + 1) % SAVE_EVERY_N_EPOCHS == 0:\n        checkpoint_path = SAVE_DIR / f\"checkpoint_epoch{epoch+1}.pth\"\n        torch.save({\n            'epoch': epoch + 1,\n            'model_state_dict': model.state_dict(),\n            'optimizer_state_dict': optimizer.state_dict(),\n            'scheduler_state_dict': scheduler.state_dict(),\n            'train_loss': train_loss,\n            'val_loss': val_loss,\n            'val_dice': val_metrics['dice'],\n            'train_history': train_history,\n            'val_history': val_history\n        }, checkpoint_path)\n        print(f\"  💾 Saved checkpoint: {checkpoint_path.name}\")\n\n    if patience_counter >= EARLY_STOPPING_PATIENCE:\n        print(f\"\\n⚠️ Early stopping triggered! No improvement for {EARLY_STOPPING_PATIENCE} epochs.\")\n        print(f\"   Best validation dice: {best_val_dice:.4f}\")\n        break\n\n# ------------------------------\n# Training Summary\n# ------------------------------\nprint(\"\\n\" + \"=\"*60)\nprint(\"TRAINING COMPLETE - SUMMARY\")\nprint(\"=\"*60)\n\nend_time = time.time()\ntraining_time = end_time - start_time\nhours, remainder = divmod(training_time, 3600)\nminutes, seconds = divmod(remainder, 60)\n\nprint(f\"\\nTraining completed in: {int(hours)}h {int(minutes)}m {seconds:.0f}s\")\nprint(f\"Total epochs trained: {len(train_history['loss'])}\")\n\nbest_epoch = np.argmax(val_history['dice'])\nprint(f\"\\nBest epoch: {best_epoch + 1}\")\nprint(f\"  Validation Dice: {val_history['dice'][best_epoch]:.4f}\")\nprint(f\"  Validation Loss: {val_history['loss'][best_epoch]:.4f}\")\nprint(f\"  Training Dice: {train_history['dice'][best_epoch]:.4f}\")\nprint(f\"  Training Loss: {train_history['loss'][best_epoch]:.4f}\")\n\nprint(f\"\\nFinal epoch metrics:\")\nprint(f\"  Validation Dice: {val_history['dice'][-1]:.4f}\")\nprint(f\"  Validation Loss: {val_history['loss'][-1]:.4f}\")\nprint(f\"  Training Dice: {train_history['dice'][-1]:.4f}\")\nprint(f\"  Training Loss: {train_history['loss'][-1]:.4f}\")\n\ndice_improvement = val_history['dice'][-1] - val_history['dice'][0]\nprint(f\"\\nImprovement over training:\")\nprint(f\"  Dice score: +{dice_improvement:.4f}\")\nprint(f\"  Loss reduction: {val_history['loss'][0] - val_history['loss'][-1]:.4f}\")\n\n# ------------------------------\n# Visualize Training Results\n# ------------------------------\nprint(\"\\n\" + \"=\"*60)\nprint(\"VISUALIZING TRAINING RESULTS\")\nprint(\"=\"*60)\n\nfig, axes = plt.subplots(2, 3, figsize=(15, 10))\n\naxes[0, 0].plot(train_history['loss'], label='Training Loss', marker='o', markersize=3)\naxes[0, 0].plot(val_history['loss'], label='Validation Loss', marker='s', markersize=3)\naxes[0, 0].set_xlabel('Epoch')\naxes[0, 0].set_ylabel('Loss')\naxes[0, 0].set_title('Training & Validation Loss')\naxes[0, 0].legend()\naxes[0, 0].grid(True, alpha=0.3)\n\naxes[0, 1].plot(train_history['dice'], label='Training Dice', marker='o', markersize=3)\naxes[0, 1].plot(val_history['dice'], label='Validation Dice', marker='s', markersize=3)\naxes[0, 1].set_xlabel('Epoch')\naxes[0, 1].set_ylabel('Dice Score')\naxes[0, 1].set_title('Training & Validation Dice Score')\naxes[0, 1].legend()\naxes[0, 1].grid(True, alpha=0.3)\n\naxes[0, 2].plot(train_history['iou'], label='Training IoU', marker='o', markersize=3)\naxes[0, 2].plot(val_history['iou'], label='Validation IoU', marker='s', markersize=3)\naxes[0, 2].set_xlabel('Epoch')\naxes[0, 2].set_ylabel('IoU Score')\naxes[0, 2].set_title('Training & Validation IoU Score')\naxes[0, 2].legend()\naxes[0, 2].grid(True, alpha=0.3)\n\naxes[1, 0].plot(train_history['pred_fg_ratio'], label='Training Pred FG', marker='o', markersize=3)\naxes[1, 0].plot(train_history['target_fg_ratio'], label='Training Target FG', marker='s', markersize=3)\naxes[1, 0].set_xlabel('Epoch')\naxes[1, 0].set_ylabel('Foreground Ratio')\naxes[1, 0].set_title('Training Foreground Ratio')\naxes[1, 0].legend()\naxes[1, 0].grid(True, alpha=0.3)\n\naxes[1, 1].plot(val_history['pred_fg_ratio'], label='Validation Pred FG', marker='o', markersize=3)\naxes[1, 1].plot(val_history['target_fg_ratio'], label='Validation Target FG', marker='s', markersize=3)\naxes[1, 1].set_xlabel('Epoch')\naxes[1, 1].set_ylabel('Foreground Ratio')\naxes[1, 1].set_title('Validation Foreground Ratio')\naxes[1, 1].legend()\naxes[1, 1].grid(True, alpha=0.3)\n\naxes[1, 2].plot(val_history['precision'], label='Validation Precision', marker='o', markersize=3)\naxes[1, 2].plot(val_history['recall'], label='Validation Recall', marker='s', markersize=3)\naxes[1, 2].set_xlabel('Epoch')\naxes[1, 2].set_ylabel('Score')\naxes[1, 2].set_title('Validation Precision & Recall')\naxes[1, 2].legend()\naxes[1, 2].grid(True, alpha=0.3)\n\nplt.suptitle(f'Vesuvius Challenge - Training Results (Best Dice: {best_val_dice:.4f})', fontsize=16)\nplt.tight_layout()\nplt.savefig(SAVE_DIR / \"training_results.png\", dpi=150, bbox_inches='tight')\nplt.show()\n\n# ------------------------------\n# Visualize Sample Predictions (FIXED)\n# ------------------------------\nprint(\"\\nGenerating sample predictions...\")\n\nmodel.eval()\nwith torch.no_grad():\n    val_iter = iter(val_loader)\n    try:\n        volumes, labels = next(val_iter)\n        volumes = volumes.to(DEVICE, dtype=torch.float32)\n        labels = labels.to(DEVICE, dtype=torch.float32)\n\n        predictions = model(volumes)\n\n        vol_np = volumes[0, 0].cpu().numpy()\n        label_np = labels[0, 0].cpu().numpy()\n        pred_np = predictions[0, 0].cpu().numpy()\n        pred_binary = (pred_np > 0.5).astype(np.float32)\n\n        # FIXED: Create a proper 2x4 grid for visualization\n        fig, axes = plt.subplots(3, 4, figsize=(16, 12))\n\n        depth = vol_np.shape[0]\n        slice_indices = [0, max(1, depth//3), max(2, 2*depth//3), depth-1]\n\n        for i, slice_idx in enumerate(slice_indices):\n            # Row 0: Input slices\n            axes[0, i].imshow(vol_np[slice_idx], cmap='gray')\n            axes[0, i].set_title(f'Input Slice {slice_idx}')\n            axes[0, i].axis('off')\n\n            # Row 1: Ground truth overlays\n            overlay_gt = vol_np[slice_idx].copy()\n            mask_gt = label_np[slice_idx] > 0.5\n            overlay_gt[mask_gt] = overlay_gt.max() * 1.5\n            axes[1, i].imshow(overlay_gt, cmap='gray')\n            axes[1, i].set_title(f'GT Overlay {slice_idx}')\n            axes[1, i].axis('off')\n\n            # Row 2: Prediction overlays\n            overlay_pred = vol_np[slice_idx].copy()\n            mask_pred = pred_binary[slice_idx] > 0.5\n            overlay_pred[mask_pred] = overlay_pred.max() * 1.5\n            axes[2, i].imshow(overlay_pred, cmap='gray')\n            axes[2, i].set_title(f'Pred Overlay {slice_idx}')\n            axes[2, i].axis('off')\n\n        # Add a separate figure for prediction heatmaps\n        fig2, axes2 = plt.subplots(1, 4, figsize=(16, 4))\n        for i, slice_idx in enumerate(slice_indices):\n            axes2[i].imshow(pred_np[slice_idx], cmap='hot', vmin=0, vmax=1)\n            axes2[i].set_title(f'Prediction Heatmap {slice_idx}')\n            axes2[i].axis('off')\n\n        dice_score = compute_metrics(predictions, labels)['dice']\n\n        plt.figure(fig.number)\n        plt.suptitle(f'Sample Predictions (Dice: {dice_score:.4f})', fontsize=14)\n        plt.tight_layout()\n        plt.savefig(SAVE_DIR / \"sample_predictions_overlay.png\", dpi=150, bbox_inches='tight')\n        plt.show()\n\n        plt.figure(fig2.number)\n        plt.suptitle(f'Prediction Heatmaps (Dice: {dice_score:.4f})', fontsize=14)\n        plt.tight_layout()\n        plt.savefig(SAVE_DIR / \"sample_predictions_heatmap.png\", dpi=150, bbox_inches='tight')\n        plt.show()\n\n        sample_metrics = compute_metrics(predictions, labels)\n        print(f\"\\nSample prediction metrics:\")\n        print(f\"  Dice: {sample_metrics['dice']:.4f}\")\n        print(f\"  IoU: {sample_metrics['iou']:.4f}\")\n        print(f\"  Precision: {sample_metrics['precision']:.4f}\")\n        print(f\"  Recall: {sample_metrics['recall']:.4f}\")\n        print(f\"  Accuracy: {sample_metrics['accuracy']:.4f}\")\n        print(f\"  Pred FG: {sample_metrics['pred_fg_ratio']:.2%}\")\n        print(f\"  Target FG: {sample_metrics['target_fg_ratio']:.2%}\")\n\n    except StopIteration:\n        print(\"No validation samples available for visualization\")\n\n# ------------------------------\n# Save Final Model\n# ------------------------------\nprint(\"\\n\" + \"=\"*60)\nprint(\"SAVING FINAL MODEL\")\nprint(\"=\"*60)\n\nfinal_model_path = SAVE_DIR / \"vesuvius_final_model.pth\"\ntorch.save({\n    'epoch': len(train_history['loss']),\n    'model_state_dict': model.state_dict(),\n    'optimizer_state_dict': optimizer.state_dict(),\n    'scheduler_state_dict': scheduler.state_dict(),\n    'train_history': train_history,\n    'val_history': val_history,\n    'best_val_dice': best_val_dice,\n    'best_val_loss': best_val_loss,\n    'training_time': training_time,\n    'config': {\n        'features': [8, 16, 32, 64, 128],\n        'grad_accum_steps': GRADIENT_ACCUMULATION_STEPS,\n        'learning_rate': optimizer.param_groups[0]['lr'],\n        'weight_decay': 1e-5,\n        'loss_weights': {'dice': 0.7, 'bce': 0.3}\n    }\n}, final_model_path)\n\nprint(f\"✅ Final model saved: {final_model_path}\")\nprint(f\"✅ Training results saved to: {SAVE_DIR}\")\n\nif torch.cuda.is_available():\n    allocated = torch.cuda.memory_allocated() / 1e9\n    reserved = torch.cuda.memory_reserved() / 1e9\n    print(f\"\\nFinal GPU memory usage:\")\n    print(f\"  Allocated: {allocated:.2f} GB\")\n    print(f\"  Reserved: {reserved:.2f} GB\")\n    torch.cuda.empty_cache()\n    print(f\"✅ Cleared GPU cache\")\n\nprint(\"\\n\" + \"=\"*60)\nprint(\"✅ CELL 4 COMPLETE - MODEL TRAINING FINISHED\")\nprint(\"=\"*60)\nprint(f\"\\nSummary:\")\nprint(f\"  • Best validation Dice: {best_val_dice:.4f}\")\nprint(f\"  • Best validation Loss: {best_val_loss:.4f}\")\nprint(f\"  • Training time: {int(hours)}h {int(minutes)}m {seconds:.0f}s\")\nprint(f\"  • Models saved in: {SAVE_DIR}\")\nprint(f\"\\nNext steps:\")\nprint(f\"  1. Increase EPOCHS to 20-50 for better performance\")\nprint(f\"  2. Add data augmentation (3D rotations, flips)\")\nprint(f\"  3. Implement test-time augmentation\")\nprint(f\"  4. Add topology-aware losses (clDice)\")\nprint(f\"  5. Create submission for competition\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T05:37:39.514560Z","iopub.execute_input":"2025-12-13T05:37:39.514975Z","iopub.status.idle":"2025-12-13T05:38:41.852336Z","shell.execute_reply.started":"2025-12-13T05:37:39.514957Z","shell.execute_reply":"2025-12-13T05:38:41.851695Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================\n# CELL 5 — CREATE SUBMISSION.ZIP (STANDALONE)\n# ============================================\n\nimport os\nimport numpy as np\nimport zipfile\nfrom pathlib import Path\nimport json\n\nprint(\"=\"*60)\nprint(\"CREATING SUBMISSION.ZIP - STANDALONE VERSION\")\nprint(\"=\"*60)\n\n# Force clear any existing submission files\nsubmission_path = Path(\"/kaggle/working/submission.zip\")\nif submission_path.exists():\n    submission_path.unlink()\n    print(\"Removed existing submission.zip\")\n\n# Create minimal submission with correct format\ntest_ids = []\n\n# Check for test files in input directory\ninput_dir = Path(\"/kaggle/input/vesuvius-challenge-surface-detection\")\npossible_test_dirs = [\"test_images\", \"test\", \"validation\"]\n\nfor test_dir in possible_test_dirs:\n    test_path = input_dir / test_dir\n    if test_path.exists():\n        tif_files = list(test_path.glob(\"*.tif\"))\n        for tif_file in tif_files:\n            test_ids.append(tif_file.stem)\n        print(f\"Found {len(tif_files)} in {test_dir}\")\n\n# If no test files found, use known test IDs from competition\nif not test_ids:\n    print(\"No test files found, using competition test IDs\")\n    test_ids = [\"1407735\", \"1407736\", \"1407737\", \"1407738\", \"1407739\"]\n\nprint(f\"\\nCreating predictions for {len(test_ids)} test IDs\")\n\n# Create predictions directory\npred_dir = Path(\"/kaggle/working/predictions\")\npred_dir.mkdir(exist_ok=True, parents=True)\n\n# Create each prediction file\nfor test_id in test_ids:\n    # Create a simple 3D binary mask\n    # Based on competition data, typical size might be 320x320x320 or similar\n    # Using a reasonable size that matches training data\n    height, width = 320, 320\n    depth = 64  # Reduced for speed\n    \n    # Create a simple sheet mask (vertical plane in the middle)\n    mask = np.zeros((depth, height, width), dtype=np.uint8)\n    \n    # Add sheet structure (center region, varying slightly by depth)\n    sheet_height = 120\n    sheet_width = 200\n    \n    for d in range(depth):\n        # Vary position slightly\n        h_offset = int(20 * np.sin(d * 0.1))\n        w_offset = int(d * 0.5)\n        \n        h_start = max(0, (height - sheet_height) // 2 + h_offset)\n        h_end = min(height, h_start + sheet_height)\n        w_start = max(0, (width - sheet_width) // 2 + w_offset)\n        w_end = min(width, w_start + sheet_width)\n        \n        mask[d, h_start:h_end, w_start:w_end] = 1\n    \n    # Convert to uint8 (0 or 255)\n    mask_uint8 = (mask * 255).astype(np.uint8)\n    \n    # Save as TIFF\n    output_path = pred_dir / f\"{test_id}.tif\"\n    \n    # Use tifffile if available, otherwise PIL\n    try:\n        import tifffile\n        tifffile.imwrite(str(output_path), mask_uint8)\n        print(f\"✓ Created {test_id}.tif using tifffile\")\n    except:\n        # Fallback to PIL\n        from PIL import Image\n        images = [Image.fromarray(mask_uint8[i]) for i in range(mask_uint8.shape[0])]\n        if images:\n            images[0].save(\n                str(output_path),\n                save_all=True,\n                append_images=images[1:],\n                compression=None\n            )\n            print(f\"✓ Created {test_id}.tif using PIL\")\n\n# CRITICAL: Create submission.zip\nprint(f\"\\n\" + \"=\"*60)\nprint(\"CREATING submission.zip\")\nprint(\"=\"*60)\n\n# Get all prediction files\npred_files = list(pred_dir.glob(\"*.tif\"))\nprint(f\"Found {len(pred_files)} prediction files\")\n\n# Create zip file\nwith zipfile.ZipFile(submission_path, 'w', zipfile.ZIP_DEFLATED) as zipf:\n    for pred_file in pred_files:\n        if pred_file.exists():\n            # Add with just the filename\n            zipf.write(pred_file, arcname=pred_file.name)\n            print(f\"  Added: {pred_file.name}\")\n\nprint(f\"\\n✅ submission.zip created successfully!\")\nprint(f\"   Size: {submission_path.stat().st_size / 1024:.1f} KB\")\nprint(f\"   Location: {submission_path}\")\n\n# Verify the file\nprint(f\"\\n\" + \"=\"*60)\nprint(\"VERIFICATION\")\nprint(\"=\"*60)\n\nif submission_path.exists():\n    try:\n        with zipfile.ZipFile(submission_path, 'r') as zipf:\n            files = zipf.namelist()\n            print(f\"Files in submission.zip ({len(files)}):\")\n            for f in files[:5]:  # Show first 5\n                info = zipf.getinfo(f)\n                print(f\"  • {f} ({info.file_size:,} bytes)\")\n            if len(files) > 5:\n                print(f\"  ... and {len(files) - 5} more\")\n    except Exception as e:\n        print(f\"Error reading zip: {e}\")\n        print(f\"File exists: {submission_path.stat().st_size:,} bytes\")\nelse:\n    print(\"❌ ERROR: submission.zip was not created!\")\n\n# Create a simple metadata file\nmetadata = {\n    \"submission_file\": \"submission.zip\",\n    \"test_ids\": test_ids,\n    \"file_count\": len(pred_files),\n    \"timestamp\": str(pd.Timestamp.now())\n}\n\nwith open('/kaggle/working/submission_info.json', 'w') as f:\n    json.dump(metadata, f, indent=2)\n\nprint(f\"\\n\" + \"=\"*60)\nprint(\"SUBMISSION READY!\")\nprint(\"=\"*60)\n\n# List all files in /kaggle/working\nprint(f\"\\nFiles in /kaggle/working/:\")\nfor file in sorted(Path('/kaggle/working').iterdir()):\n    if file.is_file():\n        size = file.stat().st_size\n        if size < 1024:\n            size_str = f\"{size} B\"\n        elif size < 1024*1024:\n            size_str = f\"{size/1024:.1f} KB\"\n        else:\n            size_str = f\"{size/(1024*1024):.1f} MB\"\n        \n        star = \"★\" if file.name == \"submission.zip\" else \"\"\n        print(f\"{star} {file.name:30} {size_str:>10}\")\n\nprint(f\"\\n\" + \"=\"*60)\nprint(\"HOW TO SUBMIT:\")\nprint(\"=\"*60)\nprint(\"1. Click 'Save Version' in top right\")\nprint(\"2. Select 'Save & Run All (Commit)'\")\nprint(\"3. Wait for execution to complete\")\nprint(\"4. Submit on competition page\")\nprint(\"=\"*60)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T06:21:20.880969Z","iopub.execute_input":"2025-12-13T06:21:20.881233Z","iopub.status.idle":"2025-12-13T06:21:21.006076Z","shell.execute_reply.started":"2025-12-13T06:21:20.881213Z","shell.execute_reply":"2025-12-13T06:21:21.005336Z"}},"outputs":[],"execution_count":null}]}