{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":61446,"databundleVersionId":6962461,"sourceType":"competition"}],"dockerImageVersionId":30886,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nfrom pathlib import Path\nfrom typing import List, Tuple, Dict\nimport tifffile\nfrom tqdm import tqdm\nimport warnings\nwarnings.filterwarnings(\"ignore\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-02-27T01:01:21.299424Z","iopub.execute_input":"2025-02-27T01:01:21.299772Z","iopub.status.idle":"2025-02-27T01:01:23.733586Z","shell.execute_reply.started":"2025-02-27T01:01:21.299748Z","shell.execute_reply":"2025-02-27T01:01:23.732936Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class KidneyVesselEDA:\n    def __init__(self, base_path: str):\n        \"\"\"\n        Initialize the EDA class\n        \n        Args:\n            base_path (str): Path to the Kaggle dataset directory\n        \"\"\"\n        self.base_path = Path(base_path)\n        self.train_path = self.base_path / 'train'\n        self.datasets = ['kidney_1_dense', 'kidney_1_voi', 'kidney_2', \n                        'kidney_3_dense', 'kidney_3_sparse']\n        \n    def load_dataset_info(self) -> Dict:\n        \"\"\"\n        Load basic information about each dataset\n        \n        Returns:\n            Dict: Dictionary containing dataset statistics\n        \"\"\"\n        dataset_info = {}\n        \n        for dataset in self.datasets:\n            dataset_path = self.train_path / dataset\n            if dataset_path.exists():\n                images_path = dataset_path / 'images'\n                labels_path = dataset_path / 'labels'\n                \n                if images_path.exists():\n                    n_images = len(list(images_path.glob('*.tif')))\n                else:\n                    n_images = 0\n                    \n                if labels_path.exists():\n                    n_labels = len(list(labels_path.glob('*.tif')))\n                else:\n                    n_labels = 0\n                \n                # Get image dimensions from first image\n                if n_images > 0:\n                    first_image = tifffile.imread(str(next(images_path.glob('*.tif'))))\n                    dimensions = first_image.shape\n                else:\n                    dimensions = None\n                \n                dataset_info[dataset] = {\n                    'n_images': n_images,\n                    'n_labels': n_labels,\n                    'dimensions': dimensions\n                }\n        \n        return dataset_info\n    \n    def analyze_class_distribution(self, dataset: str) -> Tuple[float, float]:\n        \"\"\"\n        Analyze the class distribution (vessel vs non-vessel) in a dataset\n        \n        Args:\n            dataset (str): Name of the dataset to analyze\n            \n        Returns:\n            Tuple[float, float]: Percentage of vessel and non-vessel pixels\n        \"\"\"\n        labels_path = self.train_path / dataset / 'labels'\n        if not labels_path.exists():\n            return None\n            \n        total_pixels = 0\n        vessel_pixels = 0\n        \n        for label_file in tqdm(list(labels_path.glob('*.tif')), desc=f'Analyzing {dataset}'):\n            mask = tifffile.imread(str(label_file))\n            total_pixels += mask.size\n            vessel_pixels += np.sum(mask > 0)\n        \n        vessel_percentage = (vessel_pixels / total_pixels) * 100\n        non_vessel_percentage = 100 - vessel_percentage\n        \n        return vessel_percentage, non_vessel_percentage\n    \n    def visualize_sample_slices(self, dataset: str, n_samples: int = 5) -> None:\n        \"\"\"\n        Visualize sample slices from a dataset with their corresponding masks\n        \n        Args:\n            dataset (str): Name of the dataset to visualize\n            n_samples (int): Number of samples to visualize\n        \"\"\"\n        images_path = self.train_path / dataset / 'images'\n        labels_path = self.train_path / dataset / 'labels'\n        \n        if not images_path.exists() or not labels_path.exists():\n            print(f\"Dataset {dataset} not found or incomplete\")\n            return\n            \n        image_files = sorted(list(images_path.glob('*.tif')))\n        label_files = sorted(list(labels_path.glob('*.tif')))\n        \n        # Select evenly spaced samples\n        indices = np.linspace(0, len(image_files)-1, n_samples, dtype=int)\n        \n        fig, axes = plt.subplots(n_samples, 2, figsize=(10, 3*n_samples))\n        fig.suptitle(f'Sample Slices from {dataset}')\n        \n        for idx, (ax_row, i) in enumerate(zip(axes, indices)):\n            # Load image and mask\n            image = tifffile.imread(str(image_files[i]))\n            mask = tifffile.imread(str(label_files[i]))\n            \n            # Display image\n            ax_row[0].imshow(image, cmap='gray')\n            ax_row[0].set_title(f'Slice {str(image_files[i]).split(\"/\")[-1]}')\n            ax_row[0].axis('off')\n            \n            # Display mask\n            ax_row[1].imshow(mask, cmap='binary')\n            ax_row[1].set_title(f'Mask {str(label_files[i]).split(\"/\")[-1]}')\n            ax_row[1].axis('off')\n        \n        plt.tight_layout()\n        plt.show()\n    \n    def plot_class_distribution(self, distribution_data: Dict) -> None:\n        \"\"\"\n        Plot the class distribution across datasets\n        \n        Args:\n            distribution_data (Dict): Dictionary containing class distribution data\n        \"\"\"\n        datasets = list(distribution_data.keys())\n        vessel_percentages = [d[0] for d in distribution_data.values()]\n        non_vessel_percentages = [d[1] for d in distribution_data.values()]\n        \n        fig, ax = plt.subplots(figsize=(12, 6))\n        width = 0.35\n        \n        ax.bar(datasets, vessel_percentages, width, label='Vessel')\n        ax.bar(datasets, non_vessel_percentages, width, bottom=vessel_percentages, label='Non-vessel')\n        \n        ax.set_ylabel('Percentage')\n        ax.set_title('Class Distribution Across Datasets')\n        ax.legend()\n        \n        plt.xticks(rotation=45)\n        plt.tight_layout()\n        plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T01:01:27.473861Z","iopub.execute_input":"2025-02-27T01:01:27.474720Z","iopub.status.idle":"2025-02-27T01:01:27.487760Z","shell.execute_reply.started":"2025-02-27T01:01:27.474684Z","shell.execute_reply":"2025-02-27T01:01:27.486891Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def main():\n    # Initialize EDA class with Kaggle dataset path\n    kaggle_path = '/kaggle/input/blood-vessel-segmentation'\n    eda = KidneyVesselEDA(kaggle_path)\n    \n    # Get dataset information\n    print(\"Analyzing dataset information...\")\n    dataset_info = eda.load_dataset_info()\n    \n    # Print dataset statistics\n    print(\"\\nDataset Statistics:\")\n    for dataset, info in dataset_info.items():\n        print(f\"\\n{dataset}:\")\n        print(f\"  Number of images: {info['n_images']}\")\n        print(f\"  Number of labels: {info['n_labels']}\")\n        print(f\"  Image dimensions: {info['dimensions']}\")\n    \n    # Analyze class distribution for training datasets\n    print(\"\\nAnalyzing class distribution...\")\n    distribution_data = {}\n    for dataset in eda.datasets:\n        if dataset_info[dataset]['n_labels'] > 0:\n            distribution = eda.analyze_class_distribution(dataset)\n            if distribution:\n                distribution_data[dataset] = distribution\n                print(f\"\\n{dataset}:\")\n                print(f\"  Vessel pixels: {distribution[0]:.2f}%\")\n                print(f\"  Non-vessel pixels: {distribution[1]:.2f}%\")\n    \n    # Plot class distribution\n    eda.plot_class_distribution(distribution_data)\n    \n    # Visualize sample slices from each dataset\n    print(\"\\nVisualizing sample slices...\")\n    for dataset in eda.datasets:\n        if dataset_info[dataset]['n_labels'] > 0:\n            print(f\"\\nVisualizing {dataset}...\")\n            eda.visualize_sample_slices(dataset)\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T01:01:32.915959Z","iopub.execute_input":"2025-02-27T01:01:32.916292Z","iopub.status.idle":"2025-02-27T01:05:41.648033Z","shell.execute_reply.started":"2025-02-27T01:01:32.916269Z","shell.execute_reply":"2025-02-27T01:05:41.647112Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df = pd.read_csv(\"/kaggle/input/blood-vessel-segmentation/train_rles.csv\")\ntrain_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T01:06:36.652789Z","iopub.execute_input":"2025-02-27T01:06:36.653159Z","iopub.status.idle":"2025-02-27T01:06:37.884521Z","shell.execute_reply.started":"2025-02-27T01:06:36.653127Z","shell.execute_reply":"2025-02-27T01:06:37.883548Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Data Preprocessing","metadata":{}},{"cell_type":"code","source":"!pip install -q -U albumentations","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T01:06:42.149633Z","iopub.execute_input":"2025-02-27T01:06:42.149988Z","iopub.status.idle":"2025-02-27T01:06:48.230798Z","shell.execute_reply.started":"2025-02-27T01:06:42.149959Z","shell.execute_reply":"2025-02-27T01:06:48.229668Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.model_selection import train_test_split\nimport numpy as np\nimport tifffile\nfrom pathlib import Path\nfrom typing import Tuple, List, Dict, Optional\nfrom scipy.ndimage import binary_dilation, binary_closing\nimport albumentations as A\nfrom tqdm import tqdm\nimport os\nimport gc\nimport matplotlib.pyplot as plt\nimport cv2\nfrom sklearn.model_selection import train_test_split\nimport json\nimport warnings\nwarnings.filterwarnings(\"ignore\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T01:06:51.067435Z","iopub.execute_input":"2025-02-27T01:06:51.067793Z","iopub.status.idle":"2025-02-27T01:06:56.073659Z","shell.execute_reply.started":"2025-02-27T01:06:51.067760Z","shell.execute_reply":"2025-02-27T01:06:56.072944Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def clear_memory():\n    \"\"\"Aggressively clear memory\"\"\"\n    plt.close('all')\n    gc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T01:06:56.074645Z","iopub.execute_input":"2025-02-27T01:06:56.075368Z","iopub.status.idle":"2025-02-27T01:06:56.079036Z","shell.execute_reply.started":"2025-02-27T01:06:56.075335Z","shell.execute_reply":"2025-02-27T01:06:56.078313Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class_weights = {\n    'kidney_1_dense': {'dice_weight': 0.7, 'bce_weight': 0.3},\n    'kidney_1_voi': {'dice_weight': 0.6, 'bce_weight': 0.4},\n    'kidney_2': {'dice_weight': 0.7, 'bce_weight': 0.3},\n    'kidney_3_sparse': {'dice_weight': 0.8, 'bce_weight': 0.2}\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T01:06:58.904660Z","iopub.execute_input":"2025-02-27T01:06:58.904984Z","iopub.status.idle":"2025-02-27T01:06:58.909159Z","shell.execute_reply.started":"2025-02-27T01:06:58.904961Z","shell.execute_reply":"2025-02-27T01:06:58.908258Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class VesselDataset(Dataset):\n    \"\"\"\n    Dataset class for vessel segmentation data\n    \"\"\"\n    def __init__(self, \n                 image_files: List[Path], \n                 label_files: List[Path], \n                 transform: Optional[A.Compose] = None, \n                 size: Tuple[int, int] = (512, 512)):\n        \"\"\"\n        Initialize the dataset\n        \n        Args:\n            image_files (List[Path]): List of paths to image files\n            label_files (List[Path]): List of paths to label files\n            transform (Optional[A.Compose]): Albumentations transforms to apply\n            size (Tuple[int, int]): Target size for resizing (height, width)\n        \"\"\"\n        self.image_files = image_files\n        self.label_files = label_files\n        self.transform = transform\n        self.target_h, self.target_w = size\n        \n        # Validate files\n        self._validate_files()\n        \n        # Get original dimensions from first image\n        try:\n            first_image = tifffile.imread(str(image_files[0]))\n            orig_h, orig_w = first_image.shape\n            print(f\"Dataset: Original dimensions: {orig_h}x{orig_w}, Target dimensions: {self.target_h}x{self.target_w}\")\n            \n            # Clear memory\n            del first_image\n            gc.collect()\n            \n        except Exception as e:\n            print(f\"Warning: Could not read first image: {str(e)}\")\n    \n    def _validate_files(self):\n        \"\"\"Validate that all files exist and match\"\"\"\n        if len(self.image_files) != len(self.label_files):\n            raise ValueError(f\"Number of images ({len(self.image_files)}) \"\n                           f\"!= number of labels ({len(self.label_files)})\")\n        \n        # Check all files exist\n        for img_path, label_path in zip(self.image_files, self.label_files):\n            if not Path(img_path).exists():\n                raise FileNotFoundError(f\"Image file not found: {img_path}\")\n            if not Path(label_path).exists():\n                raise FileNotFoundError(f\"Label file not found: {label_path}\")\n    \n    def __len__(self) -> int:\n        \"\"\"Return the total number of samples\"\"\"\n        return len(self.image_files)\n    \n    def preprocess_image(self, \n                        image: np.ndarray, \n                        is_mask: bool = False) -> np.ndarray:\n        \"\"\"\n        Preprocess image to target size\n        \n        Args:\n            image (np.ndarray): Input image or mask\n            is_mask (bool): Whether the input is a mask\n            \n        Returns:\n            np.ndarray: Preprocessed image or mask\n        \"\"\"\n        try:\n            if is_mask:\n                # Use nearest neighbor for masks to preserve binary values\n                processed = cv2.resize(\n                    image.astype(np.uint8),\n                    (self.target_w, self.target_h),\n                    interpolation=cv2.INTER_NEAREST\n                )\n            else:\n                # Use bilinear interpolation for images\n                processed = cv2.resize(\n                    image.astype(np.float32),\n                    (self.target_w, self.target_h),\n                    interpolation=cv2.INTER_LINEAR\n                )\n            return processed\n            \n        except Exception as e:\n            print(f\"Error in preprocessing: {str(e)}\")\n            # Return zero array of correct shape and type\n            return np.zeros((self.target_h, self.target_w), \n                          dtype=np.uint8 if is_mask else np.float32)\n    \n    def normalize_image(self, image: np.ndarray) -> np.ndarray:\n        \"\"\"\n        Normalize image to [0,1] range with safe division\n        \n        Args:\n            image (np.ndarray): Input image\n            \n        Returns:\n            np.ndarray: Normalized image\n        \"\"\"\n        try:\n            image_min = image.min()\n            image_max = image.max()\n            \n            if image_max - image_min == 0:\n                return np.zeros_like(image, dtype=np.float32)\n                \n            return (image - image_min) / (image_max - image_min)\n            \n        except Exception as e:\n            print(f\"Error in normalization: {str(e)}\")\n            return np.zeros_like(image, dtype=np.float32)\n    \n    def __getitem__(self, idx: int) -> Tuple[torch.Tensor, torch.Tensor]:\n        \"\"\"\n        Get a sample from the dataset\n        \n        Args:\n            idx (int): Index of the sample\n            \n        Returns:\n            Tuple[torch.Tensor, torch.Tensor]: Image and mask tensors\n        \"\"\"\n        try:\n            # Load image and mask with memory management\n            image = tifffile.imread(str(self.image_files[idx]))\n            mask = tifffile.imread(str(self.label_files[idx]))\n            \n            # Basic input validation\n            if image is None or mask is None:\n                raise ValueError(\"Failed to load image or mask\")\n            \n            if image.size == 0 or mask.size == 0:\n                raise ValueError(\"Empty image or mask\")\n            \n            # Normalize image\n            image = self.normalize_image(image)\n            \n            # Convert mask to binary\n            mask = (mask > 0).astype(np.uint8)\n            \n            # Resize both to target size\n            image = self.preprocess_image(image, is_mask=False)\n            mask = self.preprocess_image(mask, is_mask=True)\n            \n            # Apply transforms if specified\n            if self.transform:\n                transformed = self.transform(\n                    image=image.astype(np.float32),\n                    mask=mask\n                )\n                image = transformed['image']\n                mask = transformed['mask']\n            \n            # Convert to tensors\n            image = torch.from_numpy(image).float().unsqueeze(0)\n            mask = torch.from_numpy(mask).float().unsqueeze(0)\n            \n            # Validate output tensors\n            if torch.isnan(image).any() or torch.isnan(mask).any():\n                raise ValueError(\"NaN values in output tensors\")\n            \n            # Clear memory\n            gc.collect()\n            \n            return image, mask\n            \n        except Exception as e:\n            print(f\"Error loading sample {idx} from {self.image_files[idx]}: {str(e)}\")\n            # Return zero tensors in case of error\n            return (torch.zeros((1, self.target_h, self.target_w), dtype=torch.float32),\n                   torch.zeros((1, self.target_h, self.target_w), dtype=torch.float32))\n    \n    def get_class_weights(self) -> Tuple[float, float]:\n        \"\"\"\n        Calculate class weights based on the full dataset\n        \n        Returns:\n            Tuple[float, float]: Weights for background and vessel classes\n        \"\"\"\n        try:\n            total_pixels = 0\n            vessel_pixels = 0\n            \n            for label_file in self.label_files:\n                mask = tifffile.imread(str(label_file))\n                total_pixels += mask.size\n                vessel_pixels += np.sum(mask > 0)\n            \n            background_pixels = total_pixels - vessel_pixels\n            \n            # Calculate weights (inverse frequency)\n            background_weight = 1.0\n            vessel_weight = (background_pixels / vessel_pixels) if vessel_pixels > 0 else 1.0\n            \n            return background_weight, vessel_weight\n            \n        except Exception as e:\n            print(f\"Error calculating class weights: {str(e)}\")\n            return 1.0, 1.0","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T01:07:01.249541Z","iopub.execute_input":"2025-02-27T01:07:01.249856Z","iopub.status.idle":"2025-02-27T01:07:01.267072Z","shell.execute_reply.started":"2025-02-27T01:07:01.249821Z","shell.execute_reply":"2025-02-27T01:07:01.266128Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_train_transforms(dataset_name):\n    \"\"\"Get improved but safer augmentation transforms\"\"\"\n    # Base transforms for all datasets\n    base_transforms = [\n        A.RandomRotate90(p=0.5),\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.ShiftScaleRotate(\n            shift_limit=0.0625,\n            scale_limit=0.1, \n            rotate_limit=30,\n            p=0.5,\n            border_mode=cv2.BORDER_CONSTANT\n        )\n    ]\n    \n    # Dataset-specific transforms (simplified)\n    if 'sparse' in dataset_name:\n        base_transforms.extend([\n            A.GaussNoise(var_limit=(10.0, 30.0), p=0.3),\n            A.RandomBrightnessContrast(p=0.3)\n        ])\n    \n    if 'dense' in dataset_name:\n        base_transforms.extend([\n            A.RandomBrightnessContrast(p=0.3),\n            A.CLAHE(clip_limit=2, p=0.3)\n        ])\n    \n    # Add normalization as final transform\n    base_transforms.append(\n        A.Normalize(mean=[0.485], std=[0.229], max_pixel_value=1.0, p=1.0)\n    )\n    \n    return A.Compose(base_transforms)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T01:07:08.272374Z","iopub.execute_input":"2025-02-27T01:07:08.272683Z","iopub.status.idle":"2025-02-27T01:07:08.278497Z","shell.execute_reply.started":"2025-02-27T01:07:08.272658Z","shell.execute_reply":"2025-02-27T01:07:08.277543Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_val_transforms():\n    return A.Compose([\n        A.Normalize(mean=[0.485], std=[0.229], max_pixel_value=1.0),\n    ])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T01:07:11.814807Z","iopub.execute_input":"2025-02-27T01:07:11.815148Z","iopub.status.idle":"2025-02-27T01:07:11.819454Z","shell.execute_reply.started":"2025-02-27T01:07:11.815121Z","shell.execute_reply":"2025-02-27T01:07:11.818566Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def process_dataset(kaggle_path: Path, dataset_name: str, target_size=(512, 512)) -> Tuple[DataLoader, DataLoader]:\n    \"\"\"Process dataset and create train/val splits\"\"\"\n    print(f\"\\nProcessing {dataset_name}\")\n    \n    try:\n        # Setup paths\n        images_path = kaggle_path / 'train' / dataset_name / 'images'\n        labels_path = kaggle_path / 'train' / dataset_name / 'labels'\n        \n        if not images_path.exists() or not labels_path.exists():\n            print(f\"Required directories not found for {dataset_name}\")\n            return None, None\n        \n        # Get files and match them\n        image_files = sorted(list(images_path.glob('*.tif')))\n        label_files = []\n        \n        # Match files by index\n        for img_file in image_files:\n            img_idx = int(img_file.stem)\n            matching_label = labels_path / f\"{img_idx:04d}.tif\"\n            if matching_label.exists():\n                label_files.append(matching_label)\n        \n        # Keep only images with matching labels\n        image_files = image_files[:len(label_files)]\n        \n        if len(image_files) == 0:\n            print(f\"No matching pairs found for {dataset_name}\")\n            return None, None\n        \n        print(f\"Found {len(image_files)} matching pairs\")\n        \n        # Split into train and validation\n        train_images, val_images, train_labels, val_labels = train_test_split(\n            image_files, label_files,\n            test_size=0.2,\n            random_state=42\n        )\n        \n        # Create datasets with fixed size\n        train_dataset = VesselDataset(\n            train_images,\n            train_labels,\n            transform=get_train_transforms(dataset_name),\n            size=target_size\n        )\n        \n        val_dataset = VesselDataset(\n            val_images,\n            val_labels,\n            transform=get_val_transforms(),\n            size=target_size\n        )\n        \n        # Create dataloaders\n        train_loader = DataLoader(\n            train_dataset,\n            batch_size=4,\n            shuffle=True,\n            num_workers=0,\n            pin_memory=True\n        )\n        \n        val_loader = DataLoader(\n            val_dataset,\n            batch_size=4,\n            shuffle=False,\n            num_workers=0,\n            pin_memory=True\n        )\n        \n        print(f\"Created train dataloader with {len(train_dataset)} samples\")\n        print(f\"Created val dataloader with {len(val_dataset)} samples\")\n        \n        return train_loader, val_loader\n        \n    except Exception as e:\n        print(f\"Error processing {dataset_name}: {str(e)}\")\n        gc.collect()\n        return None, None","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T01:07:26.471884Z","iopub.execute_input":"2025-02-27T01:07:26.472244Z","iopub.status.idle":"2025-02-27T01:07:26.479938Z","shell.execute_reply.started":"2025-02-27T01:07:26.472210Z","shell.execute_reply":"2025-02-27T01:07:26.478986Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def create_overlay(image, mask, alpha=0.5):\n    \"\"\"Create overlay of mask on image\"\"\"\n    # Ensure proper types\n    image = image.astype(np.float32)\n    mask = mask.astype(bool)\n    \n    # Create RGB version of grayscale image\n    rgb_image = np.stack([image] * 3, axis=-1)\n    \n    # Create red mask overlay\n    red_mask = np.zeros_like(rgb_image)\n    red_mask[mask] = [1, 0, 0]  # Red color for mask\n    \n    # Combine image and mask\n    overlay = (1 - alpha) * rgb_image + alpha * red_mask\n    \n    # Ensure values are in valid range\n    overlay = np.clip(overlay, 0, 1)\n    \n    return overlay","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T01:07:30.436450Z","iopub.execute_input":"2025-02-27T01:07:30.436791Z","iopub.status.idle":"2025-02-27T01:07:30.442092Z","shell.execute_reply.started":"2025-02-27T01:07:30.436764Z","shell.execute_reply":"2025-02-27T01:07:30.441112Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_processed_dataset(images_path: Path, labels_path: Path, dataset_name: str, num_samples: int = 3):\n    \"\"\"Visualize samples from dataset including augmentations\"\"\"\n    try:\n        image_files = sorted(list(images_path.glob('*.tif')))\n        label_files = sorted(list(labels_path.glob('*.tif')))\n        \n        if len(image_files) == 0 or len(label_files) == 0:\n            print(f\"No images or masks found for {dataset_name}\")\n            return\n            \n        # Select evenly spaced samples\n        indices = np.linspace(0, len(image_files)-1, num_samples, dtype=int)\n        \n        # Get transforms\n        train_transform = get_train_transforms(dataset_name)\n        \n        for idx in indices:\n            # Load images with memory management\n            with warnings.catch_warnings():\n                warnings.simplefilter(\"ignore\")\n                \n                # Load and preprocess image\n                image = tifffile.imread(str(image_files[idx]))\n                image = (image - image.min()) / (image.max() - image.min())\n                \n                # Load and preprocess mask - keep as uint8\n                mask = tifffile.imread(str(label_files[idx]))\n                mask = (mask > 0).astype(np.uint8)  # Convert to uint8 instead of bool\n                \n                # Apply augmentation\n                try:\n                    augmented = train_transform(image=image.astype(np.float32), \n                                             mask=mask)\n                    aug_image = augmented['image']\n                    aug_mask = augmented['mask']\n                except Exception as e:\n                    print(f\"Augmentation error: {str(e)}\")\n                    continue\n                \n                # Create figure with two rows\n                fig, axes = plt.subplots(2, 3, figsize=(15, 10))\n                fig.suptitle(f'{dataset_name} - Sample {idx}')\n                \n                # Original images row\n                axes[0, 0].imshow(image, cmap='gray')\n                axes[0, 0].set_title('Original')\n                axes[0, 0].axis('off')\n                \n                axes[0, 1].imshow(mask, cmap='Reds')\n                axes[0, 1].set_title('Original Mask')\n                axes[0, 1].axis('off')\n                \n                # Create and show original overlay\n                overlay = create_overlay(image, mask)\n                axes[0, 2].imshow(overlay)\n                axes[0, 2].set_title('Original Overlay')\n                axes[0, 2].axis('off')\n                \n                # Augmented images row\n                axes[1, 0].imshow(aug_image, cmap='gray')\n                axes[1, 0].set_title('Augmented')\n                axes[1, 0].axis('off')\n                \n                axes[1, 1].imshow(aug_mask, cmap='Reds')\n                axes[1, 1].set_title('Augmented Mask')\n                axes[1, 1].axis('off')\n                \n                # Create and show augmented overlay\n                aug_overlay = create_overlay(aug_image, aug_mask.astype(bool))\n                axes[1, 2].imshow(aug_overlay)\n                axes[1, 2].set_title('Augmented Overlay')\n                axes[1, 2].axis('off')\n                \n                plt.tight_layout()\n                plt.show()\n                plt.close()\n                \n            # Clear memory after each sample\n            gc.collect()\n                \n    except Exception as e:\n        print(f\"Error in visualization: {str(e)}\")\n        gc.collect()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T01:07:32.972194Z","iopub.execute_input":"2025-02-27T01:07:32.972532Z","iopub.status.idle":"2025-02-27T01:07:32.982717Z","shell.execute_reply.started":"2025-02-27T01:07:32.972506Z","shell.execute_reply":"2025-02-27T01:07:32.981906Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def save_dataloader_info(dataloaders, save_dir):\n    \"\"\"\n    Save dataloader configurations and splits for each dataset\n    \"\"\"\n    save_dir = Path(save_dir)\n    save_dir.mkdir(parents=True, exist_ok=True)\n    \n    dataset_info = {}\n    \n    for dataset_name, loaders in dataloaders.items():\n        # Get file paths from datasets\n        train_dataset = loaders['train'].dataset\n        val_dataset = loaders['val'].dataset\n        \n        dataset_info[dataset_name] = {\n            'train': {\n                'image_files': [str(path) for path in train_dataset.image_files],\n                'label_files': [str(path) for path in train_dataset.label_files],\n            },\n            'val': {\n                'image_files': [str(path) for path in val_dataset.image_files],\n                'label_files': [str(path) for path in val_dataset.label_files],\n            },\n            'config': {\n                'batch_size': loaders['train'].batch_size,\n                'num_workers': loaders['train'].num_workers,\n                'pin_memory': loaders['train'].pin_memory,\n            }\n        }\n    \n    # Save to JSON file\n    with open(save_dir / 'dataloader_info.json', 'w') as f:\n        json.dump(dataset_info, f, indent=4)\n    \n    print(f\"\\nDataloader information saved to {save_dir / 'dataloader_info.json'}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T01:07:37.108205Z","iopub.execute_input":"2025-02-27T01:07:37.108507Z","iopub.status.idle":"2025-02-27T01:07:37.114860Z","shell.execute_reply.started":"2025-02-27T01:07:37.108484Z","shell.execute_reply":"2025-02-27T01:07:37.113841Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_and_recreate_dataloaders(save_dir):\n    \"\"\"\n    Recreate dataloaders from saved information\n    \"\"\"\n    save_dir = Path(save_dir)\n    \n    # Load saved information\n    with open(save_dir / 'dataloader_info.json', 'r') as f:\n        dataset_info = json.load(f)\n    \n    dataloaders = {}\n    \n    for dataset_name, info in dataset_info.items():\n        # Create train dataset\n        train_dataset = VesselDataset(\n            image_files=[Path(p) for p in info['train']['image_files']],\n            label_files=[Path(p) for p in info['train']['label_files']],\n            transform=get_train_transforms(dataset_name)\n        )\n        \n        # Create val dataset\n        val_dataset = VesselDataset(\n            image_files=[Path(p) for p in info['val']['image_files']],\n            label_files=[Path(p) for p in info['val']['label_files']],\n            transform=get_val_transforms()\n        )\n        \n        # Create dataloaders\n        config = info['config']\n        train_loader = DataLoader(\n            train_dataset,\n            batch_size=config['batch_size'],\n            shuffle=True,\n            num_workers=config['num_workers'],\n            pin_memory=config['pin_memory']\n        )\n        \n        val_loader = DataLoader(\n            val_dataset,\n            batch_size=config['batch_size'],\n            shuffle=False,\n            num_workers=config['num_workers'],\n            pin_memory=config['pin_memory']\n        )\n        \n        dataloaders[dataset_name] = {\n            'train': train_loader,\n            'val': val_loader\n        }\n    \n    print(\"\\nDataloaders recreated successfully!\")\n    return dataloaders\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T01:07:40.505638Z","iopub.execute_input":"2025-02-27T01:07:40.506020Z","iopub.status.idle":"2025-02-27T01:07:40.512757Z","shell.execute_reply.started":"2025-02-27T01:07:40.505969Z","shell.execute_reply":"2025-02-27T01:07:40.511786Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def verify_saved_dataloaders(original_loaders, recreated_loaders):\n    \"\"\"Verify that recreated dataloaders match the original ones\"\"\"\n    print(\"\\nVerifying recreated dataloaders:\")\n    print(\"=\" * 50)\n    \n    for dataset_name in original_loaders.keys():\n        print(f\"\\nDataset: {dataset_name}\")\n        \n        orig_train = original_loaders[dataset_name]['train']\n        new_train = recreated_loaders[dataset_name]['train']\n        \n        orig_val = original_loaders[dataset_name]['val']\n        new_val = recreated_loaders[dataset_name]['val']\n        \n        print(f\"Train samples - Original: {len(orig_train.dataset)}, Recreated: {len(new_train.dataset)}\")\n        print(f\"Val samples - Original: {len(orig_val.dataset)}, Recreated: {len(new_val.dataset)}\")\n        \n        # Verify first batch shapes\n        try:\n            orig_batch = next(iter(orig_train))\n            new_batch = next(iter(new_train))\n            \n            print(\"Batch shapes:\")\n            print(f\"Original - Images: {orig_batch[0].shape}, Masks: {orig_batch[1].shape}\")\n            print(f\"Recreated - Images: {new_batch[0].shape}, Masks: {new_batch[1].shape}\")\n            print(\"Shapes match:\", \n                  orig_batch[0].shape == new_batch[0].shape and \n                  orig_batch[1].shape == new_batch[1].shape)\n            \n        except Exception as e:\n            print(f\"Error checking batch shapes: {str(e)}\")\n            continue","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T01:07:44.476936Z","iopub.execute_input":"2025-02-27T01:07:44.477301Z","iopub.status.idle":"2025-02-27T01:07:44.483122Z","shell.execute_reply.started":"2025-02-27T01:07:44.477273Z","shell.execute_reply":"2025-02-27T01:07:44.482309Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def save_dataloaders_main(dataloaders):\n    \"\"\"Save and verify dataloaders\"\"\"\n    # Save directory\n    save_dir = Path('/kaggle/working/dataloader_info')\n    \n    # Save dataloader information\n    save_dataloader_info(dataloaders, save_dir)\n    \n    # Recreate dataloaders to verify\n    recreated_loaders = load_and_recreate_dataloaders(save_dir)\n    \n    # Verify recreated dataloaders\n    verify_saved_dataloaders(dataloaders, recreated_loaders)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T01:07:50.219454Z","iopub.execute_input":"2025-02-27T01:07:50.219772Z","iopub.status.idle":"2025-02-27T01:07:50.223931Z","shell.execute_reply.started":"2025-02-27T01:07:50.219747Z","shell.execute_reply":"2025-02-27T01:07:50.222881Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def main():\n    print(\"Starting Preprocessing Pipeline\")\n    print(\"=\"*50)\n    \n    # Declare global variables\n    global train_loader, val_loader\n    train_loader, val_loader = None, None\n    \n    # Setup paths\n    kaggle_path = Path('/kaggle/input/blood-vessel-segmentation')\n    save_dir = Path('/kaggle/working/dataloader_info')  \n    \n    # List of datasets\n    datasets = [\n        'kidney_1_dense',\n        'kidney_1_voi',\n        'kidney_2',\n        'kidney_3_dense',\n        'kidney_3_sparse'\n    ]\n    \n    # Process each dataset\n    dataloaders = {}\n    processed_datasets = []\n    \n    for dataset_name in datasets:\n        try:\n            print(f\"\\nProcessing {dataset_name}\")\n            print(\"-\" * 30)\n            \n            # Process dataset with memory management\n            train_loader, val_loader = process_dataset(\n                kaggle_path,\n                dataset_name\n            )\n            \n            if train_loader and val_loader:\n                dataloaders[dataset_name] = {\n                    'train': train_loader,\n                    'val': val_loader\n                }\n                processed_datasets.append(dataset_name)\n                \n                # Clear memory after successful processing\n                gc.collect()\n                \n                # Visualize samples\n                print(f\"\\nDisplaying samples from {dataset_name}\")\n                images_path = kaggle_path / 'train' / dataset_name / 'images'\n                labels_path = kaggle_path / 'train' / dataset_name / 'labels'\n                \n                if images_path.exists() and labels_path.exists():\n                    visualize_processed_dataset(images_path, labels_path, dataset_name)\n                    print(f\"Visualization completed for {dataset_name}\")\n                else:\n                    print(f\"Could not find image/label directories for {dataset_name}\")\n            else:\n                print(f\"Failed to create dataloaders for {dataset_name}\")\n                    \n        except Exception as e:\n            print(f\"Error processing {dataset_name}: {str(e)}\")\n            gc.collect()\n            continue\n    \n    print(\"\\nPreprocessing Summary:\")\n    print(\"=\" * 50)\n    print(f\"Successfully processed datasets: {len(processed_datasets)}/{len(datasets)}\")\n    for dataset in processed_datasets:\n        print(f\"- {dataset}\")\n    \n    # Save dataloader information\n    print(\"\\nSaving dataloader information...\")\n    save_dataloader_info(dataloaders, save_dir)\n    \n    # Load and verify saved dataloaders\n    print(\"\\nVerifying saved dataloaders...\")\n    recreated_loaders = load_and_recreate_dataloaders(save_dir)\n    verify_saved_dataloaders(dataloaders, recreated_loaders)\n    \n    print(\"\\nPreprocessing and dataloader creation completed!\")\n    return dataloaders\n\nif __name__ == \"__main__\":\n    dataloaders = main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T01:07:53.279413Z","iopub.execute_input":"2025-02-27T01:07:53.279764Z","iopub.status.idle":"2025-02-27T01:09:09.425219Z","shell.execute_reply.started":"2025-02-27T01:07:53.279738Z","shell.execute_reply":"2025-02-27T01:09:09.424209Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model Architecture","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torchvision.models import resnet34\nfrom torchvision.models import resnet101","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T01:12:47.688912Z","iopub.execute_input":"2025-02-27T01:12:47.689330Z","iopub.status.idle":"2025-02-27T01:12:51.799203Z","shell.execute_reply.started":"2025-02-27T01:12:47.689300Z","shell.execute_reply":"2025-02-27T01:12:51.798527Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class AttentionBlock(nn.Module):\n    def __init__(self, in_channels):\n        super().__init__()\n        self.conv = nn.Conv2d(in_channels, in_channels, kernel_size=1)\n        self.activation = nn.Sigmoid()\n        \n    def forward(self, x):\n        # Avoid operations that could collapse spatial dimensions\n        attention = self.conv(x)\n        attention = self.activation(attention)\n        return x * attention\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T01:12:53.535101Z","iopub.execute_input":"2025-02-27T01:12:53.535700Z","iopub.status.idle":"2025-02-27T01:12:53.540614Z","shell.execute_reply.started":"2025-02-27T01:12:53.535670Z","shell.execute_reply":"2025-02-27T01:12:53.539726Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ConvBlock(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super().__init__()\n        self.conv1 = nn.Sequential(\n            nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True)\n        )\n        self.conv2 = nn.Sequential(\n            nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True)\n        )\n        # Add residual connection\n        self.identity = nn.Conv2d(in_channels, out_channels, kernel_size=1) if in_channels != out_channels else nn.Identity()\n        \n    def forward(self, x):\n        identity = self.identity(x)\n        out = self.conv1(x)\n        out = self.conv2(out)\n        return F.relu(out + identity)  # Residual connection","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T01:12:55.967053Z","iopub.execute_input":"2025-02-27T01:12:55.967389Z","iopub.status.idle":"2025-02-27T01:12:55.973171Z","shell.execute_reply.started":"2025-02-27T01:12:55.967360Z","shell.execute_reply":"2025-02-27T01:12:55.972120Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ResNetUNet(nn.Module):\n    def __init__(self, n_classes=1):\n        super().__init__()\n        \n        # Load pretrained ResNet34\n        resnet = resnet34(pretrained=True)\n        \n        # Modify first layer to accept single channel\n        self.firstconv = nn.Conv2d(1, 64, kernel_size=7, stride=2, padding=3, bias=False)\n        \n        # Initialize first layer with pretrained weights\n        with torch.no_grad():\n            self.firstconv.weight[:, 0:1, :, :] = torch.sum(resnet.conv1.weight, dim=1, keepdim=True)\n        \n        # Encoder (ResNet layers)\n        self.encoder1 = nn.Sequential(\n            self.firstconv,\n            resnet.bn1,\n            resnet.relu\n        )\n        self.pool = resnet.maxpool\n        self.encoder2 = resnet.layer1  # 64 channels\n        self.encoder3 = resnet.layer2  # 128 channels\n        self.encoder4 = resnet.layer3  # 256 channels\n        self.encoder5 = resnet.layer4  # 512 channels\n        \n        # Simple attention at deepest level\n        self.attention = AttentionBlock(512)\n        \n        # Decoder\n        self.decoder5 = ConvBlock(512, 512)\n        self.decoder4 = ConvBlock(512 + 256, 256)\n        self.decoder3 = ConvBlock(256 + 128, 128)\n        self.decoder2 = ConvBlock(128 + 64, 64)\n        self.decoder1 = ConvBlock(64 + 64, 32)\n        \n        # Final layers\n        self.upsample = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)\n        self.final_conv = nn.Conv2d(32, n_classes, kernel_size=1)\n        \n        self.dropout = nn.Dropout2d(0.25)\n        \n        # Initialize weights of decoder\n        self._initialize_weights()\n    \n    def _initialize_weights(self):\n        for m in [self.decoder1, self.decoder2, self.decoder3, self.decoder4, self.decoder5, self.final_conv]:\n            if isinstance(m, nn.Conv2d):\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            elif isinstance(m, nn.BatchNorm2d):\n                nn.init.constant_(m.weight, 1)\n                nn.init.constant_(m.bias, 0)\n    \n    def forward(self, x):\n        # Encoder\n        e1 = self.encoder1(x)\n        e1_pool = self.pool(e1)\n        e2 = self.encoder2(e1_pool)\n        e3 = self.encoder3(e2)\n        e4 = self.encoder4(e3)\n        e5 = self.encoder5(e4)\n        \n        # Apply simple attention only at the deepest level\n        e5 = self.attention(e5)\n        \n        # Apply dropout for regularization\n        e5 = self.dropout(e5)\n        \n        # Decoder with skip connections\n        d5 = self.decoder5(e5)\n        d4 = self.decoder4(torch.cat([self.upsample(d5), e4], dim=1))\n        d3 = self.decoder3(torch.cat([self.upsample(d4), e3], dim=1))\n        d2 = self.decoder2(torch.cat([self.upsample(d3), e2], dim=1))\n        d1 = self.decoder1(torch.cat([self.upsample(d2), e1], dim=1))\n        \n        # Final output\n        final_out = self.final_conv(self.upsample(d1))\n        \n        return torch.sigmoid(final_out)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T01:12:58.893498Z","iopub.execute_input":"2025-02-27T01:12:58.893859Z","iopub.status.idle":"2025-02-27T01:12:58.907232Z","shell.execute_reply.started":"2025-02-27T01:12:58.893827Z","shell.execute_reply":"2025-02-27T01:12:58.906252Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class DiceLoss(nn.Module):\n    def __init__(self, smooth=1.0):\n        super().__init__()\n        self.smooth = smooth\n        \n    def forward(self, pred, target):\n        pred_flat = pred.view(-1)\n        target_flat = target.view(-1)\n        \n        intersection = (pred_flat * target_flat).sum()\n        union = pred_flat.sum() + target_flat.sum()\n        \n        dice = (2. * intersection + self.smooth) / (union + self.smooth)\n        return 1 - dice","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T01:13:03.426043Z","iopub.execute_input":"2025-02-27T01:13:03.426382Z","iopub.status.idle":"2025-02-27T01:13:03.431470Z","shell.execute_reply.started":"2025-02-27T01:13:03.426354Z","shell.execute_reply":"2025-02-27T01:13:03.430621Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class WeightedBCELoss(nn.Module):\n    def __init__(self, pos_weight=None):\n        super().__init__()\n        self.pos_weight = pos_weight\n        \n    def forward(self, pred, target):\n        if self.pos_weight is None:\n            # Calculate weights based on inverse class frequency\n            neg_count = (target == 0).float().sum()\n            pos_count = (target == 1).float().sum()\n            total = neg_count + pos_count\n            self.pos_weight = (neg_count / total) / (pos_count / total)\n        \n        return F.binary_cross_entropy_with_logits(\n            pred, target, \n            pos_weight=self.pos_weight * torch.ones_like(target)\n        )\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T01:13:05.753601Z","iopub.execute_input":"2025-02-27T01:13:05.753947Z","iopub.status.idle":"2025-02-27T01:13:05.758991Z","shell.execute_reply.started":"2025-02-27T01:13:05.753917Z","shell.execute_reply":"2025-02-27T01:13:05.758102Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CombinedLoss(nn.Module):\n    def __init__(self, dice_weight=0.5, bce_weight=0.5):\n        super().__init__()\n        self.dice_weight = dice_weight\n        self.bce_weight = bce_weight\n        self.dice_loss = DiceLoss()\n        self.weighted_bce = WeightedBCELoss()\n        \n    def forward(self, pred, target):\n        dice = self.dice_loss(pred, target)\n        weighted_bce = self.weighted_bce(pred, target)\n        \n        # Focal Loss component\n        pt = torch.exp(-weighted_bce)\n        focal = (1 - pt) ** 2 * weighted_bce\n        \n        return (self.dice_weight * dice + \n                self.bce_weight * weighted_bce \n                )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T01:13:24.861352Z","iopub.execute_input":"2025-02-27T01:13:24.861686Z","iopub.status.idle":"2025-02-27T01:13:24.867358Z","shell.execute_reply.started":"2025-02-27T01:13:24.861661Z","shell.execute_reply":"2025-02-27T01:13:24.866295Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def initialize_model():\n    \"\"\"Initialize model, loss function, and print architecture summary\"\"\"\n    print(\"Initializing Model Architecture\")\n    print(\"=\" * 50)\n    \n    try:\n        # Initialize model\n        model = ResNetUNet(n_classes=1)\n        \n        # Create loss function\n        criterion = CombinedLoss(dice_weight=0.7, bce_weight=0.3)\n        \n        # Move to GPU if available\n        device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n        print(f\"\\nUsing device: {device}\")\n        \n        model = model.to(device)\n        \n        # Print model summary\n        print(\"\\nModel Architecture:\")\n        print(\"-\" * 30)\n        \n        # Calculate total parameters\n        total_params = sum(p.numel() for p in model.parameters())\n        trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\n        \n        print(f\"\\nTotal Parameters: {total_params:,}\")\n        print(f\"Trainable Parameters: {trainable_params:,}\")\n        \n        # Test forward pass\n        print(\"\\nTesting forward pass...\")\n        model.eval()  # Set to evaluation mode\n        test_input = torch.randn(1, 1, 512, 512).to(device)\n        with torch.no_grad():\n            test_output = model(test_input)\n        print(f\"Input shape: {test_input.shape}\")\n        print(f\"Output shape: {test_output.shape}\")\n        \n        print(\"\\nModel architecture initialization completed successfully!\")\n        return model, criterion, device\n        \n    except Exception as e:\n        print(f\"Error in model initialization: {str(e)}\")\n        return None, None, None\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T01:13:29.504705Z","iopub.execute_input":"2025-02-27T01:13:29.505102Z","iopub.status.idle":"2025-02-27T01:13:29.511604Z","shell.execute_reply.started":"2025-02-27T01:13:29.505071Z","shell.execute_reply":"2025-02-27T01:13:29.510766Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def main():\n    print(\"Starting Model Architecture Setup\")\n    print(\"=\" * 50)\n    \n    try:\n        # Initialize model, criterion, and device\n        model, criterion, device = initialize_model()\n        \n        if model is None:\n            print(\"Model initialization failed!\")\n            return None, None, None\n        \n        # Save model architecture (optional)\n        try:\n            torch.save(model.state_dict(), '/kaggle/working/initial_model.pth')\n            print(\"\\nSaved initial model state\")\n        except Exception as e:\n            print(f\"Error saving model: {str(e)}\")\n        \n        print(\"\\nModel Architecture Setup Completed!\")\n        return model, criterion, device\n        \n    except Exception as e:\n        print(f\"Error in main: {str(e)}\")\n        return None, None, None\n\nif __name__ == \"__main__\":\n    model, criterion, device = main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T01:13:33.628977Z","iopub.execute_input":"2025-02-27T01:13:33.629346Z","iopub.status.idle":"2025-02-27T01:13:36.031669Z","shell.execute_reply.started":"2025-02-27T01:13:33.629319Z","shell.execute_reply":"2025-02-27T01:13:36.030663Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training Pipeline","metadata":{}},{"cell_type":"code","source":"import torch\nfrom torch.optim import Adam\nimport numpy as np\nfrom tqdm import tqdm\nimport time\nimport json\nfrom pathlib import Path\nimport matplotlib.pyplot as plt\nfrom datetime import datetime\nfrom scipy.ndimage import distance_transform_edt, binary_erosion\nfrom copy import deepcopy\nimport gc\nimport traceback","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T01:13:42.526906Z","iopub.execute_input":"2025-02-27T01:13:42.527241Z","iopub.status.idle":"2025-02-27T01:13:42.531913Z","shell.execute_reply.started":"2025-02-27T01:13:42.527216Z","shell.execute_reply":"2025-02-27T01:13:42.530895Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class AverageMeter:\n    \"\"\"Computes and stores the average and current value\"\"\"\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.val = 0\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\n\n    def update(self, val, n=1):\n        self.val = val\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T01:13:45.573803Z","iopub.execute_input":"2025-02-27T01:13:45.574178Z","iopub.status.idle":"2025-02-27T01:13:45.579122Z","shell.execute_reply.started":"2025-02-27T01:13:45.574149Z","shell.execute_reply":"2025-02-27T01:13:45.578259Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def calculate_surface_dice(pred, target, tolerance=1):\n    \"\"\"Calculate Surface Dice score with memory-efficient operations\"\"\"\n    from scipy.ndimage import distance_transform_edt, binary_erosion\n    import numpy as np\n    \n    # Ensure inputs are boolean\n    pred = pred.astype(bool)\n    target = target.astype(bool)\n    \n    # Get surface points using XOR operation\n    pred_eroded = binary_erosion(pred)\n    target_eroded = binary_erosion(target)\n    \n    pred_surface = np.logical_xor(pred, pred_eroded)\n    target_surface = np.logical_xor(target, target_eroded)\n    \n    # Calculate distance maps\n    pred_distance = distance_transform_edt(~pred_surface)\n    target_distance = distance_transform_edt(~target_surface)\n    \n    # Get surface points within tolerance\n    pred_tolerant = pred_surface & (target_distance <= tolerance)\n    target_tolerant = target_surface & (pred_distance <= tolerance)\n    \n    # Calculate Surface Dice\n    surface_dice = (2.0 * pred_tolerant.sum() + 1e-7) / (pred_surface.sum() + target_surface.sum() + 1e-7)\n    \n    return surface_dice.item()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T01:13:48.317283Z","iopub.execute_input":"2025-02-27T01:13:48.317608Z","iopub.status.idle":"2025-02-27T01:13:48.323672Z","shell.execute_reply.started":"2025-02-27T01:13:48.317580Z","shell.execute_reply":"2025-02-27T01:13:48.322672Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def calculate_metrics(pred, target):\n    \"\"\"Calculate Dice and Surface Dice metrics\"\"\"\n    pred = pred.float()\n    target = target.float()\n    \n    # Move tensors to CPU and convert to numpy for Surface Dice\n    pred_np = pred.detach().cpu().numpy()\n    target_np = target.detach().cpu().numpy()\n    \n    # Regular Dice\n    intersection = (pred * target).sum()\n    union = pred.sum() + target.sum()\n    dice = (2. * intersection + 1e-7) / (union + 1e-7)\n    \n    # Surface Dice (calculate for each sample in batch)\n    surface_dice_scores = []\n    for p, t in zip(pred_np, target_np):\n        surface_dice_scores.append(calculate_surface_dice(p.squeeze(), t.squeeze()))\n    avg_surface_dice = np.mean(surface_dice_scores)\n    \n    # Clear memory\n    del pred_np, target_np\n    gc.collect()\n    \n    return {\n        'dice': dice.item(),\n        'surface_dice': avg_surface_dice\n    }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T01:13:51.022467Z","iopub.execute_input":"2025-02-27T01:13:51.022776Z","iopub.status.idle":"2025-02-27T01:13:51.028442Z","shell.execute_reply.started":"2025-02-27T01:13:51.022752Z","shell.execute_reply":"2025-02-27T01:13:51.027605Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def mixup_batch(images, masks, alpha=0.2):\n    lam = np.random.beta(alpha, alpha)\n    batch_size = images.size(0)\n    index = torch.randperm(batch_size)\n    \n    mixed_images = lam * images + (1 - lam) * images[index]\n    mixed_masks = lam * masks + (1 - lam) * masks[index]\n    \n    return mixed_images, mixed_masks","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T01:13:54.082669Z","iopub.execute_input":"2025-02-27T01:13:54.082981Z","iopub.status.idle":"2025-02-27T01:13:54.087695Z","shell.execute_reply.started":"2025-02-27T01:13:54.082956Z","shell.execute_reply":"2025-02-27T01:13:54.086847Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Trainer:\n    def __init__(self, model, criterion, optimizer, device, save_dir,\n                save_every=1, patience=10):\n        self.model = model\n        self.criterion = criterion\n        self.optimizer = optimizer\n        self.device = device\n        self.save_dir = Path(save_dir)\n        self.save_every = save_every\n        self.patience = patience\n        \n        # Initialize best metrics\n        self.best_dice = -1\n        self.best_surface_dice = -1\n        self.patience_counter = 0\n        \n        # Initialize history\n        self.history = {\n            'train_loss': [], 'train_dice': [], 'train_surface_dice': [],\n            'val_loss': [], 'val_dice': [], 'val_surface_dice': [],\n            'lr': []\n        }\n        \n        # Create save directory if it doesn't exist\n        self.save_dir.mkdir(parents=True, exist_ok=True)\n        \n        # Setup scaler for mixed precision\n        self.scaler = torch.cuda.amp.GradScaler()\n    \n    def train_epoch(self, train_loader):\n        self.model.train()\n        losses = AverageMeter()\n        dice_scores = AverageMeter()\n        surface_dice_scores = AverageMeter()\n        \n        pbar = tqdm(train_loader, desc=f'Training')\n        \n        for batch_idx, (images, masks) in enumerate(pbar):\n            try:\n                # Move to device\n                images = images.to(self.device)\n                masks = masks.to(self.device)\n                \n                # Forward pass with mixed precision\n                self.optimizer.zero_grad()\n                \n                with torch.cuda.amp.autocast():\n                    outputs = self.model(images)\n                    loss = self.criterion(outputs, masks)\n                \n                # Backward pass with gradient scaling\n                self.scaler.scale(loss).backward()\n                self.scaler.step(self.optimizer)\n                self.scaler.update()\n                \n                # Calculate metrics\n                with torch.no_grad():\n                    metrics = calculate_metrics((outputs > 0.5).float(), masks)\n                \n                # Update meters\n                losses.update(loss.item(), images.size(0))\n                dice_scores.update(metrics['dice'], images.size(0))\n                surface_dice_scores.update(metrics['surface_dice'], images.size(0))\n                \n                # Update progress bar\n                pbar.set_postfix({\n                    'Loss': f'{losses.avg:.4f}',\n                    'Dice': f'{dice_scores.avg:.4f}',\n                    'SurfDice': f'{surface_dice_scores.avg:.4f}',\n                    'LR': f\"{self.optimizer.param_groups[0]['lr']:.6f}\"\n                })\n                \n                # Clear memory\n                del images, masks, outputs, loss\n                if batch_idx % 10 == 0:\n                    gc.collect()\n                    if torch.cuda.is_available():\n                        torch.cuda.empty_cache()\n            \n            except Exception as e:\n                print(f\"Error in training batch: {str(e)}\")\n                continue\n        \n        return {\n            'loss': losses.avg,\n            'dice': dice_scores.avg,\n            'surface_dice': surface_dice_scores.avg\n        }\n    \n    def validate(self, val_loader):\n        self.model.eval()\n        \n        losses = AverageMeter()\n        dice_scores = AverageMeter()\n        surface_dice_scores = AverageMeter()\n        \n        with torch.no_grad():\n            pbar = tqdm(val_loader, desc='Validating')\n            for images, masks in pbar:\n                try:\n                    # Move to device\n                    images = images.to(self.device)\n                    masks = masks.to(self.device)\n                    \n                    # Forward pass\n                    outputs = self.model(images)\n                    loss = self.criterion(outputs, masks)\n                    \n                    # Calculate metrics\n                    pred_binary = (outputs > 0.5).float()\n                    metrics = calculate_metrics(pred_binary, masks)\n                    \n                    # Update meters\n                    losses.update(loss.item(), images.size(0))\n                    dice_scores.update(metrics['dice'], images.size(0))\n                    surface_dice_scores.update(metrics['surface_dice'], images.size(0))\n                    \n                    # Update progress bar\n                    pbar.set_postfix({\n                        'Loss': f'{losses.avg:.4f}',\n                        'Dice': f'{dice_scores.avg:.4f}',\n                        'SurfDice': f'{surface_dice_scores.avg:.4f}'\n                    })\n                    \n                    # Clear memory\n                    del images, masks, outputs, pred_binary\n                    gc.collect()\n                    if torch.cuda.is_available():\n                        torch.cuda.empty_cache()\n                \n                except Exception as e:\n                    print(f\"Error in validation batch: {str(e)}\")\n                    continue\n        \n        return {\n            'loss': losses.avg,\n            'dice': dice_scores.avg,\n            'surface_dice': surface_dice_scores.avg\n        }\n    \n    def train(self, train_loader, val_loader, num_epochs, scheduler=None, early_stopping=True):\n        \"\"\"Train the model with safer approach\"\"\"\n        print(f\"Starting safer improved training with {num_epochs} epochs\")\n        \n        for epoch in range(1, num_epochs + 1):\n            print(f'\\nEpoch {epoch}/{num_epochs}')\n            print('-' * 20)\n            \n            # Train\n            train_metrics = self.train_epoch(train_loader)\n            \n            # Validate\n            val_metrics = self.validate(val_loader)\n            \n            # Update scheduler with validation metrics for ReduceLROnPlateau\n            if scheduler is not None and isinstance(scheduler, torch.optim.lr_scheduler.ReduceLROnPlateau):\n                # Use dice score as the metric to monitor\n                scheduler.step(val_metrics['dice'])\n            elif scheduler is not None:\n                # For other schedulers that don't need metrics\n                scheduler.step()\n            \n            # Update history\n            self.history['train_loss'].append(train_metrics['loss'])\n            self.history['train_dice'].append(train_metrics['dice'])\n            self.history['train_surface_dice'].append(train_metrics['surface_dice'])\n            self.history['val_loss'].append(val_metrics['loss'])\n            self.history['val_dice'].append(val_metrics['dice'])\n            self.history['val_surface_dice'].append(val_metrics['surface_dice'])\n            self.history['lr'].append(self.optimizer.param_groups[0]['lr'])\n            \n            # Print epoch summary\n            print(f\"\\nEpoch {epoch} Summary:\")\n            print(f\"Train - Loss: {train_metrics['loss']:.4f}, \"\n                 f\"Dice: {train_metrics['dice']:.4f}, \"\n                 f\"Surface Dice: {train_metrics['surface_dice']:.4f}\")\n            print(f\"Val - Loss: {val_metrics['loss']:.4f}, \"\n                 f\"Dice: {val_metrics['dice']:.4f}, \"\n                 f\"Surface Dice: {val_metrics['surface_dice']:.4f}\")\n            \n            # Check for improvement\n            improved = False\n            if val_metrics['dice'] > self.best_dice:\n                self.best_dice = val_metrics['dice']\n                improved = True\n                print(f\"New best Dice: {self.best_dice:.4f}\")\n            \n            if val_metrics['surface_dice'] > self.best_surface_dice:\n                self.best_surface_dice = val_metrics['surface_dice']\n                improved = True\n                print(f\"New best Surface Dice: {self.best_surface_dice:.4f}\")\n            \n            # Save checkpoint if improved\n            if improved:\n                self.patience_counter = 0\n                self.save_checkpoint(epoch, val_metrics, is_best=True)\n                print(\"New best model saved!\")\n            else:\n                self.patience_counter += 1\n                print(f\"No improvement for {self.patience_counter} epochs\")\n                if early_stopping and self.patience_counter >= self.patience:\n                    print(f\"\\nEarly stopping triggered after {self.patience} epochs without improvement\")\n                    break\n            \n            # Plot current progress\n            if epoch % 2 == 0 or epoch == num_epochs:\n                self.plot_training_history()\n            \n            # Clear memory\n            gc.collect()\n            if torch.cuda.is_available():\n                torch.cuda.empty_cache()\n        \n        print(\"\\nTraining completed!\")\n        print(f\"Best Dice Score: {self.best_dice:.4f}\")\n        print(f\"Best Surface Dice Score: {self.best_surface_dice:.4f}\")\n        \n        return self.best_dice, self.best_surface_dice\n           \n    \n    def save_checkpoint(self, epoch, metrics, is_best=False):\n        \"\"\"Save training checkpoint\"\"\"\n        checkpoint = {\n            'epoch': epoch,\n            'model_state_dict': self.model.state_dict(),\n            'optimizer_state_dict': self.optimizer.state_dict(),\n            'metrics': metrics,\n            'best_dice': self.best_dice,\n            'best_surface_dice': self.best_surface_dice,\n            'history': self.history\n        }\n        \n        # Save periodic checkpoint\n        if epoch % self.save_every == 0:\n            checkpoint_path = self.save_dir / f'checkpoint_epoch_{epoch}.pth'\n            torch.save(checkpoint, checkpoint_path)\n            print(f\"Saved checkpoint to {checkpoint_path}\")\n        \n        # Save best model\n        if is_best:\n            best_model_path = self.save_dir / 'best_model.pth'\n            torch.save(checkpoint, best_model_path)\n            print(f\"Saved best model to {best_model_path}\")\n        \n        # Clear memory\n        del checkpoint\n        gc.collect()\n        if torch.cuda.is_available():\n            torch.cuda.empty_cache()\n    \n\n    def load_checkpoint(self, checkpoint_path):\n        \"\"\"Load training checkpoint\"\"\"\n        try:\n            checkpoint = torch.load(checkpoint_path, map_location=self.device)\n            self.model.load_state_dict(checkpoint['model_state_dict'])\n            self.optimizer.load_state_dict(checkpoint['optimizer_state_dict'])\n            self.best_dice = checkpoint['best_dice']\n            self.best_surface_dice = checkpoint['best_surface_dice']\n            self.history = checkpoint['history']\n            start_epoch = checkpoint['epoch']\n            \n            print(f\"Loaded checkpoint from epoch {start_epoch}\")\n            return start_epoch\n        \n        except Exception as e:\n            print(f\"Error loading checkpoint: {str(e)}\")\n            return 0\n    \n    def save_history(self):\n        \"\"\"Save training history to JSON\"\"\"\n        history_path = self.save_dir / 'training_history.json'\n        with open(history_path, 'w') as f:\n            json.dump(self.history, f)\n        print(f\"Saved training history to {history_path}\")\n        \n    def plot_training_history(self):\n        \"\"\"Display training history plots without saving\"\"\"\n        if len(self.history['train_loss']) < 2:\n            print(\"Not enough data points to plot (need at least 2 epochs)\")\n            return\n        \n        plt.figure(figsize=(15, 10))\n        \n        # Plot loss\n        plt.subplot(2, 1, 1)\n        plt.plot(self.history['train_loss'], 'b-', label='Train', marker='o')\n        plt.plot(self.history['val_loss'], 'r-', label='Validation', marker='s')\n        plt.title('Loss History')\n        plt.xlabel('Epoch')\n        plt.ylabel('Loss')\n        plt.grid(True)\n        plt.legend()\n        \n        # Plot metrics\n        plt.subplot(2, 1, 2)\n        plt.plot(self.history['train_dice'], 'b-', label='Train Dice', marker='o')\n        plt.plot(self.history['val_dice'], 'b--', label='Val Dice', marker='s')\n        plt.plot(self.history['train_surface_dice'], 'r-', label='Train Surface Dice', marker='^')\n        plt.plot(self.history['val_surface_dice'], 'r--', label='Val Surface Dice', marker='v')\n        plt.title('Metrics History')\n        plt.xlabel('Epoch')\n        plt.ylabel('Score')\n        plt.grid(True)\n        plt.legend()\n        \n        plt.tight_layout()\n        plt.show()\n        plt.close()\n        \n        # Clear memory\n        gc.collect()\n        if torch.cuda.is_available():\n            torch.cuda.empty_cache()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T01:34:39.286895Z","iopub.execute_input":"2025-02-27T01:34:39.287225Z","iopub.status.idle":"2025-02-27T01:34:39.312779Z","shell.execute_reply.started":"2025-02-27T01:34:39.287201Z","shell.execute_reply":"2025-02-27T01:34:39.311947Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_and_recreate_dataloaders(save_dir):\n    \"\"\"Recreate dataloaders from saved information\"\"\"\n    save_dir = Path(save_dir)\n    \n    # Load saved information\n    with open(save_dir / 'dataloader_info.json', 'r') as f:\n        dataset_info = json.load(f)\n    \n    dataloaders = {}\n    \n    for dataset_name, info in dataset_info.items():\n        # Create train dataset\n        train_dataset = VesselDataset(\n            image_files=[Path(p) for p in info['train']['image_files']],\n            label_files=[Path(p) for p in info['train']['label_files']],\n            transform=get_train_transforms(dataset_name)\n        )\n        \n        # Create val dataset\n        val_dataset = VesselDataset(\n            image_files=[Path(p) for p in info['val']['image_files']],\n            label_files=[Path(p) for p in info['val']['label_files']],\n            transform=get_val_transforms()\n        )\n        \n        # Create dataloaders\n        config = info['config']\n        train_loader = DataLoader(\n            train_dataset,\n            batch_size=config['batch_size'],\n            shuffle=True,\n            num_workers=config['num_workers'],\n            pin_memory=config['pin_memory']\n        )\n        \n        val_loader = DataLoader(\n            val_dataset,\n            batch_size=config['batch_size'],\n            shuffle=False,\n            num_workers=config['num_workers'],\n            pin_memory=config['pin_memory']\n        )\n        \n        dataloaders[dataset_name] = {\n            'train': train_loader,\n            'val': val_loader\n        }\n    \n    return dataloaders\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T01:34:43.259764Z","iopub.execute_input":"2025-02-27T01:34:43.260118Z","iopub.status.idle":"2025-02-27T01:34:43.266478Z","shell.execute_reply.started":"2025-02-27T01:34:43.260090Z","shell.execute_reply":"2025-02-27T01:34:43.265763Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def create_subset_loader(loader, subset_size=50):\n    \"\"\"Create a smaller dataloader for testing\"\"\"\n    subset_dataset = torch.utils.data.Subset(\n        loader.dataset, \n        indices=range(min(subset_size, len(loader.dataset)))\n    )\n    return DataLoader(\n        subset_dataset,\n        batch_size=loader.batch_size,\n        shuffle=True,\n        num_workers=loader.num_workers,\n        pin_memory=loader.pin_memory\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T01:34:47.104243Z","iopub.execute_input":"2025-02-27T01:34:47.104608Z","iopub.status.idle":"2025-02-27T01:34:47.108800Z","shell.execute_reply.started":"2025-02-27T01:34:47.104584Z","shell.execute_reply":"2025-02-27T01:34:47.108098Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def main(test_mode=False, subset_size=50, num_epochs=15):\n    \"\"\"\n    Safer improved training pipeline with minimal changes to your original model\n    \n    Args:\n        test_mode (bool): If True, use subset of data for testing\n        subset_size (int): Number of samples to use in test mode\n        num_epochs (int): Number of epochs to train\n    \"\"\"\n    print(\"Starting Safer Improved Training Pipeline\")\n    print(\"=\" * 50)\n    \n    try:\n        # Clear initial memory\n        gc.collect()\n        if torch.cuda.is_available():\n            torch.cuda.empty_cache()\n        \n        # Load dataloaders\n        print(\"\\nLoading dataloaders...\")\n        dataloader_path = Path('/kaggle/working/dataloader_info')\n        dataloaders = load_and_recreate_dataloaders(dataloader_path)\n        print(f\"Successfully loaded dataloaders for {len(dataloaders)} datasets\")\n        \n        if test_mode:\n            print(f\"\\nTest Mode: Creating subset of {subset_size} samples\")\n            subset_loaders = {}\n            for dataset_name, loaders in dataloaders.items():\n                subset_loaders[dataset_name] = {\n                    'train': create_subset_loader(loaders['train'], subset_size),\n                    'val': create_subset_loader(loaders['val'], subset_size//5)\n                }\n            dataloaders = subset_loaders\n            print(\"Created subset dataloaders for testing\")\n            gc.collect()\n        \n        # Load improved model\n        print(\"\\nLoading model with minimal improvements...\")\n        model = ResNetUNet(n_classes=1)\n        \n        # Setup device\n        device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n        print(f\"\\nUsing device: {device}\")\n        model = model.to(device)\n        \n        # Setup save directory\n        timestamp = datetime.now().strftime('%Y%m%d_%H%M%S')\n        mode_prefix = 'test' if test_mode else 'safer_improved'\n        base_save_dir = Path(f'/kaggle/working/{mode_prefix}_training_results_{timestamp}')\n        base_save_dir.mkdir(parents=True, exist_ok=True)\n        print(f\"\\nCreated base save directory: {base_save_dir}\")\n        \n        # Print training configuration\n        print(\"\\nTraining Configuration:\")\n        print(f\"Mode: {'Test' if test_mode else 'Safer Improved'}\")\n        print(f\"Number of epochs: {num_epochs}\")\n        print(f\"Device: {device}\")\n        if test_mode:\n            print(f\"Subset size: {subset_size}\")\n        \n        # Initialize results dictionary to store best scores\n        results = {}\n        \n        # Start training for each dataset\n        for dataset_name, loaders in dataloaders.items():\n            print(f\"\\nTraining on dataset: {dataset_name}\")\n            print(f\"Train samples: {len(loaders['train'].dataset)}\")\n            print(f\"Val samples: {len(loaders['val'].dataset)}\")\n            print(\"-\" * 30)\n            \n            # Get dataset-specific weights\n            weights = class_weights.get(dataset_name, {'dice_weight': 0.7, 'bce_weight': 0.3})\n            \n            criterion = CombinedLoss(\n                dice_weight=weights['dice_weight'], \n                bce_weight=weights['bce_weight']\n            )\n            \n            # Setup optimizer - keeping original learning rate\n            optimizer = torch.optim.Adam(\n                model.parameters(), \n                lr=1e-4,\n                weight_decay=1e-5\n            )\n            \n            # Setup learning rate scheduler - simple ReduceLROnPlateau\n            scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n                optimizer,\n                mode='max',  \n                factor=0.5,  \n                patience=3,  \n                verbose=True,\n                min_lr=1e-6\n            )\n            \n            # Create dataset-specific save directory\n            dataset_save_dir = base_save_dir / dataset_name\n            dataset_save_dir.mkdir(parents=True, exist_ok=True)\n            print(f\"Save directory for {dataset_name}: {dataset_save_dir}\")\n            \n            # Create trainer\n            trainer = Trainer(\n                model=model,\n                criterion=criterion,\n                optimizer=optimizer,\n                device=device,\n                save_dir=dataset_save_dir,\n                save_every=1 if test_mode else 3,\n                patience=3 if test_mode else 8\n            )\n            \n            # Train\n            best_dice, best_surface_dice = trainer.train(\n                train_loader=loaders['train'],\n                val_loader=loaders['val'],\n                num_epochs=num_epochs,\n                scheduler=scheduler,\n                early_stopping=True\n            )\n            \n            # Store results\n            results[dataset_name] = {\n                'best_dice': best_dice,\n                'best_surface_dice': best_surface_dice\n            }\n            \n            print(f\"\\nCompleted training for {dataset_name}\")\n            print(f\"Best Dice Score: {best_dice:.4f}\")\n            print(f\"Best Surface Dice Score: {best_surface_dice:.4f}\")\n            \n            # Save history\n            trainer.save_history()\n            \n            # Clear memory\n            gc.collect()\n            if torch.cuda.is_available():\n                torch.cuda.empty_cache()\n            \n            # Clear some variables to free memory\n            del loaders, trainer\n            gc.collect()\n        \n        print(\"\\nTraining Pipeline Completed!\")\n        print(\"\\nFinal Results Summary:\")\n        print(\"=\" * 50)\n        for dataset_name, metrics in results.items():\n            print(f\"\\n{dataset_name}:\")\n            print(f\"Best Dice Score: {metrics['best_dice']:.4f}\")\n            print(f\"Best Surface Dice Score: {metrics['best_surface_dice']:.4f}\")\n        \n        # Save final results\n        results_path = base_save_dir / 'final_results.json'\n        with open(results_path, 'w') as f:\n            json.dump(results, f, indent=4)\n        print(f\"\\nSaved final results to: {results_path}\")\n        \n        # Final memory cleanup\n        gc.collect()\n        if torch.cuda.is_available():\n            torch.cuda.empty_cache()\n        \n        return results\n    \n    except Exception as e:\n        print(f\"Error in training pipeline: {str(e)}\")\n        import traceback\n        traceback.print_exc()\n        gc.collect()\n        if torch.cuda.is_available():\n            torch.cuda.empty_cache()\n        return None\n\n\nif __name__ == \"__main__\":\n    # For testing\n    #results = safer_improved_main(test_mode=True, num_epochs=2, subset_size=50)\n    \n    # For full training (uncomment to use)\n    results = main(test_mode=False, num_epochs=2)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T01:34:49.344398Z","iopub.execute_input":"2025-02-27T01:34:49.344739Z","iopub.status.idle":"2025-02-27T03:03:38.660093Z","shell.execute_reply.started":"2025-02-27T01:34:49.344715Z","shell.execute_reply":"2025-02-27T03:03:38.658333Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Visualization","metadata":{}},{"cell_type":"code","source":"import torch\nimport matplotlib.pyplot as plt\nimport numpy as np\nfrom pathlib import Path\nimport random\nimport gc","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T03:04:40.489267Z","iopub.execute_input":"2025-02-27T03:04:40.489644Z","iopub.status.idle":"2025-02-27T03:04:40.531241Z","shell.execute_reply.started":"2025-02-27T03:04:40.489615Z","shell.execute_reply":"2025-02-27T03:04:40.530260Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_predictions(model, dataset_name, val_loader, device, num_samples=3):\n    \"\"\"\n    Visualize model predictions for a given dataset\n    \"\"\"\n    model.eval()\n    \n    # Get random batch\n    try:\n        dataiter = iter(val_loader)\n        images, masks = next(dataiter)\n        \n        # Select random indices\n        batch_size = images.size(0)\n        indices = random.sample(range(batch_size), min(num_samples, batch_size))\n        \n        # Create figure\n        fig, axes = plt.subplots(num_samples, 3, figsize=(15, 5*num_samples))\n        fig.suptitle(f'Model Predictions for {dataset_name}', fontsize=16)\n        \n        with torch.no_grad():\n            # Move to device\n            images = images.to(device)\n            masks = masks.to(device)\n            \n            # Get predictions\n            outputs = model(images)\n            predictions = (outputs > 0.5).float()\n            \n            # Display images\n            for idx, sample_idx in enumerate(indices):\n                # Get single sample\n                image = images[sample_idx].cpu().squeeze().numpy()\n                mask = masks[sample_idx].cpu().squeeze().numpy()\n                pred = predictions[sample_idx].cpu().squeeze().numpy()\n                \n                # Original image\n                axes[idx, 0].imshow(image, cmap='gray')\n                axes[idx, 0].set_title('Original Image')\n                axes[idx, 0].axis('off')\n                \n                # Ground truth mask\n                axes[idx, 1].imshow(mask, cmap='Reds')  \n                axes[idx, 1].set_title('Ground Truth')\n                axes[idx, 1].axis('off')\n                \n                # Prediction\n                axes[idx, 2].imshow(pred, cmap='Reds')  \n                axes[idx, 2].set_title('Prediction')\n                axes[idx, 2].axis('off')\n        \n        plt.tight_layout()\n        plt.show()\n        plt.close()\n        \n        # Clear memory\n        del images, masks, outputs, predictions\n        gc.collect()\n        if torch.cuda.is_available():\n            torch.cuda.empty_cache()\n            \n    except Exception as e:\n        print(f\"Error visualizing predictions for {dataset_name}: {str(e)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T03:04:43.727818Z","iopub.execute_input":"2025-02-27T03:04:43.728219Z","iopub.status.idle":"2025-02-27T03:04:43.736337Z","shell.execute_reply.started":"2025-02-27T03:04:43.728189Z","shell.execute_reply":"2025-02-27T03:04:43.735573Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def main():\n    \"\"\"\n    Main function to visualize predictions for all datasets\n    \"\"\"\n    try:\n        print(\"Starting visualization...\")\n        \n        # Load dataloaders\n        print(\"\\nLoading dataloaders...\")\n        dataloader_path = Path('/kaggle/working/dataloader_info')\n        dataloaders = load_and_recreate_dataloaders(dataloader_path)\n        \n        # Initialize model\n        print(\"\\nInitializing model...\")\n        model = ResNetUNet(n_classes=1)  # Using original ResNetUNet\n        device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n        model = model.to(device)\n        \n        # Find the latest training directory with a more flexible pattern\n        training_dirs = list(Path('/kaggle/working').glob('*training_results_*'))\n        if not training_dirs:\n            raise ValueError(\"No training results directories found. Make sure you've run training first.\")\n        \n        # Sort by modification time to get the latest\n        latest_dir = max(training_dirs, key=lambda x: x.stat().st_mtime)\n        print(f\"\\nUsing results from: {latest_dir}\")\n        \n        # Visualize for each dataset\n        for dataset_name, loaders in dataloaders.items():\n            print(f\"\\nVisualizing predictions for {dataset_name}\")\n            \n            # Load best model for this dataset\n            model_path = latest_dir / dataset_name / 'best_model.pth'\n            print(f\"Looking for model at: {model_path}\")\n            \n            if model_path.exists():\n                try:\n                    checkpoint = torch.load(model_path, map_location=device)\n                    model.load_state_dict(checkpoint['model_state_dict'])\n                    print(f\"Loaded model with Best Dice: {checkpoint['best_dice']:.4f}\")\n                    \n                    # Visualize predictions\n                    visualize_predictions(\n                        model=model,\n                        dataset_name=dataset_name,\n                        val_loader=loaders['val'],\n                        device=device\n                    )\n                except Exception as e:\n                    print(f\"Error loading model for {dataset_name}: {str(e)}\")\n            else:\n                print(f\"No model found at {model_path}\")\n        \n        # Clear memory\n        del model\n        gc.collect()\n        if torch.cuda.is_available():\n            torch.cuda.empty_cache()\n            \n    except Exception as e:\n        print(f\"Error in visualization: {str(e)}\")\n        import traceback\n        traceback.print_exc()\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T03:11:51.818353Z","iopub.execute_input":"2025-02-27T03:11:51.818728Z","iopub.status.idle":"2025-02-27T03:12:18.858627Z","shell.execute_reply.started":"2025-02-27T03:11:51.818702Z","shell.execute_reply":"2025-02-27T03:12:18.857670Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Segmentation","metadata":{}},{"cell_type":"code","source":"def segment_and_visualize(model, image_tensor, device):\n    \"\"\"\n    Enhanced segmentation visualization with probability map\n    \"\"\"\n    model.eval()\n    with torch.no_grad():\n        # Move image to device and get prediction\n        image = image_tensor.to(device)\n        output = model(image.unsqueeze(0))\n        prediction = (output > 0.5).float()\n        probability = output.cpu().squeeze().numpy()  # Get raw probability\n        \n        # Convert to numpy for visualization\n        original = image.cpu().squeeze().numpy()\n        segmentation = prediction.cpu().squeeze().numpy()\n        \n        # Create overlay with alpha based on probability\n        overlay = np.zeros((*original.shape, 3))\n        overlay[..., 0] = original  # Gray channel\n        overlay[..., 1] = original  # Gray channel\n        overlay[..., 2] = original  # Gray channel\n        \n        # Create probability-weighted mask (red for vessels, more intense = higher probability)\n        prob_mask = np.zeros((*original.shape, 4))  # RGBA\n        prob_mask[..., 0] = 1.0  # Red channel\n        prob_mask[..., 3] = probability  # Alpha channel based on probability\n        \n        # Add red highlight for segmented vessels\n        mask_region = segmentation > 0\n        overlay[mask_region, 0] = 1.0  # Red channel\n        overlay[mask_region, 1] = 0.0  # Green channel\n        overlay[mask_region, 2] = 0.0  # Blue channel\n        \n        return original, segmentation, probability, overlay","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T03:12:50.219904Z","iopub.execute_input":"2025-02-27T03:12:50.220246Z","iopub.status.idle":"2025-02-27T03:12:50.226475Z","shell.execute_reply.started":"2025-02-27T03:12:50.220223Z","shell.execute_reply":"2025-02-27T03:12:50.225567Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def perform_vessel_segmentation(model, dataset_name, val_loader, device, num_samples=3):\n    \"\"\"\n    Perform and visualize vessel segmentation with probability maps\n    \"\"\"\n    # Get random batch\n    dataiter = iter(val_loader)\n    images, ground_truth = next(dataiter)\n    \n    # Select random indices\n    batch_size = images.size(0)\n    indices = random.sample(range(batch_size), min(num_samples, batch_size))\n    \n    # Create figure\n    fig, axes = plt.subplots(num_samples, 5, figsize=(25, 5*num_samples))\n    fig.suptitle(f'Enhanced Blood Vessel Segmentation Results for {dataset_name}', fontsize=16)\n    \n    for idx, sample_idx in enumerate(indices):\n        image = images[sample_idx]\n        gt_mask = ground_truth[sample_idx]\n        \n        # Get segmentation results\n        original, segmentation, probability, overlay = segment_and_visualize(model, image, device)\n        \n        # Display results\n        # Original image\n        axes[idx, 0].imshow(original, cmap='gray')\n        axes[idx, 0].set_title('Original Image')\n        axes[idx, 0].axis('off')\n        \n        # Ground truth\n        axes[idx, 1].imshow(gt_mask.squeeze().cpu().numpy(), cmap='Reds')\n        axes[idx, 1].set_title('Ground Truth Vessels')\n        axes[idx, 1].axis('off')\n        \n        # Probability map\n        axes[idx, 2].imshow(probability, cmap='hot')\n        axes[idx, 2].set_title('Probability Map')\n        axes[idx, 2].axis('off')\n        \n        # Segmented vessels\n        axes[idx, 3].imshow(segmentation, cmap='Reds')\n        axes[idx, 3].set_title('Segmented Vessels')\n        axes[idx, 3].axis('off')\n        \n        # Vessel overlay\n        axes[idx, 4].imshow(overlay)\n        axes[idx, 4].set_title('Vessel Overlay')\n        axes[idx, 4].axis('off')\n    \n    plt.tight_layout()\n    plt.subplots_adjust(top=0.95)\n    plt.show()\n    plt.close()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T03:19:43.811674Z","iopub.execute_input":"2025-02-27T03:19:43.812061Z","iopub.status.idle":"2025-02-27T03:19:43.820943Z","shell.execute_reply.started":"2025-02-27T03:19:43.812032Z","shell.execute_reply":"2025-02-27T03:19:43.820090Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualization_main():\n    \"\"\"\n    Main function to perform enhanced vessel segmentation visualization on all datasets\n    \"\"\"\n    try:\n        print(\"Starting enhanced vessel segmentation visualization...\")\n        \n        # Load dataloaders\n        print(\"\\nLoading dataloaders...\")\n        dataloader_path = Path('/kaggle/working/dataloader_info')\n        dataloaders = load_and_recreate_dataloaders(dataloader_path)\n        \n        # Initialize model\n        print(\"\\nInitializing model...\")\n        model = ResNetUNet(n_classes=1)\n        device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n        model = model.to(device)\n        \n        # Find the latest training directory with a more flexible pattern\n        training_dirs = list(Path('/kaggle/working').glob('*training_results_*'))\n        if not training_dirs:\n            raise ValueError(\"No training results directories found. Make sure you've run training first.\")\n        \n        # Sort by modification time to get the latest\n        latest_dir = max(training_dirs, key=lambda x: x.stat().st_mtime)\n        print(f\"\\nUsing models from: {latest_dir}\")\n        \n        # Process each dataset\n        for dataset_name, loaders in dataloaders.items():\n            print(f\"\\nProcessing {dataset_name}\")\n            \n            # Check for regular model\n            model_path = latest_dir / dataset_name / 'best_model.pth'\n            \n            print(f\"Loading model from: {model_path}\")\n            \n            if model_path.exists():\n                try:\n                    # Load model weights\n                    checkpoint = torch.load(model_path, map_location=device)\n                    model.load_state_dict(checkpoint['model_state_dict'])\n                    print(f\"Loaded model with Best Dice: {checkpoint['best_dice']:.4f}\")\n                    \n                    # Perform enhanced segmentation\n                    print(\"Performing enhanced vessel segmentation...\")\n                    perform_vessel_segmentation(\n                        model=model,\n                        dataset_name=dataset_name,\n                        val_loader=loaders['val'],\n                        device=device\n                    )\n                except Exception as e:\n                    print(f\"Error processing {dataset_name}: {str(e)}\")\n            else:\n                print(f\"No model found at {model_path}\")\n            \n            # Clear memory\n            gc.collect()\n            if torch.cuda.is_available():\n                torch.cuda.empty_cache()\n        \n        print(\"\\nEnhanced vessel segmentation visualization completed!\")\n        \n    except Exception as e:\n        print(f\"Error in enhanced vessel segmentation: {str(e)}\")\n        import traceback\n        traceback.print_exc()\n\nif __name__ == \"__main__\":\n    visualization_main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T03:19:48.344472Z","iopub.execute_input":"2025-02-27T03:19:48.344793Z","iopub.status.idle":"2025-02-27T03:20:11.668879Z","shell.execute_reply.started":"2025-02-27T03:19:48.344769Z","shell.execute_reply":"2025-02-27T03:20:11.667924Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Segmentation (Surface Dice)","metadata":{}},{"cell_type":"code","source":"def visualize_using_surface_dice():\n    \"\"\"\n    Visualization function that uses models with best Surface Dice instead of best Dice\n    \"\"\"\n    try:\n        print(\"Starting vessel segmentation based on best Surface Dice...\")\n        \n        # Load dataloaders\n        print(\"\\nLoading dataloaders...\")\n        dataloader_path = Path('/kaggle/working/dataloader_info')\n        dataloaders = load_and_recreate_dataloaders(dataloader_path)\n        \n        # Initialize model\n        print(\"\\nInitializing model...\")\n        model = ResNetUNet(n_classes=1)\n        device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n        model = model.to(device)\n        \n        # Find the latest training directory\n        training_dirs = list(Path('/kaggle/working').glob('*training_results_*'))\n        if not training_dirs:\n            raise ValueError(\"No training results directories found. Make sure you've run training first.\")\n        \n        # Sort by modification time to get the latest\n        latest_dir = max(training_dirs, key=lambda x: x.stat().st_mtime)\n        print(f\"\\nUsing models from: {latest_dir}\")\n        \n        # Process each dataset\n        for dataset_name, loaders in dataloaders.items():\n            print(f\"\\nProcessing {dataset_name}\")\n            \n            # Check for best model\n            model_path = latest_dir / dataset_name / 'best_model.pth'\n            \n            if model_path.exists():\n                try:\n                    # Load checkpoint and check metrics\n                    checkpoint = torch.load(model_path, map_location=device)\n                    best_dice = checkpoint['best_dice']\n                    best_surface_dice = checkpoint['best_surface_dice']\n                    \n                    print(f\"Found model with Best Dice: {best_dice:.4f} and Best Surface Dice: {best_surface_dice:.4f}\")\n                    \n                    # Load model weights\n                    model.load_state_dict(checkpoint['model_state_dict'])\n                    \n                    # Perform enhanced segmentation\n                    print(\"Performing vessel segmentation using best Surface Dice model...\")\n                    \n                    # Create figure title highlighting Surface Dice\n                    title = f'Surface Dice Optimized Vessel Segmentation for {dataset_name} (SD: {best_surface_dice:.4f})'\n                    \n                    # Get random batch\n                    dataiter = iter(loaders['val'])\n                    images, ground_truth = next(dataiter)\n                    \n                    # Select random indices\n                    batch_size = images.size(0)\n                    num_samples = 3\n                    indices = random.sample(range(batch_size), min(num_samples, batch_size))\n                    \n                    # Create figure\n                    fig, axes = plt.subplots(num_samples, 4, figsize=(20, 5*num_samples))\n                    fig.suptitle(title, fontsize=16)\n                    \n                    for idx, sample_idx in enumerate(indices):\n                        image = images[sample_idx]\n                        gt_mask = ground_truth[sample_idx]\n                        \n                        # Get segmentation results\n                        original, segmentation, probability, overlay = segment_and_visualize(model, image, device)\n                        \n                        # Display results\n                        # Original image\n                        axes[idx, 0].imshow(original, cmap='gray')\n                        axes[idx, 0].set_title('Original Image')\n                        axes[idx, 0].axis('off')\n                        \n                        # Ground truth\n                        axes[idx, 1].imshow(gt_mask.squeeze().cpu().numpy(), cmap='Reds')\n                        axes[idx, 1].set_title('Ground Truth Vessels')\n                        axes[idx, 1].axis('off')\n                        \n                        # Probability map\n                        axes[idx, 2].imshow(probability, cmap='hot')\n                        axes[idx, 2].set_title('Vessel Probability')\n                        axes[idx, 2].axis('off')\n                        \n                        # Overlay\n                        axes[idx, 3].imshow(overlay)\n                        axes[idx, 3].set_title('Vessel Overlay')\n                        axes[idx, 3].axis('off')\n                    \n                    plt.tight_layout()\n                    plt.subplots_adjust(top=0.92)\n                    plt.show()\n                    plt.close()\n                    \n                except Exception as e:\n                    print(f\"Error processing {dataset_name}: {str(e)}\")\n                    traceback.print_exc()\n            else:\n                print(f\"No model found at {model_path}\")\n            \n            # Clear memory\n            gc.collect()\n            if torch.cuda.is_available():\n                torch.cuda.empty_cache()\n        \n        print(\"\\nSurface Dice based vessel segmentation completed!\")\n        \n    except Exception as e:\n        print(f\"Error in vessel segmentation: {str(e)}\")\n        import traceback\n        traceback.print_exc()\n\nif __name__ == \"__main__\":\n    visualize_using_surface_dice()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T03:26:43.230295Z","iopub.execute_input":"2025-02-27T03:26:43.230688Z","iopub.status.idle":"2025-02-27T03:27:04.398453Z","shell.execute_reply.started":"2025-02-27T03:26:43.230660Z","shell.execute_reply":"2025-02-27T03:27:04.397748Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}