{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":91249,"databundleVersionId":11294684,"sourceType":"competition"}],"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom pathlib import Path\nimport glob\nimport random\nimport cv2\nfrom tqdm import tqdm\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.optim import Adam\nfrom sklearn.model_selection import train_test_split\nfrom torchvision import transforms\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nfrom sklearn.metrics import precision_recall_curve, average_precision_score\nfrom PIL import Image\n\n# FIX 1: Added proper error handling throughout the code\n\n# Setup device detection with better error handling\ndef setup_device():\n    \"\"\"\n    Set up the best available device: TPU > GPU > CPU\n    Returns device type and appropriate PyTorch device\n    \"\"\"\n    # Check for TPU\n    try:\n        import torch_xla.core.xla_model as xm\n        print(\"TPU available, using TPU acceleration\")\n        device = xm.xla_device()\n        return \"tpu\", device\n    except (ImportError, NameError) as e:\n        print(f\"TPU not available: {e}\")\n    \n    # Check for GPU\n    if torch.cuda.is_available():\n        device_count = torch.cuda.device_count()\n        print(f\"Found {device_count} GPU device(s)\")\n        for i in range(device_count):\n            gpu_name = torch.cuda.get_device_name(i)\n            print(f\"GPU {i}: {gpu_name}\")\n        \n        # FIX 2: Added memory check to prevent OOM errors\n        total_memory = torch.cuda.get_device_properties(0).total_memory\n        print(f\"GPU memory: {total_memory / 1e9:.2f} GB\")\n        \n        device = torch.device(\"cuda:0\")\n        # Set performance optimizations\n        torch.backends.cudnn.benchmark = True\n        print(f\"Using GPU: {torch.cuda.get_device_name(0)}\")\n        return \"gpu\", device\n    \n    # If no accelerator available, use CPU\n    print(\"No GPU or TPU found, using CPU\")\n    return \"cpu\", torch.device(\"cpu\")\n\n# Set seed for reproducibility\ndef set_seed(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed(seed)\n        torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\nset_seed()\n\n# FIX 3: Added try-except for path definition to handle Kaggle path issues\ntry:\n    # Define data paths\n    BASE_PATH = Path(\"/kaggle/input/byu-locating-bacterial-flagellar-motors-2025\")\n    if not BASE_PATH.exists():\n        raise FileNotFoundError(f\"Base path {BASE_PATH} not found\")\n    \n    TRAIN_PATH = BASE_PATH / \"train\"\n    TEST_PATH = BASE_PATH / \"test\"\n    SAMPLE_SUBMISSION = BASE_PATH / \"sample_submission.csv\"\n    TRAIN_LABELS = BASE_PATH / \"train_labels.csv\"\n    \n    # Verify paths exist\n    for path in [TRAIN_PATH, TEST_PATH, SAMPLE_SUBMISSION, TRAIN_LABELS]:\n        if not path.exists():\n            print(f\"Warning: Path {path} does not exist\")\nexcept Exception as e:\n    print(f\"Error setting up paths: {e}\")\n    # FIX: Provide fallback paths\n    BASE_PATH = Path(\".\")\n    TRAIN_PATH = BASE_PATH / \"train\"\n    TEST_PATH = BASE_PATH / \"test\"\n    SAMPLE_SUBMISSION = BASE_PATH / \"sample_submission.csv\"\n    TRAIN_LABELS = BASE_PATH / \"train_labels.csv\"\n\n# Set up device\naccelerator_type, device = setup_device()\nprint(f\"Using accelerator type: {accelerator_type}\")\nprint(f\"Using device: {device}\")\n\n# Read training labels with better error handling\ntrain_labels_df = None\ntry:\n    train_labels_df = pd.read_csv(TRAIN_LABELS)\n    print(f\"Labels data shape: {train_labels_df.shape}\")\n    print(\"Label file columns:\")\n    print(train_labels_df.columns.tolist())\n    print(\"Label data first 5 rows:\")\n    print(train_labels_df.head())\nexcept Exception as e:\n    print(f\"Error reading label file: {e}\")\n\n# Read sample submission file\nsample_submission_df = None\ntry:\n    sample_submission_df = pd.read_csv(SAMPLE_SUBMISSION)\n    print(f\"Sample submission template shape: {sample_submission_df.shape}\")\n    print(\"Sample submission file columns:\")\n    print(sample_submission_df.columns.tolist())\n    print(\"Sample submission first 5 rows:\")\n    print(sample_submission_df.head())\nexcept Exception as e:\n    print(f\"Error reading sample submission file: {e}\")\n\n# Get all tomogram folders and slices\ndef get_data_paths():\n    # Get training data tomogram folders\n    train_tomogram_folders = []\n    test_tomogram_folders = []\n    train_slices = []\n    test_slices = []\n    \n    try:\n        # Get training data tomogram folders\n        if TRAIN_PATH.exists():\n            train_tomogram_folders = [f for f in TRAIN_PATH.iterdir() if f.is_dir()]\n            train_tomogram_folders.sort()  # Ensure consistent order\n            \n            # Get training slices\n            for folder in train_tomogram_folders:\n                slices = list(folder.glob(\"*.jpg\"))\n                slices.sort()  # Ensure slices are ordered\n                train_slices.extend(slices)\n        else:\n            print(f\"Warning: Training path {TRAIN_PATH} does not exist\")\n        \n        # Get test data tomogram folders\n        if TEST_PATH.exists():\n            test_tomogram_folders = [f for f in TEST_PATH.iterdir() if f.is_dir()]\n            test_tomogram_folders.sort()\n            \n            # Get test slices\n            for folder in test_tomogram_folders:\n                slices = list(folder.glob(\"*.jpg\"))\n                slices.sort()\n                test_slices.extend(slices)\n        else:\n            print(f\"Warning: Test path {TEST_PATH} does not exist\")\n            \n    except Exception as e:\n        print(f\"Error getting data paths: {e}\")\n    \n    return train_tomogram_folders, test_tomogram_folders, train_slices, test_slices\n\ntrain_tomogram_folders, test_tomogram_folders, train_slices, test_slices = get_data_paths()\nprint(f\"Number of training tomograms: {len(train_tomogram_folders)}\")\nprint(f\"Number of test tomograms: {len(test_tomogram_folders)}\")\nprint(f\"Total number of training slices: {len(train_slices)}\")\nprint(f\"Total number of test slices: {len(test_slices)}\")\n\n# Example view of a training folder name\nif train_tomogram_folders:\n    print(f\"Example training tomogram folder name: {train_tomogram_folders[0].name}\")\n    # View slice filenames in this folder\n    slices = list(train_tomogram_folders[0].glob(\"*.jpg\"))\n    if slices:\n        print(f\"Example slice filename: {slices[0].name}\")\n\n# Extract tomogram ID and slice index from file path\ndef extract_tomo_slice_info(file_path):\n    # Extract tomogram ID from path\n    tomo_id = file_path.parent.name\n    # Extract slice index from filename\n    try:\n        # FIX 4: More robust extraction of slice index\n        filename = file_path.stem\n        if '_' in filename:\n            slice_idx = int(filename.split('_')[1])\n        else:\n            # Fallback if filename format is different\n            slice_idx = int(''.join(filter(str.isdigit, filename)))\n    except (ValueError, IndexError) as e:\n        print(f\"Error extracting slice index from {file_path}: {e}\")\n        slice_idx = 0  # Default value\n    \n    return tomo_id, slice_idx\n\n# FIX 5: More robust label preprocessing\ndef preprocess_labels(labels_df):\n    \"\"\"\n    Process labels dataframe into a dictionary for quick lookup\n    Handle different column naming conventions\n    \"\"\"\n    if labels_df is None:\n        print(\"Warning: No labels data provided\")\n        return {}\n    \n    # Check columns and adapt accordingly\n    columns = labels_df.columns.tolist()\n    \n    # Try to find the right column names\n    tomo_col = None\n    slice_col = None\n    x_col = None\n    y_col = None\n    \n    # Look for tomogram ID column\n    for col in columns:\n        if 'tomo' in col.lower():\n            tomo_col = col\n            break\n    \n    # Look for slice/row ID column\n    for col in columns:\n        if 'row' in col.lower() or 'slice' in col.lower():\n            slice_col = col\n            break\n    \n    # Look for X coordinate column\n    for col in columns:\n        if 'axis 0' in col.lower() or 'x' in col.lower():\n            x_col = col\n            break\n    \n    # Look for Y coordinate column\n    for col in columns:\n        if 'axis 1' in col.lower() or 'y' in col.lower():\n            y_col = col\n            break\n    \n    # Use default column names if not found\n    if tomo_col is None:\n        print(\"Warning: Tomogram ID column not found, using 'tomo_id'\")\n        tomo_col = 'tomo_id'\n    \n    if slice_col is None:\n        print(\"Warning: Slice index column not found, using 'row_id'\")\n        slice_col = 'row_id'\n    \n    if x_col is None:\n        print(\"Warning: X coordinate column not found, using 'Motor axis 0'\")\n        x_col = 'Motor axis 0'\n    \n    if y_col is None:\n        print(\"Warning: Y coordinate column not found, using 'Motor axis 1'\")\n        y_col = 'Motor axis 1'\n    \n    print(f\"Using columns: {tomo_col}, {slice_col}, {x_col}, {y_col}\")\n    \n    # Create dictionary to store labels\n    labels_dict = {}\n    \n    # Iterate through label data\n    for _, row in labels_df.iterrows():\n        try:\n            # Convert to proper types with error handling\n            tomo_id = str(row[tomo_col])\n            \n            # Handle different slice index formats\n            try:\n                slice_idx = int(row[slice_col])\n            except (ValueError, TypeError):\n                # Try to extract numeric part if not a clean integer\n                slice_idx = int(''.join(filter(str.isdigit, str(row[slice_col]))))\n            \n            # Extract coordinates\n            try:\n                x = float(row[x_col])\n                y = float(row[y_col])\n            except (ValueError, TypeError) as e:\n                print(f\"Error extracting coordinates: {e}\")\n                continue\n            \n            key = (tomo_id, slice_idx)\n            if key not in labels_dict:\n                labels_dict[key] = []\n            \n            # Add (x, y) coordinates to corresponding (tomo_id, slice_idx) key\n            labels_dict[key].append((x, y))\n        \n        except Exception as e:\n            print(f\"Error processing row: {e}\")\n    \n    return labels_dict\n\n# FIX 6: Check data before processing\nif train_labels_df is not None:\n    labels_dict = preprocess_labels(train_labels_df)\n    print(f\"Number of slices with labels: {len(labels_dict)}\")\n    # Show examples of first few labels\n    if labels_dict:\n        count = 0\n        for key, points in labels_dict.items():\n            print(f\"Slice {key}: {len(points)} marked points\")\n            count += 1\n            if count >= 3:\n                break\nelse:\n    print(\"Warning: No label data available\")\n    labels_dict = {}\n\n# FIX 7: Improved heatmap creation function\ndef create_heatmap(img_shape, points, sigma=10):\n    \"\"\"\n    Create Gaussian heatmap for given points\n    \n    Parameters:\n    - img_shape: Original image shape (height, width) for PIL images\n    - points: List of coordinates [(x1, y1), (x2, y2), ...]\n    - sigma: Standard deviation of Gaussian kernel\n    \n    Returns:\n    - Heatmap\n    \"\"\"\n    # For PIL images, shape is (width, height), we need (height, width)\n    height, width = img_shape[1], img_shape[0]\n    heatmap = np.zeros((height, width), dtype=np.float32)\n    \n    # If no points, return zero heatmap\n    if not points or len(points) == 0:\n        return heatmap\n    \n    # Create meshgrid for the image\n    y_grid, x_grid = np.mgrid[0:height, 0:width]\n    \n    for x, y in points:\n        # Ensure coordinates are within image bounds\n        if 0 <= x < width and 0 <= y < height:\n            # Compute Gaussian values\n            gaussian = np.exp(-((x_grid - x)**2 + (y_grid - y)**2) / (2 * sigma**2))\n            \n            # Update heatmap, taking maximum value to avoid overlap issues\n            heatmap = np.maximum(heatmap, gaussian)\n    \n    return heatmap\n\n# FIX 8: Improved dataset class with better error handling\nclass BacterialMotorDataset(Dataset):\n    def __init__(self, image_paths, labels_dict=None, transform=None, is_test=False):\n        self.image_paths = image_paths\n        self.labels_dict = labels_dict if labels_dict is not None else {}\n        self.transform = transform\n        self.is_test = is_test\n        self.target_size = (256, 256)  # Default target size\n        \n    def __len__(self):\n        return len(self.image_paths)\n    \n    def __getitem__(self, idx):\n        img_path = self.image_paths[idx]\n        \n        try:\n            # Read image with error handling\n            img = cv2.imread(str(img_path), cv2.IMREAD_GRAYSCALE)\n            if img is None:\n                print(f\"Error reading image: {img_path}\")\n                # Create a blank image as fallback\n                img = np.zeros((256, 256), dtype=np.uint8)\n            \n            # Convert to PIL Image for transformations\n            img_pil = Image.fromarray(img)\n            \n            # Get original dimensions\n            original_size = img_pil.size  # (width, height)\n            \n            # Extract tomogram ID and slice index\n            tomo_id, slice_idx = extract_tomo_slice_info(img_path)\n            \n            # Apply transforms\n            if self.transform:\n                img_transformed = self.transform(img_pil)\n            else:\n                # Default transform if none provided\n                img_transformed = transforms.ToTensor()(img_pil)\n            \n            # For test set, return minimal information\n            if self.is_test:\n                return {\n                    'image': img_transformed,\n                    'tomo_id': tomo_id,\n                    'slice_idx': slice_idx,\n                    'image_path': str(img_path)\n                }\n            \n            # Get label points or empty list if not found\n            key = (tomo_id, slice_idx)\n            points = self.labels_dict.get(key, [])\n            \n            # Create heatmap label before resize\n            heatmap = create_heatmap(original_size, points)\n            \n            # Resize heatmap to match transformed image size\n            heatmap_resized = cv2.resize(\n                heatmap, \n                (self.target_size[1], self.target_size[0]),  # (width, height)\n                interpolation=cv2.INTER_LINEAR\n            )\n            \n            # Convert to tensor\n            heatmap_tensor = torch.tensor(heatmap_resized, dtype=torch.float32).unsqueeze(0)\n            \n            return {\n                'image': img_transformed,\n                'heatmap': heatmap_tensor,\n                'points': points,\n                'tomo_id': tomo_id,\n                'slice_idx': slice_idx,\n                'image_path': str(img_path)\n            }\n            \n        except Exception as e:\n            print(f\"Error processing image {img_path}: {e}\")\n            # Return a default item to avoid breaking the DataLoader\n            default_img = torch.zeros((1, self.target_size[0], self.target_size[1]), dtype=torch.float32)\n            default_heatmap = torch.zeros((1, self.target_size[0], self.target_size[1]), dtype=torch.float32)\n            \n            return {\n                'image': default_img,\n                'heatmap': default_heatmap,\n                'points': [],\n                'tomo_id': \"error\",\n                'slice_idx': -1,\n                'image_path': str(img_path)\n            }\n\n# FIX 9: Simplified transforms with fixed size\ntrain_transform = transforms.Compose([\n    transforms.Resize((256, 256)),\n    transforms.ToTensor(),\n])\n\nval_transform = transforms.Compose([\n    transforms.Resize((256, 256)),\n    transforms.ToTensor(),\n])\n\n# FIX 10: Simplified and more robust collate function\ndef custom_collate_fn(batch):\n    \"\"\"\n    Collate function that handles potential errors in the batch\n    \"\"\"\n    # Filter out any None values\n    valid_batch = [item for item in batch if item is not None]\n    \n    if not valid_batch:\n        print(\"Warning: Empty batch after filtering\")\n        # Return empty tensors as fallback\n        return {\n            'image': torch.zeros((0, 1, 256, 256), dtype=torch.float32),\n            'heatmap': torch.zeros((0, 1, 256, 256), dtype=torch.float32),\n            'tomo_id': [],\n            'slice_idx': [],\n            'image_path': []\n        }\n    \n    # Extract elements from valid items\n    images = [item['image'] for item in valid_batch if 'image' in item]\n    \n    # Handle test batch differently\n    is_test = 'heatmap' not in valid_batch[0]\n    \n    if not is_test:\n        heatmaps = [item['heatmap'] for item in valid_batch if 'heatmap' in item]\n        tomo_ids = [item['tomo_id'] for item in valid_batch if 'tomo_id' in item]\n        slice_idxs = [item['slice_idx'] for item in valid_batch if 'slice_idx' in item]\n        image_paths = [item['image_path'] for item in valid_batch if 'image_path' in item]\n        \n        # Stack tensors\n        if images:\n            images = torch.stack(images, dim=0)\n        else:\n            images = torch.zeros((0, 1, 256, 256), dtype=torch.float32)\n            \n        if heatmaps:\n            heatmaps = torch.stack(heatmaps, dim=0)\n        else:\n            heatmaps = torch.zeros((0, 1, 256, 256), dtype=torch.float32)\n        \n        return {\n            'image': images,\n            'heatmap': heatmaps,\n            'tomo_id': tomo_ids,\n            'slice_idx': slice_idxs,\n            'image_path': image_paths\n        }\n    else:\n        # Test batch\n        tomo_ids = [item['tomo_id'] for item in valid_batch if 'tomo_id' in item]\n        slice_idxs = [item['slice_idx'] for item in valid_batch if 'slice_idx' in item]\n        image_paths = [item['image_path'] for item in valid_batch if 'image_path' in item]\n        \n        # Stack tensors\n        if images:\n            images = torch.stack(images, dim=0)\n        else:\n            images = torch.zeros((0, 1, 256, 256), dtype=torch.float32)\n        \n        return {\n            'image': images,\n            'tomo_id': tomo_ids,\n            'slice_idx': slice_idxs,\n            'image_path': image_paths\n        }\n\n# FIX 11: Improved UNet architecture with better initialization\nclass DoubleConv(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super(DoubleConv, self).__init__()\n        self.double_conv = nn.Sequential(\n            nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True)\n        )\n        \n        # Initialize weights\n        for m in self.modules():\n            if isinstance(m, nn.Conv2d):\n                nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')\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        return self.double_conv(x)\n\nclass UNet(nn.Module):\n    def __init__(self, in_channels=1, out_channels=1, init_features=32):\n        super(UNet, self).__init__()\n        \n        features = init_features\n        self.encoder1 = DoubleConv(in_channels, features)\n        self.pool1 = nn.MaxPool2d(kernel_size=2, stride=2)\n        \n        self.encoder2 = DoubleConv(features, features * 2)\n        self.pool2 = nn.MaxPool2d(kernel_size=2, stride=2)\n        \n        self.encoder3 = DoubleConv(features * 2, features * 4)\n        self.pool3 = nn.MaxPool2d(kernel_size=2, stride=2)\n        \n        self.encoder4 = DoubleConv(features * 4, features * 8)\n        self.pool4 = nn.MaxPool2d(kernel_size=2, stride=2)\n        \n        self.bottleneck = DoubleConv(features * 8, features * 16)\n        \n        self.upconv4 = nn.ConvTranspose2d(features * 16, features * 8, kernel_size=2, stride=2)\n        self.decoder4 = DoubleConv(features * 16, features * 8)\n        \n        self.upconv3 = nn.ConvTranspose2d(features * 8, features * 4, kernel_size=2, stride=2)\n        self.decoder3 = DoubleConv(features * 8, features * 4)\n        \n        self.upconv2 = nn.ConvTranspose2d(features * 4, features * 2, kernel_size=2, stride=2)\n        self.decoder2 = DoubleConv(features * 4, features * 2)\n        \n        self.upconv1 = nn.ConvTranspose2d(features * 2, features, kernel_size=2, stride=2)\n        self.decoder1 = DoubleConv(features * 2, features)\n        \n        self.conv = nn.Conv2d(features, out_channels, kernel_size=1)\n        \n    def forward(self, x):\n        # FIX 12: Add input shape check\n        if x.dim() != 4:\n            raise ValueError(f\"Expected 4D tensor, got {x.dim()}D tensor\")\n        \n        enc1 = self.encoder1(x)\n        enc2 = self.encoder2(self.pool1(enc1))\n        enc3 = self.encoder3(self.pool2(enc2))\n        enc4 = self.encoder4(self.pool3(enc3))\n        \n        bottleneck = self.bottleneck(self.pool4(enc4))\n        \n        dec4 = self.upconv4(bottleneck)\n        # FIX 13: Add size check for concatenation\n        if dec4.size()[2:] != enc4.size()[2:]:\n            dec4 = F.interpolate(dec4, size=enc4.size()[2:], mode='bilinear', align_corners=False)\n        dec4 = torch.cat((dec4, enc4), dim=1)\n        dec4 = self.decoder4(dec4)\n        \n        dec3 = self.upconv3(dec4)\n        if dec3.size()[2:] != enc3.size()[2:]:\n            dec3 = F.interpolate(dec3, size=enc3.size()[2:], mode='bilinear', align_corners=False)\n        dec3 = torch.cat((dec3, enc3), dim=1)\n        dec3 = self.decoder3(dec3)\n        \n        dec2 = self.upconv2(dec3)\n        if dec2.size()[2:] != enc2.size()[2:]:\n            dec2 = F.interpolate(dec2, size=enc2.size()[2:], mode='bilinear', align_corners=False)\n        dec2 = torch.cat((dec2, enc2), dim=1)\n        dec2 = self.decoder2(dec2)\n        \n        dec1 = self.upconv1(dec2)\n        if dec1.size()[2:] != enc1.size()[2:]:\n            dec1 = F.interpolate(dec1, size=enc1.size()[2:], mode='bilinear', align_corners=False)\n        dec1 = torch.cat((dec1, enc1), dim=1)\n        dec1 = self.decoder1(dec1)\n        \n        output = self.conv(dec1)\n        output = torch.sigmoid(output)  # Use sigmoid for [0,1] range\n        \n        return output\n\n# FIX 14: Training function with proper error handling and GPU memory optimization\ndef train_model(model, train_loader, val_loader, criterion, optimizer, scheduler, num_epochs=10, accelerator_type='cpu'):\n    best_val_loss = float('inf')\n    history = {'train_loss': [], 'val_loss': []}\n    \n    # For TPU acceleration\n    if accelerator_type == 'tpu':\n        try:\n            import torch_xla.core.xla_model as xm\n            import torch_xla.distributed.parallel_loader as pl\n        except ImportError as e:\n            print(f\"Error importing TPU modules: {e}\")\n            print(\"Falling back to CPU training\")\n            accelerator_type = 'cpu'\n    \n    # For mixed precision training on GPU\n    use_amp = False\n    scaler = None\n    if accelerator_type == 'gpu':\n        try:\n            from torch.cuda.amp import GradScaler, autocast\n            use_amp = True\n            scaler = GradScaler()\n            print(\"Mixed precision training enabled\")\n        except ImportError:\n            print(\"Mixed precision not available, using full precision\")\n    \n    for epoch in range(num_epochs):\n        # Training phase\n        model.train()\n        train_loss = 0.0\n        \n        # Choose appropriate dataloader based on accelerator\n        if accelerator_type == 'tpu':\n            try:\n                train_device_loader = pl.ParallelLoader(train_loader, [device]).per_device_loader(device)\n                loader_to_use = train_device_loader\n            except Exception as e:\n                print(f\"Error setting up TPU dataloader: {e}\")\n                loader_to_use = train_loader\n        else:\n            loader_to_use = train_loader\n        \n        # Training loop with error handling\n        pbar = tqdm(loader_to_use, desc=f'Epoch {epoch+1}/{num_epochs} [Train]')\n        for batch_idx, batch in enumerate(pbar):\n            try:\n                images = batch['image'].to(device)\n                heatmaps = batch['heatmap'].to(device)\n                \n                # Check for empty batch\n                if images.size(0) == 0:\n                    print(\"Warning: Empty batch encountered, skipping\")\n                    continue\n                \n                # Forward pass with appropriate precision\n                if use_amp:\n                    with autocast():\n                        outputs = model(images)\n                        loss = criterion(outputs, heatmaps)\n                else:\n                    outputs = model(images)\n                    loss = criterion(outputs, heatmaps)\n                \n                # Backward pass\n                optimizer.zero_grad()\n                \n                if use_amp:\n                    scaler.scale(loss).backward()\n                    scaler.step(optimizer)\n                    scaler.update()\n                elif accelerator_type == 'tpu':\n                    loss.backward()\n                    xm.optimizer_step(optimizer, barrier=True)\n                else:\n                    loss.backward()\n                    optimizer.step()\n                \n                # Update statistics\n                train_loss += loss.item() * images.size(0)\n                \n                # Update progress bar\n                pbar.set_postfix(loss=loss.item())\n                \n                # FIX 15: Periodic GPU memory cleanup\n                if accelerator_type == 'gpu' and batch_idx % 10 == 0:\n                    torch.cuda.empty_cache()\n                \n            except Exception as e:\n                print(f\"Error in training batch {batch_idx}: {e}\")\n                continue\n        \n        # Calculate average loss\n        train_loss /= len(train_loader.dataset)\n        \n        # Validation phase\n        model.eval()\n        val_loss = 0.0\n        \n        # Choose appropriate validation dataloader\n        if accelerator_type == 'tpu':\n            try:\n                val_device_loader = pl.ParallelLoader(val_loader, [device]).per_device_loader(device)\n                val_loader_to_use = val_device_loader\n            except Exception as e:\n                print(f\"Error setting up TPU validation dataloader: {e}\")\n                val_loader_to_use = val_loader\n        \n        with torch.no_grad():\n            pbar = tqdm(val_loader_to_use, desc=f'Epoch {epoch+1}/{num_epochs} [Val]')\n            for batch_idx, batch in enumerate(pbar):\n                try:\n                    images = batch['image'].to(device)\n                    heatmaps = batch['heatmap'].to(device)\n                    \n                    # Check for empty batch\n                    if images.size(0) == 0:\n                        print(\"Warning: Empty validation batch encountered, skipping\")\n                        continue\n                    \n                    # Forward pass\n                    outputs = model(images)\n                    loss = criterion(outputs, heatmaps)\n                    \n                    # Update statistics\n                    val_loss += loss.item() * images.size(0)\n                    \n                    # Update progress bar\n                    pbar.set_postfix(loss=loss.item())\n                    \n                except Exception as e:\n                    print(f\"Error in validation batch {batch_idx}: {e}\")\n                    continue\n        \n        # Calculate average validation loss\n        val_loss /= len(val_loader.dataset)\n        \n        # Update learning rate scheduler\n        if scheduler is not None:\n            if isinstance(scheduler, ReduceLROnPlateau):\n                scheduler.step(val_loss)\n            else:\n                scheduler.step()\n        \n        # Save model if validation loss improved\n        if val_loss < best_val_loss:\n            best_val_loss = val_loss\n            # Save model\n            save_path = Path(\"./best_model.pth\")\n            try:\n                if accelerator_type == 'tpu':\n                    # TPU-specific saving\n                    xm.save(model.state_dict(), str(save_path))\n                else:\n                    torch.save(model.state_dict(), save_path)\n                print(f\"Model saved to {save_path}\")\n            except Exception as e:\n                print(f\"Error saving model: {e}\")\n        \n        # Update history\n        history['train_loss'].append(train_loss)\n        history['val_loss'].append(val_loss)\n        \n        # Print epoch results\n        print(f'Epoch {epoch+1}/{num_epochs}: '\n              f'Train Loss: {train_loss:.4f}, '\n              f'Val Loss: {val_loss:.4f}')\n        \n        # Plot training progress\n        if (epoch + 1) % 5 == 0 or epoch == num_epochs - 1:\n            try:\n                plt.figure(figsize=(10, 5))\n                plt.subplot(1, 2, 1)\n                plt.plot(history['train_loss'], label='Train Loss')\n                plt.plot(history['val_loss'], label='Val Loss')\n                plt.xlabel('Epoch')\n                plt.ylabel('Loss')\n                plt.legend()\n                plt.title('Training Progress')\n                plt.savefig(f'training_progress_epoch_{epoch+1}.png')\n                plt.close()\n            except Exception as e:\n                print(f\"Error plotting training progress: {e}\")\n    \n    return model, history\n\n# FIX 16: Enhanced post-processing function to detect peaks in the heatmap\ndef detect_points_from_heatmap(heatmap, threshold=0.3, min_distance=10, original_size=None):\n    \"\"\"\n    Detect points from heatmap by finding local maxima\n    \n    Parameters:\n    - heatmap: Predicted heatmap tensor\n    - threshold: Minimum value to consider as a potential motor\n    - min_distance: Minimum distance between detected points\n    - original_size: Tuple (height, width) for scaling back to original image size\n    \n    Returns:\n    - List of (x, y) coordinates of detected motors\n    \"\"\"\n    from scipy.ndimage import gaussian_filter, maximum_filter\n    from scipy.ndimage.morphology import generate_binary_structure, binary_erosion\n    \n    # Convert tensor to numpy if needed\n    if isinstance(heatmap, torch.Tensor):\n        heatmap = heatmap.squeeze().cpu().numpy()\n    \n    # Ensure heatmap is 2D\n    if heatmap.ndim > 2:\n        heatmap = heatmap.squeeze()\n    \n    # Apply Gaussian filter to smooth the heatmap\n    heatmap_smoothed = gaussian_filter(heatmap, sigma=1)\n    \n    # Find local maxima\n    # Create a mask of local maxima\n    neighborhood = generate_binary_structure(2, 2)\n    local_max = maximum_filter(heatmap_smoothed, footprint=neighborhood) == heatmap_smoothed\n    background = (heatmap_smoothed < threshold)\n    eroded_background = binary_erosion(background, structure=neighborhood, border_value=1)\n    detected_maxima = local_max & ~eroded_background\n    \n    # Extract coordinates of maxima\n    y_indices, x_indices = np.where(detected_maxima)\n    \n    # Get heatmap values at maxima\n    intensities = heatmap_smoothed[y_indices, x_indices]\n    \n    # Sort points by intensity (highest to lowest)\n    sorted_indices = np.argsort(intensities)[::-1]\n    y_indices = y_indices[sorted_indices]\n    x_indices = x_indices[sorted_indices]\n    \n    # Filter close points (keep the stronger one)\n    points = []\n    for i in range(len(x_indices)):\n        # Check if this point is far enough from all accepted points\n        is_far_enough = True\n        for px, py in points:\n            dist = np.sqrt((px - x_indices[i])**2 + (py - y_indices[i])**2)\n            if dist < min_distance:\n                is_far_enough = False\n                break\n        \n        if is_far_enough:\n            points.append((x_indices[i], y_indices[i]))\n    \n    # Rescale points to original image size if provided\n    if original_size is not None:\n        h_scale = original_size[0] / heatmap.shape[0]\n        w_scale = original_size[1] / heatmap.shape[1]\n        \n        points = [(int(x * w_scale), int(y * h_scale)) for x, y in points]\n    \n    return points\n\n# FIX 17: Add evaluation metrics\ndef calculate_metrics(true_points, pred_points, distance_threshold=20):\n    \"\"\"\n    Calculate precision, recall, and F1 score\n    \n    Parameters:\n    - true_points: List of ground truth points [(x1, y1), (x2, y2), ...]\n    - pred_points: List of predicted points [(x1, y1), (x2, y2), ...]\n    - distance_threshold: Maximum distance to consider a prediction correct\n    \n    Returns:\n    - Dictionary with precision, recall, and F1 metrics\n    \"\"\"\n    # Handle empty cases\n    if not true_points and not pred_points:\n        return {'precision': 1.0, 'recall': 1.0, 'f1': 1.0}\n    elif not true_points:\n        return {'precision': 0.0, 'recall': 1.0, 'f1': 0.0}\n    elif not pred_points:\n        return {'precision': 1.0, 'recall': 0.0, 'f1': 0.0}\n    \n    # Convert to numpy arrays for easier computation\n    true_points = np.array(true_points)\n    pred_points = np.array(pred_points)\n    \n    # Calculate distances between all pairs of points\n    true_count = len(true_points)\n    pred_count = len(pred_points)\n    \n    # Initialize match matrices\n    true_matched = np.zeros(true_count, dtype=bool)\n    pred_matched = np.zeros(pred_count, dtype=bool)\n    \n    # For each predicted point, find the closest true point\n    for i in range(pred_count):\n        min_dist = float('inf')\n        closest_idx = -1\n        \n        for j in range(true_count):\n            if true_matched[j]:\n                continue  # Skip already matched true points\n                \n            # Calculate Euclidean distance\n            dist = np.sqrt(np.sum((pred_points[i] - true_points[j])**2))\n            \n            if dist < min_dist and dist <= distance_threshold:\n                min_dist = dist\n                closest_idx = j\n        \n        # If a match was found, mark both points as matched\n        if closest_idx != -1:\n            true_matched[closest_idx] = True\n            pred_matched[i] = True\n    \n    # Calculate metrics\n    true_positives = np.sum(true_matched)\n    false_positives = pred_count - true_positives\n    false_negatives = true_count - true_positives\n    \n    precision = true_positives / (true_positives + false_positives) if (true_positives + false_positives) > 0 else 0\n    recall = true_positives / (true_positives + false_negatives) if (true_positives + false_negatives) > 0 else 0\n    f1 = 2 * precision * recall / (precision + recall) if (precision + recall) > 0 else 0\n    \n    return {\n        'precision': precision,\n        'recall': recall,\n        'f1': f1,\n        'true_positives': int(true_positives),\n        'false_positives': int(false_positives),\n        'false_negatives': int(false_negatives)\n    }\n\n# FIX 18: Improved evaluation function\ndef evaluate_model(model, test_loader, device, detection_threshold=0.3, distance_threshold=20):\n    \"\"\"\n    Evaluate the model on test data\n    \n    Parameters:\n    - model: Trained UNet model\n    - test_loader: DataLoader for test data\n    - device: Device to run evaluation on\n    - detection_threshold: Threshold for peak detection\n    - distance_threshold: Distance threshold for considering a detection correct\n    \n    Returns:\n    - Dictionary with evaluation metrics\n    \"\"\"\n    model.eval()\n    all_metrics = []\n    all_predictions = []\n    \n    with torch.no_grad():\n        for batch in tqdm(test_loader, desc=\"Evaluating\"):\n            try:\n                images = batch['image'].to(device)\n                \n                # Skip empty batches\n                if images.size(0) == 0:\n                    continue\n                \n                # Forward pass\n                outputs = model(images)\n                \n                # Process each image in the batch\n                for i in range(images.size(0)):\n                    heatmap = outputs[i].squeeze().cpu().numpy()\n                    \n                    # Get ground truth points if available\n                    if 'points' in batch:\n                        true_points = batch['points'][i]\n                    else:\n                        true_points = []\n                    \n                    # Detect points from heatmap\n                    pred_points = detect_points_from_heatmap(\n                        heatmap, \n                        threshold=detection_threshold, \n                        min_distance=10\n                    )\n                    \n                    # Calculate metrics if ground truth is available\n                    if true_points:\n                        metrics = calculate_metrics(\n                            true_points, \n                            pred_points, \n                            distance_threshold=distance_threshold\n                        )\n                        all_metrics.append(metrics)\n                    \n                    # Save prediction info\n                    all_predictions.append({\n                        'tomo_id': batch['tomo_id'][i],\n                        'slice_idx': batch['slice_idx'][i],\n                        'points': pred_points,\n                    })\n            \n            except Exception as e:\n                print(f\"Error in evaluation batch: {e}\")\n                continue\n    \n    # Calculate average metrics\n    if all_metrics:\n        avg_metrics = {\n            'precision': np.mean([m['precision'] for m in all_metrics]),\n            'recall': np.mean([m['recall'] for m in all_metrics]),\n            'f1': np.mean([m['f1'] for m in all_metrics])\n        }\n        print(f\"Average Metrics - Precision: {avg_metrics['precision']:.4f}, \"\n              f\"Recall: {avg_metrics['recall']:.4f}, F1: {avg_metrics['f1']:.4f}\")\n    else:\n        avg_metrics = {}\n    \n    return all_predictions, avg_metrics\n\n# FIX 19: Function to prepare submission file\ndef prepare_submission(predictions, template_path):\n    \"\"\"\n    Create a submission file from predictions\n    \n    Parameters:\n    - predictions: List of dictionaries with prediction info\n    - template_path: Path to sample submission template\n    \n    Returns:\n    - DataFrame with submission format\n    \"\"\"\n    try:\n        # Read template\n        template_df = pd.read_csv(template_path)\n        \n        # Convert predictions to DataFrame format\n        rows = []\n        \n        for pred in predictions:\n            tomo_id = pred['tomo_id']\n            slice_idx = pred['slice_idx']\n            points = pred['points']\n            \n            # For each detected point, create a row\n            for x, y in points:\n                rows.append({\n                    'tomo_id': tomo_id,\n                    'row_id': slice_idx,\n                    'Motor axis 0': x,\n                    'Motor axis 1': y\n                })\n        \n        # Create DataFrame\n        submission_df = pd.DataFrame(rows)\n        \n        # If no predictions, create empty DataFrame with correct columns\n        if len(rows) == 0:\n            submission_df = pd.DataFrame(columns=['tomo_id', 'row_id', 'Motor axis 0', 'Motor axis 1'])\n        \n        return submission_df\n        \n    except Exception as e:\n        print(f\"Error preparing submission: {e}\")\n        return None\n\n# FIX 20: Visualization functions\ndef visualize_predictions(image_path, true_points=None, pred_points=None, save_path=None):\n    \"\"\"\n    Visualize image with ground truth and predicted points\n    \n    Parameters:\n    - image_path: Path to image\n    - true_points: List of ground truth points [(x1, y1), (x2, y2), ...]\n    - pred_points: List of predicted points [(x1, y1), (x2, y2), ...]\n    - save_path: Path to save visualization, if None, display only\n    \"\"\"\n    try:\n        # Read image\n        image = cv2.imread(str(image_path), cv2.IMREAD_GRAYSCALE)\n        \n        if image is None:\n            print(f\"Error reading image: {image_path}\")\n            return\n        \n        # Convert to RGB for visualization\n        image_rgb = cv2.cvtColor(image, cv2.COLOR_GRAY2RGB)\n        \n        # Draw ground truth points (green)\n        if true_points:\n            for x, y in true_points:\n                x, y = int(x), int(y)\n                if 0 <= x < image.shape[1] and 0 <= y < image.shape[0]:\n                    cv2.circle(image_rgb, (x, y), 5, (0, 255, 0), -1)\n                    cv2.circle(image_rgb, (x, y), 7, (0, 255, 0), 2)\n        \n        # Draw predicted points (red)\n        if pred_points:\n            for x, y in pred_points:\n                x, y = int(x), int(y)\n                if 0 <= x < image.shape[1] and 0 <= y < image.shape[0]:\n                    cv2.circle(image_rgb, (x, y), 5, (255, 0, 0), -1)\n                    cv2.circle(image_rgb, (x, y), 7, (255, 0, 0), 2)\n        \n        # Create figure\n        plt.figure(figsize=(10, 10))\n        plt.imshow(image_rgb)\n        plt.title(f\"Image: {Path(image_path).name}\")\n        \n        # Add legend\n        if true_points and pred_points:\n            from matplotlib.patches import Patch\n            legend_elements = [\n                Patch(facecolor='green', edgecolor='green', label='Ground Truth'),\n                Patch(facecolor='red', edgecolor='red', label='Prediction')\n            ]\n            plt.legend(handles=legend_elements, loc='upper right')\n        \n        # Save or display\n        if save_path:\n            plt.savefig(save_path, dpi=300, bbox_inches='tight')\n            plt.close()\n        else:\n            plt.show()\n            \n    except Exception as e:\n        print(f\"Error in visualization: {e}\")\n\n# FIX 21: Main execution function\ndef run_pipeline(train_ratio=0.8, batch_size=16, num_epochs=15, learning_rate=1e-4):\n    \"\"\"\n    Run the complete training and evaluation pipeline\n    \n    Parameters:\n    - train_ratio: Ratio of data to use for training vs validation\n    - batch_size: Batch size for training\n    - num_epochs: Number of training epochs\n    - learning_rate: Learning rate for optimizer\n    \"\"\"\n    try:\n        print(\"Setting up data...\")\n        \n        # Split training data into train/val sets\n        if train_slices:\n            train_paths, val_paths = train_test_split(\n                train_slices, \n                test_size=1-train_ratio, \n                random_state=42\n            )\n            \n            print(f\"Training on {len(train_paths)} images, validating on {len(val_paths)} images\")\n            \n            # Create datasets\n            train_dataset = BacterialMotorDataset(\n                image_paths=train_paths,\n                labels_dict=labels_dict,\n                transform=train_transform\n            )\n            \n            val_dataset = BacterialMotorDataset(\n                image_paths=val_paths,\n                labels_dict=labels_dict,\n                transform=val_transform\n            )\n            \n            # Create dataloaders\n            train_loader = DataLoader(\n                train_dataset,\n                batch_size=batch_size,\n                shuffle=True,\n                num_workers=4,\n                collate_fn=custom_collate_fn,\n                pin_memory=(accelerator_type == 'gpu')\n            )\n            \n            val_loader = DataLoader(\n                val_dataset,\n                batch_size=batch_size,\n                shuffle=False,\n                num_workers=4,\n                collate_fn=custom_collate_fn,\n                pin_memory=(accelerator_type == 'gpu')\n            )\n            \n            # Create test dataset\n            test_dataset = BacterialMotorDataset(\n                image_paths=test_slices,\n                transform=val_transform,\n                is_test=True\n            )\n            \n            test_loader = DataLoader(\n                test_dataset,\n                batch_size=batch_size,\n                shuffle=False,\n                num_workers=4,\n                collate_fn=custom_collate_fn,\n                pin_memory=(accelerator_type == 'gpu')\n            )\n            \n            # Create model\n            model = UNet(in_channels=1, out_channels=1, init_features=32)\n            model = model.to(device)\n            \n            # Define loss function and optimizer\n            criterion = nn.BCELoss()\n            optimizer = Adam(model.parameters(), lr=learning_rate)\n            \n            # Define scheduler\n            scheduler = ReduceLROnPlateau(\n                optimizer, \n                mode='min', \n                factor=0.5, \n                patience=5, \n                verbose=True\n            )\n            \n            # Train model\n            print(\"Starting training...\")\n            model, history = train_model(\n                model=model,\n                train_loader=train_loader,\n                val_loader=val_loader,\n                criterion=criterion,\n                optimizer=optimizer,\n                scheduler=scheduler,\n                num_epochs=num_epochs,\n                accelerator_type=accelerator_type\n            )\n            \n            # Evaluate model\n            print(\"Evaluating model...\")\n            predictions, metrics = evaluate_model(\n                model=model,\n                test_loader=test_loader,\n                device=device\n            )\n            \n            # Create submission file\n            print(\"Creating submission file...\")\n            submission_df = prepare_submission(\n                predictions=predictions,\n                template_path=SAMPLE_SUBMISSION\n            )\n            \n            if submission_df is not None:\n                submission_path = Path(\"./submission.csv\")\n                submission_df.to_csv(submission_path, index=False)\n                print(f\"Submission saved to {submission_path}\")\n            \n            # Visualize some predictions\n            print(\"Creating visualizations...\")\n            # Visualize a few test predictions\n            for i in range(min(5, len(test_slices))):\n                try:\n                    image_path = test_slices[i]\n                    # Get prediction for this image\n                    model.eval()\n                    with torch.no_grad():\n                        # Load and preprocess image\n                        img = Image.open(image_path).convert('L')\n                        img_tensor = val_transform(img).unsqueeze(0).to(device)\n                        \n                        # Forward pass\n                        output = model(img_tensor)\n                        \n                        # Get predictions\n                        heatmap = output.squeeze().cpu().numpy()\n                        pred_points = detect_points_from_heatmap(\n                            heatmap, \n                            threshold=0.3,\n                            min_distance=10,\n                            original_size=img.size\n                        )\n                        \n                        # Save visualization\n                        save_path = Path(f\"./visualization_test_{i}.png\")\n                        visualize_predictions(\n                            image_path=image_path,\n                            pred_points=pred_points,\n                            save_path=save_path\n                        )\n                        print(f\"Visualization saved to {save_path}\")\n                        \n                except Exception as e:\n                    print(f\"Error visualizing test image {i}: {e}\")\n                    continue\n            \n            # Also visualize some training images with ground truth\n            for i in range(min(5, len(train_paths))):\n                try:\n                    image_path = train_paths[i]\n                    # Extract tomogram ID and slice index\n                    tomo_id, slice_idx = extract_tomo_slice_info(image_path)\n                    # Get ground truth points\n                    true_points = labels_dict.get((tomo_id, slice_idx), [])\n                    \n                    # Get prediction for this image\n                    model.eval()\n                    with torch.no_grad():\n                        # Load and preprocess image\n                        img = Image.open(image_path).convert('L')\n                        img_tensor = val_transform(img).unsqueeze(0).to(device)\n                        \n                        # Forward pass\n                        output = model(img_tensor)\n                        \n                        # Get predictions\n                        heatmap = output.squeeze().cpu().numpy()\n                        pred_points = detect_points_from_heatmap(\n                            heatmap, \n                            threshold=0.3,\n                            min_distance=10,\n                            original_size=img.size\n                        )\n                        \n                        # Save visualization\n                        save_path = Path(f\"./visualization_train_{i}.png\")\n                        visualize_predictions(\n                            image_path=image_path,\n                            true_points=true_points,\n                            pred_points=pred_points,\n                            save_path=save_path\n                        )\n                        print(f\"Visualization saved to {save_path}\")\n                        \n                except Exception as e:\n                    print(f\"Error visualizing training image {i}: {e}\")\n                    continue\n                \n            return model, history, metrics\n        else:\n            print(\"No training data found!\")\n            return None, None, None\n            \n    except Exception as e:\n        print(f\"Error in pipeline execution: {e}\")\n        return None, None, None\n\n# Run the pipeline if executing as main script\nif __name__ == \"__main__\":\n    # Set hyperparameters\n    HYPERPARAMS = {\n        'train_ratio': 0.8,\n        'batch_size': 16,\n        'num_epochs': 30,\n        'learning_rate': 1e-4\n    }\n    \n    # Print system info\n    print(\"System information:\")\n    print(f\"PyTorch version: {torch.__version__}\")\n    print(f\"CUDA available: {torch.cuda.is_available()}\")\n    if torch.cuda.is_available():\n        print(f\"CUDA version: {torch.version.cuda}\")\n        print(f\"GPU count: {torch.cuda.device_count()}\")\n        print(f\"GPU name: {torch.cuda.get_device_name(0)}\")\n    \n    # Run pipeline\n    model, history, metrics = run_pipeline(\n        train_ratio=HYPERPARAMS['train_ratio'],\n        batch_size=HYPERPARAMS['batch_size'],\n        num_epochs=HYPERPARAMS['num_epochs'],\n        learning_rate=HYPERPARAMS['learning_rate']\n    )","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-05-11T03:20:56.028965Z","iopub.execute_input":"2025-05-11T03:20:56.029220Z"}},"outputs":[],"execution_count":null}]}