{"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"}],"dockerImageVersionId":31011,"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\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\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\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\")\n            if path == SAMPLE_SUBMISSION:\n                print(\"Warning: Sample submission file might be needed later.\")\n            if path == TRAIN_LABELS:\n                print(\"Warning: Train labels file might be needed later.\")\n\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    print(\"Using fallback paths in current directory.\")\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    if TRAIN_LABELS.exists():\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())\n    else:\n        print(f\"Training labels file not found at: {TRAIN_LABELS}\")\nexcept Exception as e:\n    print(f\"Error reading label file {TRAIN_LABELS}: {e}\")\n\n# Read sample submission file\nsample_submission_df = None\ntry:\n    if SAMPLE_SUBMISSION.exists():\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())\n    else:\n        print(f\"Sample submission file not found at: {SAMPLE_SUBMISSION}\")\nexcept Exception as e:\n    print(f\"Error reading sample submission file {SAMPLE_SUBMISSION}: {e}\")\n\n\n# Get all tomogram folders and slices\ndef get_data_paths():\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            print(f\"Found {len(train_tomogram_folders)} training tomogram folders.\")\n            print(f\"Found {len(train_slices)} total training 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            print(f\"Found {len(test_tomogram_folders)} test tomogram folders.\")\n            print(f\"Found {len(test_slices)} total test 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\n\ntrain_tomogram_folders, test_tomogram_folders, train_slices, test_slices = get_data_paths()\n# print(f\"Number of training tomograms: {len(train_tomogram_folders)}\") # Redundant with prints inside function\n# print(f\"Number of test tomograms: {len(test_tomogram_folders)}\")\n# print(f\"Total number of training slices: {len(train_slices)}\")\n# print(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_example = list(train_tomogram_folders[0].glob(\"*.jpg\"))\n    if slices_example:\n        print(f\"Example slice filename: {slices_example[0].name}\")\n\n\n# Extract tomogram ID and slice index from file path\ndef extract_tomo_slice_info(file_path):\n    # Ensure file_path is a Path object\n    file_path = Path(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_str = filename.split('_')[-1]  # Take the part after the last underscore\n            slice_idx = int(slice_idx_str)\n        else:\n            # Fallback if filename format is different (e.g., only digits)\n            slice_idx_str = ''.join(filter(str.isdigit, filename))  # Removed the backslash here\n            if slice_idx_str:\n                slice_idx = int(slice_idx_str)\n            else:\n                raise ValueError(\"Could not extract slice index from filename\")\n    except (ValueError, IndexError, TypeError) as e:\n        print(f\"Error extracting slice index from {file_path}: {e}. Using default 0.\")\n        slice_idx = 0  # Default value\n\n    return tomo_id, slice_idx\n\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 for preprocessing.\")\n        return {}\n\n    # Check columns and adapt accordingly\n    columns = labels_df.columns.tolist()\n\n    # Try to find the right column names based on common patterns\n    tomo_col = next((col for col in columns if 'tomo' in col.lower()), None)\n    slice_col = next((col for col in columns if 'row' in col.lower() or 'slice' in col.lower()), None)\n    x_col = next((col for col in columns if 'axis 0' in col.lower() or col.lower() == 'x'), None)\n    y_col = next((col for col in columns if 'axis 1' in col.lower() or col.lower() == 'y'), None)\n\n    # Use default column names if not found and issue warnings\n    if tomo_col is None:\n        tomo_col = 'tomo_id'\n        print(f\"Warning: Tomogram ID column not found, defaulting to '{tomo_col}'\")\n    if slice_col is None:\n        slice_col = 'row_id'\n        print(f\"Warning: Slice index column not found, defaulting to '{slice_col}'\")\n    if x_col is None:\n        x_col = 'Motor axis 0'\n        print(f\"Warning: X coordinate column not found, defaulting to '{x_col}'\")\n    if y_col is None:\n        y_col = 'Motor axis 1'\n        print(f\"Warning: Y coordinate column not found, defaulting to '{y_col}'\")\n\n    # Verify that the chosen columns exist in the DataFrame\n    required_cols = [tomo_col, slice_col, x_col, y_col]\n    if not all(col in labels_df.columns for col in required_cols):\n        print(\n            f\"Error: One or more required columns {required_cols} not found in labels DataFrame. Columns available: {labels_df.columns}\")\n        return {}\n\n    print(f\"Using label columns: Tomo='{tomo_col}', Slice='{slice_col}', X='{x_col}', Y='{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            slice_val = row[slice_col]\n            if isinstance(slice_val, (int, float)) and not np.isnan(slice_val):\n                slice_idx = int(slice_val)\n            elif isinstance(slice_val, str):\n                # Try extracting numeric part if not a clean integer string\n                slice_idx_str = ''.join(filter(str.isdigit, slice_val))\n                if slice_idx_str:\n                    slice_idx = int(slice_idx_str)\n                else:\n                    print(\n                        f\"Warning: Could not parse slice index from value '{slice_val}' for tomo '{tomo_id}'. Skipping row.\")\n                    continue\n            else:\n                print(\n                    f\"Warning: Unexpected type for slice index '{slice_val}' (type: {type(slice_val)}) for tomo '{tomo_id}'. Skipping row.\")\n                continue\n\n            # Extract coordinates\n            x_val = row[x_col]\n            y_val = row[y_col]\n            if pd.isna(x_val) or pd.isna(y_val):\n                print(f\"Warning: NaN coordinate found for tomo '{tomo_id}', slice {slice_idx}. Skipping point.\")\n                continue\n            x = float(x_val)\n            y = float(y_val)\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 (ValueError, TypeError) as e:\n            print(f\"Error processing row: {row.to_dict()} -> {e}. Skipping row.\")\n        except Exception as e:\n            print(f\"Unexpected error processing row: {row.to_dict()} -> {e}. Skipping row.\")\n\n    return labels_dict\n\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 unique slices with labels processed: {len(labels_dict)}\")\n    # Show examples of first few labels\n    if labels_dict:\n        count = 0\n        total_points = 0\n        for key, points in labels_dict.items():\n            if count < 3:\n                print(f\"Example Label - Slice {key}: {len(points)} marked points: {points[:3]}...\")\n            count += 1\n            total_points += len(points)\n        print(f\"Total labeled points across all slices: {total_points}\")\nelse:\n    print(\"Warning: No label data (train_labels_df) loaded. labels_dict will be empty.\")\n    labels_dict = {}\n\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: Target heatmap shape (height, width)\n    - points: List of coordinates [(x1, y1), (x2, y2), ...] in original image space\n    - sigma: Standard deviation of Gaussian kernel in heatmap space\n\n    Returns:\n    - Heatmap numpy array (float32)\n    \"\"\"\n    height, width = img_shape  # Expecting (height, width)\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 heatmap\n    y_grid, x_grid = np.mgrid[0:height, 0:width]\n\n    for x, y in points:\n        # Coordinates (x,y) need to be mapped to the heatmap space if different from original\n        # Assuming for now points are already scaled/relative to heatmap size if needed\n        # Ensure coordinates are valid for the heatmap grid\n        # Note: OpenCV/numpy uses (row, col) which is (y, x)\n        int_x, int_y = int(round(x)), int(round(y))\n\n        # Check bounds carefully (use heatmap shape)\n        if 0 <= int_x < width and 0 <= int_y < height:\n            # Compute Gaussian values centered at the point (use float coords for center)\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\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, target_size=(256, 256)):\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 = target_size  # Store target size (height, width)\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 using PIL (handles paths better, integrates with transforms)\n            img_pil = Image.open(img_path).convert('L')\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 (input for model)\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)  # Might need resize here too if transform is None\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                    'original_size': original_size[::-1]  # Return as (height, width)\n                }\n\n            # --- Training/Validation specific part ---\n            key = (tomo_id, slice_idx)\n            points_original = self.labels_dict.get(key, [])  # Points in original image coords\n\n            # Scale points to the target heatmap size (which matches the transformed image size)\n            # Original image size: original_size (width, height)\n            # Target heatmap size: self.target_size (height, width)\n            points_scaled = []\n            if points_original:\n                orig_w, orig_h = original_size\n                target_h, target_w = self.target_size\n                if orig_w > 0 and orig_h > 0:  # Avoid division by zero\n                    scale_x = target_w / orig_w\n                    scale_y = target_h / orig_h\n                    points_scaled = [(p[0] * scale_x, p[1] * scale_y) for p in points_original]\n                else:\n                    print(f\"Warning: Original image size is zero for {img_path}. Cannot scale points.\")\n\n            # Create heatmap label using scaled points and target size\n            # Pass target size (height, width) to create_heatmap\n            heatmap = create_heatmap(self.target_size, points_scaled, sigma=5)  # Adjust sigma as needed for target size\n\n            # Convert heatmap to tensor\n            heatmap_tensor = torch.tensor(heatmap, dtype=torch.float32).unsqueeze(0)\n\n            return {\n                'image': img_transformed,\n                'heatmap': heatmap_tensor,\n                'points_original': points_original,  # Keep original points if needed for eval\n                'tomo_id': tomo_id,\n                'slice_idx': slice_idx,\n                'image_path': str(img_path),\n                'original_size': original_size[::-1]  # Return as (height, width)\n            }\n\n        except FileNotFoundError:\n            print(f\"Error: Image file not found: {img_path}\")\n            # Return None or handle appropriately based on collate_fn\n            return None  # Collate fn needs to handle None\n        except Exception as e:\n            print(f\"Error processing item {idx}, path {img_path}: {e}\")\n            # Return None or a default item to avoid breaking the DataLoader\n            # Collate fn needs to handle None\n            return None\n\n\n# FIX 9: Simplified transforms with fixed size\ntarget_height, target_width = 256, 256\ntrain_transform = transforms.Compose([\n    transforms.Resize((target_height, target_width)),\n    # Add Augmentations if desired (e.g., RandomHorizontalFlip)\n    # transforms.RandomHorizontalFlip(),\n    transforms.ToTensor(),\n    # Add Normalization if desired (calculate mean/std from dataset or use imagenet defaults)\n    # transforms.Normalize(mean=[0.5], std=[0.5]) # Example for grayscale\n])\n\nval_transform = transforms.Compose([\n    transforms.Resize((target_height, target_width)),\n    transforms.ToTensor(),\n    # transforms.Normalize(mean=[0.5], std=[0.5]) # Use same normalization as training\n])\n\n\n# FIX 10: Simplified and more robust collate function (Revised)\ndef custom_collate_fn(batch):\n    \"\"\"\n    Collate function that handles potential None items in the batch resulted from dataset errors,\n    and correctly collates varying length lists/tuples.\n    \"\"\"\n    # Filter out None items\n    batch = [item for item in batch if item is not None]\n\n    if not batch:\n        # Return an empty dictionary with expected keys if batch is empty after filtering\n        print(\"Warning: Empty batch encountered in collate_fn\")\n        # IMPORTANT: Make sure to return all possible keys a batch might have\n        return {\n            'image': torch.Tensor(0),\n            'heatmap': torch.Tensor(0),  # Include heatmap, even if empty\n            'points_original': [],\n            'tomo_id': [],\n            'slice_idx': [],\n            'image_path': [],\n            'original_size': []\n        }\n\n    # Manually collate items\n    collated_batch = {}\n\n    # Stack image and heatmap tensors\n    collated_batch['image'] = torch.stack([item['image'] for item in batch])\n\n    # Conditionally add heatmap and points_original if present (for training/validation)\n    if 'heatmap' in batch[0]:\n        collated_batch['heatmap'] = torch.stack([item['heatmap'] for item in batch])\n    else:\n        # If heatmap is not present, it's likely a test batch, so don't add it to the collated batch\n        # Or, if you need this key always, you could return an empty tensor of appropriate shape.\n        # For prediction, you typically don't need the ground truth heatmap.\n        pass  # No need to add this key if it's not meant to be there for test\n\n    if 'points_original' in batch[0]:  # ADD THIS CONDITIONAL CHECK\n        collated_batch['points_original'] = [item['points_original'] for item in batch]  # This will be a list of lists\n    else:\n        # For test batches, points_original is not needed.\\\n        # You can still add an empty list for consistency if downstream code expects the key.\n        collated_batch['points_original'] = []\n\n    # Collect other items into lists (they don't need stacking into tensors)\n    collated_batch['tomo_id'] = [item['tomo_id'] for item in batch]  # List of strings\n    collated_batch['slice_idx'] = [item['slice_idx'] for item in batch]  # List of integers/tensors\n    collated_batch['image_path'] = [item['image_path'] for item in batch]  # List of strings\n    collated_batch['original_size'] = [item['original_size'] for item in batch]  # List of tuples\n\n    return collated_batch\n\n\n# FIX 11: Improved UNet architecture with better initialization\nclass DoubleConv(nn.Module):\n    def __init__(self, in_channels, out_channels, mid_channels=None):\n        super().__init__()\n        if not mid_channels:\n            mid_channels = out_channels\n        self.double_conv = nn.Sequential(\n            nn.Conv2d(in_channels, mid_channels, kernel_size=3, padding=1, bias=False),  # Bias=False with BN\n            nn.BatchNorm2d(mid_channels), \\\n            nn.ReLU(inplace=True), \\\n            nn.Conv2d(mid_channels, out_channels, kernel_size=3, padding=1, bias=False),  # Bias=False with BN\n            nn.BatchNorm2d(out_channels), \\\n            nn.ReLU(inplace=True)\n        )\n        # Initialize weights (optional but often good practice)\n        self._initialize_weights()\n\n    def _initialize_weights(self):\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\n\nclass Down(nn.Module):\n    \"\"\"Downscaling with maxpool then double conv\"\"\"\n\n    def __init__(self, in_channels, out_channels):\n        super().__init__()\n        self.maxpool_conv = nn.Sequential(\n            nn.MaxPool2d(2),\n            DoubleConv(in_channels, out_channels)\n        )\n\n    def forward(self, x):\n        return self.maxpool_conv(x)\n\n\nclass Up(nn.Module):\n    \"\"\"Upscaling then double conv\"\"\"\n\n    def __init__(self, in_channels, out_channels, bilinear=True):\n        super().__init__()\n\n        # if bilinear, use the normal convolutions to reduce the number of channels\n        if bilinear:\n            self.up = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)\n            self.conv = DoubleConv(in_channels, out_channels, in_channels // 2)\n        else:\n            self.up = nn.ConvTranspose2d(in_channels, in_channels // 2, kernel_size=2, stride=2)\n            self.conv = DoubleConv(in_channels, out_channels)\n\n    def forward(self, x1, x2):\n        x1 = self.up(x1)\n        # input is CHW\n        diffY = x2.size()[2] - x1.size()[2]\n        diffX = x2.size()[3] - x1.size()[3]\n\n        x1 = F.pad(x1, [diffX // 2, diffX - diffX // 2,\n                        diffY // 2, diffY - diffY // 2])\n        # if you have padding issues, see\n        # https://github.com/HaiyongJiang/U-Net-Pytorch-Unstructured-Buggy/commit/0e854509c2cea854e247a9c615f175176fdd0f70\n        # https://github.com/xiaopeng-liao/Pytorch-UNet/commit/8ebac70e633bac59fc22bb5195e513d5832fb3bd\n        x = torch.cat([x2, x1], dim=1)\n        return self.conv(x)\n\n\nclass OutConv(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super(OutConv, self).__init__()\n        self.conv = nn.Conv2d(in_channels, out_channels, kernel_size=1)\n        # Initialize weights (optional)\n        nn.init.kaiming_normal_(self.conv.weight, mode='fan_out', nonlinearity='relu')\n        if self.conv.bias is not None:\n            nn.init.constant_(self.conv.bias, 0)\n\n    def forward(self, x):\n        return self.conv(x)\n\n\nclass UNet(nn.Module):\n    def __init__(self, n_channels=1, n_classes=1, bilinear=False, init_features=32):  # Added init_features\n        super(UNet, self).__init__()\n        self.n_channels = n_channels\n        self.n_classes = n_classes\n        self.bilinear = bilinear\n        factor = 2 if bilinear else 1\n\n        # Use init_features to control the channel sizes\n        f = init_features\n        self.inc = DoubleConv(n_channels, f)\n        self.down1 = Down(f, f * 2)\n        self.down2 = Down(f * 2, f * 4)\n        self.down3 = Down(f * 4, f * 8)\n        self.down4 = Down(f * 8, f * 16 // factor)  # Adjusted for bilinear option\n        self.up1 = Up(f * 16, f * 8 // factor, bilinear)  # Adjusted for bilinear option\n        self.up2 = Up(f * 8, f * 4 // factor, bilinear)  # Adjusted for bilinear option\n        self.up3 = Up(f * 4, f * 2 // factor, bilinear)  # Adjusted for bilinear option\n        self.up4 = Up(f * 2, f, bilinear)\n        self.outc = OutConv(f, n_classes)\n\n    def forward(self, x):\n        # FIX 12: Add input shape check\n        if x.dim() != 4 or x.shape[1] != self.n_channels:\n            raise ValueError(f\"Expected 4D tensor N-{self.n_channels}-H-W, got {x.shape}\")\n\n        x1 = self.inc(x)\n        x2 = self.down1(x1)\n        x3 = self.down2(x2)\n        x4 = self.down3(x3)\n        x5 = self.down4(x4)\n        x = self.up1(x5, x4)\n        x = self.up2(x, x3)\n        x = self.up3(x, x2)\n        x = self.up4(x, x1)\n        logits = self.outc(x)  # Output logits directly\n\n        # REMOVED torch.sigmoid - Use BCEWithLogitsLoss instead\n        # output = torch.sigmoid(logits)\n\n        return logits  # Return logits\n\n\n# FIX 14: Training function with proper error handling and GPU memory optimization\n# UPDATED val_loader_to_use logic\ndef train_model(model, train_loader, val_loader, criterion, optimizer, scheduler, num_epochs=10, accelerator_type='cpu',\n                device=torch.device('cpu')):\n    best_val_loss = float('inf')\n    history = {'train_loss': [], 'val_loss': []}\n\n    # For TPU acceleration (Keep structure but might need specific imports if actually running on TPU)\n    is_tpu = accelerator_type == 'tpu'\n    if is_tpu:\n        try:\n            import torch_xla.core.xla_model as xm\n            import torch_xla.distributed.parallel_loader as pl\n            print(\"TPU modules imported successfully for training.\")\n        except ImportError as e:\n            print(f\"Warning: Error importing TPU modules, training might proceed on CPU/GPU if available: {e}\")\n            is_tpu = False  # Fallback if imports fail\n            # Update accelerator_type based on fallback? Or assume setup_device handled it?\n            # Re-check device if needed, though setup_device should be source of truth\n\n    # For mixed precision training on GPU\n    use_amp = False\n    scaler = None\n    if accelerator_type == 'gpu' and not is_tpu:  # Only use CUDA AMP if GPU is selected and not TPU\n        try:\n            from torch.cuda.amp import GradScaler, autocast\n            use_amp = True\n            scaler = GradScaler()\n            print(\"Mixed precision training (AMP) enabled for GPU.\")\n        except ImportError:\n            print(\"CUDA AMP (GradScaler, autocast) not available, using full precision.\")\n\n    for epoch in range(num_epochs):\n        print(f\"\\n--- Epoch {epoch + 1}/{num_epochs} ---\")\n        # Training phase\n        model.train()\n        train_loss = 0.0\n\n        # Choose appropriate dataloader based on accelerator\n        if is_tpu:\n            try:\n                train_device_loader = pl.ParallelLoader(train_loader, [device]).per_device_loader(device)\n                loader_to_use = train_device_loader\n                print(\"Using TPU ParallelLoader for training.\")\n            except Exception as e:\n                print(f\"Error setting up TPU dataloader for training: {e}. Using standard loader.\")\n                loader_to_use = train_loader\n        else:\n            loader_to_use = train_loader\n            print(\"Using standard DataLoader for training.\")\n\n        # Training loop with error handling\n        pbar_train = tqdm(loader_to_use, desc=f'Epoch {epoch + 1}/{num_epochs} [Train]', leave=False)\n        # batch_count = 0 # 这一行不再需要，因为我们将使用 enumerate 来获取 batch_idx\n        processed_items = 0\n        # --- FIX: 关键修改在这里 ---\n        for batch_idx, batch in enumerate(pbar_train):  # <--- 添加了 enumerate(pbar_train) 来获取 batch_idx\n            # --- 结束关键修改 ---\n            # Check if collate_fn returned None (due to dataset errors)\n            if batch is None:\n                print(\n                    f\"Warning: Skipping None batch in training loop (batch index approx {batch_idx}).\")  # 使用 batch_idx\n                continue\n\n            try:\n                images = batch['image'].to(device)\n                heatmaps = batch['heatmap'].to(device)\n\n                # Check for empty batch after potential filtering in collate_fn\n                if images.size(0) == 0:\n                    print(f\"Warning: Empty batch encountered at index {batch_idx}, skipping.\")  # 使用 batch_idx\n                    continue\n\n                current_batch_size = images.size(0)\n\n                # Forward pass with appropriate precision\n                if use_amp:\n                    with autocast():\n                        outputs = model(images)  # Expecting logits\n                        loss = criterion(outputs, heatmaps)  # criterion is BCEWithLogitsLoss\n                else:\n                    outputs = model(images)  # Expecting logits\n                    loss = criterion(outputs, heatmaps)  # criterion is BCEWithLogitsLoss\n\n                # Check for NaN loss\n                if torch.isnan(loss):\n                    print(f\"Warning: NaN loss encountered at batch {batch_idx}. Skipping batch.\")  # 使用 batch_idx\n                    # Optionally zero gradients here if optimizer state might be affected\n                    optimizer.zero_grad()\n                    continue\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 is_tpu:\n                    loss.backward()\n                    xm.optimizer_step(optimizer, barrier=True)  # Barrier=True ensures step completion\n                else:  # Standard CPU/GPU without AMP\n                    loss.backward()\n                    optimizer.step()\n\n                # Update statistics\n                train_loss += loss.item() * current_batch_size  # Use actual batch size\n                processed_items += current_batch_size\n\n                # Update progress bar\n                pbar_train.set_postfix(loss=f\"{loss.item():.4f}\")\n\n                # FIX 15: Periodic GPU memory cleanup (only if using GPU)\n                if accelerator_type == 'gpu' and batch_idx % 50 == 0:  # 这里的 batch_idx 现在被正确定义了\n                    torch.cuda.empty_cache()\n\n            except Exception as e:\n                # Log specific error for the batch and continue\n                print(f\"\\nError during training batch {batch_idx}: {e}\")  # 使用 batch_idx\n                # Consider adding traceback print here for debugging:\n                # import traceback\n                # traceback.print_exc()\n                continue  # Continue to the next batch\n\n        # Calculate average loss for the epoch\n        if processed_items > 0:\n            avg_train_loss = train_loss / processed_items\n        else:\n            avg_train_loss = 0.0\n            print(\"Warning: No items processed in training epoch.\")\n\n        # Validation phase\n        model.eval()\n        val_loss = 0.0\n        val_processed_items = 0\n\n        # FIX: Choose appropriate validation dataloader (Corrected Logic)\n        if is_tpu:\n            try:\n                # Make sure pl is available from the TPU check earlier\n                val_device_loader = pl.ParallelLoader(val_loader, [device]).per_device_loader(device)\n                val_loader_to_use = val_device_loader\n                print(\"Using TPU ParallelLoader for validation.\")\n            except Exception as e:\n                print(f\"Error setting up TPU validation dataloader: {e}. Using standard loader.\")\n                val_loader_to_use = val_loader  # Fallback for TPU error\n        else:\n            val_loader_to_use = val_loader  # Defined for CPU/GPU\n            print(\"Using standard DataLoader for validation.\")\n\n        pbar_val = tqdm(val_loader_to_use, desc=f'Epoch {epoch + 1}/{num_epochs} [Val]', leave=False)\n        val_batch_count = 0\n        with torch.no_grad():  # Ensure no gradients are computed\n            for batch in pbar_val:  # 验证循环中没有使用 batch_idx，所以无需修改\n                # Check if collate_fn returned None\n                if batch is None:\n                    print(f\"Warning: Skipping None batch in validation loop (batch index approx {val_batch_count}).\")\n                    val_batch_count += 1\n                    continue\n\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(f\"Warning: Empty validation batch encountered at index {val_batch_count}, skipping.\")\n                        val_batch_count += 1\n                        continue\n\n                    current_batch_size = images.size(0)\n\n                    # Forward pass (no autocast needed for validation with no_grad usually, but doesn't hurt)\n                    if use_amp:\n                        with autocast():\n                            outputs = model(images)  # Expecting logits\n                            loss = criterion(outputs, heatmaps)  # criterion is BCEWithLogitsLoss\n                    else:\n                        outputs = model(images)  # Expecting logits\n                        loss = criterion(outputs, heatmaps)  # criterion is BCEWithLogitsLoss\n\n                    # Check for NaN loss\n                    if torch.isnan(loss):\n                        print(\n                            f\"Warning: NaN loss encountered during validation batch {val_batch_count}. Skipping loss accumulation for this batch.\")\n                        val_batch_count += 1\n                        continue\n\n                    # Update statistics\n                    val_loss += loss.item() * current_batch_size\n                    val_processed_items += current_batch_size\n\n                    # Update progress bar\n                    pbar_val.set_postfix(loss=f\"{loss.item():.4f}\")\n\n                    val_batch_count += 1\n\n                except Exception as e:\n                    print(f\"\\nError during validation batch {val_batch_count}: {e}\")\n                    # import traceback\n                    # traceback.print_exc()\n                    val_batch_count += 1\n                    continue  # Continue to the next batch\n\n        # Calculate average validation loss\n        if val_processed_items > 0:\n            avg_val_loss = val_loss / val_processed_items\n        else:\n            avg_val_loss = float('inf')  # Or handle as error / zero?\n            print(\"Warning: No items processed in validation epoch.\")\n\n        # Print epoch results\n        print(f'Epoch {epoch + 1}/{num_epochs}: \\n' \\\n              f'\\tTrain Loss: {avg_train_loss:.4f}\\n' \\\n              f'\\tVal Loss:   {avg_val_loss:.4f}')\n\n        # Update learning rate scheduler based on validation loss\n        current_lr = optimizer.param_groups[0]['lr']\n        if scheduler is not None:\n            if isinstance(scheduler, ReduceLROnPlateau):\n                scheduler.step(avg_val_loss)\n            else:\n                scheduler.step()\n        new_lr = optimizer.param_groups[0]['lr']\n        if new_lr != current_lr:\n            print(f\"\\tLearning rate reduced to {new_lr:.6f}\")\n\n        # Save model if validation loss improved\n        if avg_val_loss < best_val_loss:\n            print(f\"\\tValidation loss improved from {best_val_loss:.4f} to {avg_val_loss:.4f}. Saving model...\")\n            best_val_loss = avg_val_loss\n            save_path = Path(\"./best_model.pth\")  # Save in working directory\n            try:\n                model_save = model  # Default save\n                # Handle TPU model saving if applicable (might need state dict conversion)\n                if is_tpu:\n                    # TPU-specific saving often involves getting state_dict\n                    # xm.save(model.state_dict(), str(save_path)) might be needed\n                    # Or save the CPU version if model was wrapped\n                    torch.save(model.state_dict(), save_path)  # Saving state_dict is generally safer\n                    print(f\"\\tTPU Model state_dict saved to {save_path}\")\n                else:\n                    torch.save(model.state_dict(), save_path)  # Save state_dict for flexibility\n                    print(f\"\\tModel state_dict saved to {save_path}\")\n            except Exception as e:\n                print(f\"Error saving model: {e}\")\n        else:\n            print(f\"\\tValidation loss did not improve from {best_val_loss:.4f}.\")\n\n        # Update history\n        history['train_loss'].append(avg_train_loss)\n        history['val_loss'].append(avg_val_loss)\n\n        # Plot training progress periodically\n        if (epoch + 1) % 5 == 0 or epoch == num_epochs - 1:\n            try:\n                plt.figure(figsize=(10, 5))\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.title('Training and Validation Loss')\n                plt.legend()\n                plt.grid(True)\n                plot_save_path = Path(f\"./training_progress_epoch_{epoch + 1}.png\")\n                plt.savefig(plot_save_path)\n                print(f\"\\tTraining progress plot saved to {plot_save_path}\")\n                plt.close()  # Close plot to free memory\n            except Exception as e:\n                print(f\"Error plotting training progress: {e}\")\n\n    print(\"\\nTraining finished.\")\n    return model, history\n\n\n# FIX 16: Enhanced post-processing function to detect peaks in the heatmap\ndef detect_points_from_heatmap(heatmap, threshold=0.5, min_distance=10, original_size=None, current_size=None):\n    \"\"\"\n    Detect points from heatmap by finding local maxima above a threshold.\n\n    Parameters:\n    - heatmap: Predicted heatmap tensor (C, H, W) or numpy array (H, W)\n    - threshold: Minimum value to consider as a potential motor peak\n    - min_distance: Minimum distance between detected peaks in heatmap pixel coordinates\n    - original_size: Tuple (height, width) of the original image for scaling back\n    - current_size: Tuple (height, width) of the heatmap if scaling needed\n\n    Returns:\n    - List of (x, y) coordinates of detected motors in ORIGINAL image space\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, ensure it's on CPU\n    if isinstance(heatmap, torch.Tensor):\n        heatmap = heatmap.squeeze().cpu().numpy()\n\n    # Ensure heatmap is 2D\n    if heatmap.ndim != 2:\n        print(f\"Warning: Expected 2D heatmap, got shape {heatmap.shape}. Attempting to squeeze.\")\n        heatmap = np.squeeze(heatmap)\n        if heatmap.ndim != 2:\n            print(\"Error: Could not convert heatmap to 2D.\")\n            return []\n\n    # Check if current_size is provided, otherwise use heatmap's shape\n    if current_size is None:\n        current_size = heatmap.shape  # (height, width)\n    current_h, current_w = current_size\n\n    # Apply Gaussian filter to smooth the heatmap (helps find stable peaks)\n    # Sigma=1 is often reasonable for peak detection\n    heatmap_smoothed = gaussian_filter(heatmap, sigma=1)\n\n    # Find local maxima using maximum_filter\n    # Create a footprint for the neighborhood (e.g., 3x3)\n    neighborhood = generate_binary_structure(2, 2)  # 8-connectivity\n    local_max = maximum_filter(heatmap_smoothed, footprint=neighborhood) == heatmap_smoothed\n\n    # Apply threshold: Only consider maxima above the threshold\n    detected_peaks = (heatmap_smoothed > threshold) & local_max\n\n    # Extract coordinates of peaks\n    y_indices, x_indices = np.where(detected_peaks)\n\n    # Get heatmap values (intensities) at peaks\n    intensities = heatmap_smoothed[y_indices, x_indices]\n\n    # Combine coordinates and intensities\n    points_with_intensities = list(zip(x_indices, y_indices, intensities))\n\n    # Sort points by intensity (highest to lowest) - helps in non-maximum suppression step\n    points_with_intensities.sort(key=lambda p: p[2], reverse=True)\n\n    # Non-maximum suppression based on min_distance\n    # Keep track of points to include in the final list\n    final_points_heatmap = []\n    suppressed = np.zeros(len(points_with_intensities), dtype=bool)\n\n    for i in range(len(points_with_intensities)):\n        if suppressed[i]:\n            continue  # Skip if already suppressed\n\n        # Add this point (it's the strongest in its neighborhood so far)\n        xi, yi, _ = points_with_intensities[i]\n        final_points_heatmap.append((xi, yi))\n\n        # Suppress other points within min_distance\n        for j in range(i + 1, len(points_with_intensities)):\n            if suppressed[j]:\n                continue\n            xj, yj, _ = points_with_intensities[j]\n            dist = np.sqrt((xi - xj) ** 2 + (yi - yj) ** 2)\n            if dist < min_distance:\n                suppressed[j] = True\n\n    # --- Scaling back to original image size ---\n    final_points_original = []\n    if original_size is not None:\n        orig_h, orig_w = original_size  # (height, width)\n        if current_w > 0 and current_h > 0:  # Check for division by zero\n            scale_x = orig_w / current_w\n            scale_y = orig_h / current_h\n            # Scale and round to nearest integer coordinates\n            final_points_original = [(int(round(x * scale_x)), int(round(y * scale_y))) for x, y in\n                                     final_points_heatmap]\n        else:\n            print(\"Warning: Current heatmap size is zero. Cannot scale points back.\")\n            final_points_original = final_points_heatmap  # Return heatmap coords if scaling fails\n    else:\n        # If original_size not provided, return points in heatmap coordinates\n        final_points_original = final_points_heatmap\n\n    return final_points_original\n\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 based on distance matching.\n\n    Parameters:\n    - true_points: List of ground truth points [(x1, y1), (x2, y2), ...] in original coordinates\n    - pred_points: List of predicted points [(x1, y1), (x1, y1), ...] in original coordinates\n    - distance_threshold: Maximum distance (in pixels) to consider a prediction a true positive match\n\n    Returns:\n    - Dictionary with precision, recall, F1, TP, FP, FN counts\n    \"\"\"\n    # Handle edge cases with empty lists\n    if not true_points and not pred_points:\n        return {'precision': 1.0, 'recall': 1.0, 'f1': 1.0, 'tp': 0, 'fp': 0, 'fn': 0}\n    elif not true_points:  # Only predictions (all False Positives)\n        return {'precision': 0.0, 'recall': 0.0, 'f1': 0.0, 'tp': 0, 'fp': len(pred_points), 'fn': 0}\n    elif not pred_points:  # Only ground truth (all False Negatives)\n        return {'precision': 0.0, 'recall': 0.0, 'f1': 0.0, 'tp': 0, 'fp': 0, 'fn': len(true_points)}\n\n    # Convert to numpy arrays for efficient distance calculation\n    true_points_arr = np.array(true_points)\n    pred_points_arr = np.array(pred_points)\n\n    num_true = len(true_points_arr)\n    num_pred = len(pred_points_arr)\n\n    # Calculate pairwise distances\n    # Using scipy's cdist is efficient for this\n    from scipy.spatial.distance import cdist\n    distances = cdist(pred_points_arr, true_points_arr)  # Shape: (num_pred, num_true)\n\n    # Matching using Hungarian algorithm (or greedy matching)\n    # Greedy matching: Iterate through predictions, find closest valid true point, mark both as matched.\n\n    true_matched = np.zeros(num_true, dtype=bool)\n    pred_matched_indices = -np.ones(num_pred,\n                                    dtype=int)  # Store index of matched true point for each pred, -1 if no match\n    true_positives = 0\n\n    # Iterate through predictions (can sort by confidence if available, but not here)\n    for i in range(num_pred):\n        pred_point = pred_points_arr[i]\n        min_dist = float('inf')\n        best_true_idx = -1\n\n        # Find the closest *unmatched* true point within the threshold\n        for j in range(num_true):\n            if not true_matched[j]:  # Only consider unmatched true points\n                dist = distances[i, j]\n                if dist < min_dist and dist <= distance_threshold:\n                    min_dist = dist\n                    best_true_idx = j\n\n        # If a valid match is found\n        if best_true_idx != -1:\n            true_positives += 1\n            true_matched[best_true_idx] = True  # Mark true point as matched\n            pred_matched_indices[i] = best_true_idx  # Mark prediction as matched (implicitly via TP count)\n\n    # Calculate metrics\n    false_positives = num_pred - true_positives\n    false_negatives = num_true - true_positives\n\n    precision = true_positives / (true_positives + false_positives) if (true_positives + false_positives) > 0 else 0.0\n    recall = true_positives / (true_positives + false_negatives) if (\n                                                                                true_positives + false_negatives) > 0 else 0.0  # num_true is denominator here\n    f1 = 2 * (precision * recall) / (precision + recall) if (precision + recall) > 0 else 0.0\n\n    return {\n        'precision': precision,\n        'recall': recall,\n        'f1': f1,\n        'tp': int(true_positives),\n        'fp': int(false_positives),\n        'fn': int(false_negatives)\n    }\n\n\n# FIX 18: Improved evaluation function (renamed predict_and_evaluate)\ndef predict_and_evaluate(model, data_loader, device, detection_threshold=0.5, distance_threshold=20, is_test_set=False):\n    \"\"\"\n    Run model predictions on data_loader and optionally evaluate if labels are available.\n\n    Parameters:\n    - model: Trained UNet model\n    - data_loader: DataLoader for test or validation data\n    - device: Device to run evaluation on\n    - detection_threshold: Threshold for peak detection from heatmap\n    - distance_threshold: Distance threshold for matching predictions to ground truth\n    - is_test_set: Boolean indicating if this is the test set (no ground truth)\n\n    Returns:\n    - all_predictions: List of dictionaries containing prediction info for each image.\n    - avg_metrics: Dictionary with average evaluation metrics (if not is_test_set).\\\n    \"\"\"\n    model.eval()\n    all_predictions_output = []  # Store predictions in submission format structure\n    all_metrics_list = []\n\n    with torch.no_grad():\n        pbar = tqdm(data_loader, desc=\"Predicting/Evaluating\", leave=False)\n        batch_count = 0\n        for batch in pbar:\n            # Handle potential None batches from dataset/collate errors\n            if batch is None:\n                print(f\"Warning: Skipping None batch in predict_and_evaluate (batch index approx {batch_count}).\")\n                batch_count += 1\n                continue\n\n            try:\n                images = batch['image'].to(device)\n                tomo_ids = batch['tomo_id']\n                slice_idxs = batch['slice_idx']\n                original_sizes = batch['original_size']  # Expecting N x (H, W)\n\n                # Skip empty batches\n                if images.size(0) == 0:\n                    print(f\"Warning: Empty batch encountered at index {batch_count} in predict_and_evaluate, skipping.\")\n                    batch_count += 1\n                    continue\n\n                # Forward pass - get logits\n                logits = model(images)\n                # Apply sigmoid to get probabilities [0, 1] for heatmap\n                heatmaps = torch.sigmoid(logits)\n\n                # Process each image in the batch\n                for i in range(images.size(0)):\n                    single_heatmap = heatmaps[i]  # Shape (1, H, W) or (H, W)\n                    tomo_id = tomo_ids[i]\n                    # --- FIX: 关键修改在这里 ---\n                    slice_idx = slice_idxs[i]  # <--- 移除了 .item()\n                    # --- 结束关键修改 ---\n                    orig_size_hw = original_sizes[i] # original_sizes[i] 已经是一个形如 (H, W) 的元组\n                    current_heatmap_size = tuple(single_heatmap.shape[-2:])  # (H, W) of heatmap\n\n                    # Detect points from heatmap, scale back to original image size\n                    pred_points_orig = detect_points_from_heatmap(\n                        heatmap=single_heatmap,\n                        threshold=detection_threshold,\n                        min_distance=10,  # Min distance in heatmap pixel space\n                        original_size=orig_size_hw,  # Target original image size (H, W)\n                        current_size=current_heatmap_size  # Current heatmap size (H, W)\n                    )\n\n                    # Store predictions\n                    # Format ready for prepare_submission\n                    for x, y in pred_points_orig:\n                        all_predictions_output.append({\n                            'tomo_id': tomo_id,\n                            'row_id': slice_idx,  # Assuming slice_idx corresponds to row_id\n                            'Motor axis 0': x,\n                            'Motor axis 1': y,\n                            # Optional: add confidence score if needed\n                        })\n\n                    # --- Evaluation part (if not test set) ---\n                    if not is_test_set:\n                        if 'points_original' in batch:\n                            # Ensure points_original is handled correctly (might be list of tensors/lists)\n                            # This expects a list of (x, y) tuples for the i-th item\n                            true_points_orig = batch['points_original'][i]\n                            # If points_original is stored differently (e.g., padded tensor), adjust access\n\n                            metrics = calculate_metrics(\n                                true_points=true_points_orig,\n                                pred_points=pred_points_orig,\n                                distance_threshold=distance_threshold\n                            )\n                            all_metrics_list.append(metrics)\n                        else:\n                            print(f\"Warning: 'points_original' key not found in batch for evaluation.\")\n\n                batch_count += 1\n\n            except Exception as e:\n                print(f\"\\nError processing batch {batch_count} in predict_and_evaluate: {e}\")\n                # import traceback\n                # traceback.print_exc()\n                batch_count += 1\n                continue  # Continue to next batch\n\n    # Calculate average metrics if evaluation was performed\n    avg_metrics = {}\n    if not is_test_set and all_metrics_list:\n        # Sum up TP, FP, FN across all images for micro-average, or average per-image scores for macro-average\n        # Example: Macro-average\n        avg_precision = np.mean([m['precision'] for m in all_metrics_list])\n        avg_recall = np.mean([m['recall'] for m in all_metrics_list])\n        avg_f1 = np.mean([m['f1'] for m in all_metrics_list])\n\n        avg_metrics = {\n            'precision': avg_precision,\n            'recall': avg_recall,\n            'f1': avg_f1\n        }\n        print(f\"\\nAverage Validation Metrics (Macro) - Precision: {avg_metrics['precision']:.4f}, \" \\\n              f\"Recall: {avg_metrics['recall']:.4f}, F1: {avg_metrics['f1']:.4f}\")\n    elif not is_test_set:\n        print(\"\\nNo metrics calculated (all_metrics_list is empty).\")\n\n    return all_predictions_output, avg_metrics\n\n\n# FIX 19: Function to prepare submission file (Revised - simpler, uses direct output from predict_and_evaluate)\ndef prepare_submission(predictions_list):\n    \"\"\"\n    Create a submission DataFrame from predictions.\n\n    Parameters:\n    - predictions_list: List of dictionaries, where each dict represents one detected point\n                      and has keys 'tomo_id', 'row_id', 'Motor axis 0', 'Motor axis 1'.\n                      This list is generated directly by predict_and_evaluate.\n\n    Returns:\n    - DataFrame in submission format. Returns an empty DataFrame with correct columns if predictions_list is empty.\n    \"\"\"\n    submission_columns = ['tomo_id', 'row_id', 'Motor axis 0', 'Motor axis 1']\n\n    if not predictions_list:\n        print(\"Warning: No predictions were generated. Creating empty submission file.\")\n        return pd.DataFrame(columns=submission_columns)\n\n    try:\n        submission_df = pd.DataFrame(predictions_list)\n\n        # Ensure correct columns and order\n        submission_df = submission_df[submission_columns]\n\n        # Optional: Convert coordinate columns to integers if required by competition\n        submission_df['Motor axis 0'] = submission_df['Motor axis 0'].round().astype(int)\n        submission_df['Motor axis 1'] = submission_df['Motor axis 1'].round().astype(int)\n        # Ensure row_id is integer\n        submission_df['row_id'] = submission_df['row_id'].astype(int)\n\n        return submission_df\n\n    except KeyError as e:\n        print(f\"Error preparing submission: Missing expected key {e} in predictions list.\")\n        print(\"Returning empty submission DataFrame.\")\n        return pd.DataFrame(columns=submission_columns)\n    except Exception as e:\n        print(f\"Unexpected error preparing submission: {e}\")\n        print(\"Returning empty submission DataFrame.\")\n        return pd.DataFrame(columns=submission_columns)\n\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 (coordinates in original image space).\n\n    Parameters:\n    - image_path: Path to image file\n    - true_points: List of ground truth points [(x1, y1), (x2, y2), ...]\\\n    - pred_points: List of predicted points [(x1, y1), (x2, y1), ...]\\\n    - save_path: Path object or string to save visualization. If None, display only.\n    \"\"\"\n    try:\n        # Read image using PIL\n        img = Image.open(image_path).convert('RGB')  # Convert to RGB for color circles\n        img_draw = np.array(img)  # Convert to numpy array for drawing\n\n        # Draw ground truth points (green circles)\n        if true_points:\n            for x, y in true_points:\n                # Ensure coordinates are within bounds\n                x, y = int(round(x)), int(round(y))\n                if 0 <= x < img_draw.shape[1] and 0 <= y < img_draw.shape[0]:\n                    cv2.circle(img_draw, (x, y), radius=7, color=(0, 255, 0), thickness=2)  # Green outline\n                    # cv2.circle(img_draw, (x, y), radius=1, color=(0, 255, 0), thickness=-1) # Small center dot\n\n        # Draw predicted points (red crosses or circles)\n        if pred_points:\n            for x, y in pred_points:\n                # Ensure coordinates are within bounds\n                x, y = int(round(x)), int(round(y))\n                if 0 <= x < img_draw.shape[1] and 0 <= y < img_draw.shape[0]:\n                    # Draw a cross\n                    cv2.line(img_draw, (x - 5, y), (x + 5, y), color=(255, 0, 0), thickness=2)  # Red horizontal\n                    cv2.line(img_draw, (x, y - 5), (x, y + 5), color=(255, 0, 0), thickness=2)  # Red vertical\n                    # Alternatively, draw circles\n                    # cv2.circle(img_draw, (x, y), radius=7, color=(255, 0, 0), thickness=2) # Red outline\n\n        # Create figure\n        plt.figure(figsize=(10, 10))\n        plt.imshow(img_draw)\n        plt.title(f\"Image: {Path(image_path).name}\")\n        plt.axis('off')  # Hide axes\n\n        # Add legend (simple version using plot handles)\n        legend_elements = []\n        if true_points:\n            legend_elements.append(\n                plt.Line2D([0], [0], marker='o', color='w', label='Ground Truth', markerfacecolor='g',\n                           markeredgecolor='g', markersize=10))\n        if pred_points:\n            legend_elements.append(\n                plt.Line2D([0], [0], marker='+', color='r', label='Prediction', linestyle='None', markersize=10,\n                           markeredgewidth=2))\n\n        if legend_elements:\n            plt.legend(handles=legend_elements, loc='upper right', bbox_to_anchor=(1.15, 1))\n\n        # Save or display\n        if save_path:\n            plt.savefig(str(save_path), dpi=150, bbox_inches='tight')  # Lower DPI for speed if needed\n            print(f\"Visualization saved to {save_path}\")\n            plt.close()  # Close the figure to free memory\n        else:\n            plt.show()\n\n    except FileNotFoundError:\n        print(f\"Error in visualization: Image not found at {image_path}\")\n    except Exception as e:\n        print(f\"Error during visualization for {image_path}: {e}\")\n        # import traceback\n        # traceback.print_exc()\n        if plt.gcf().get_axes():  # Close plot if it was created but failed\n            plt.close()\n\n\n# FIX 21: Main execution function (run_pipeline)\n# UPDATED loss function and prepare_submission call\ndef run_pipeline(train_ratio=0.8, batch_size=16, num_epochs=30, learning_rate=1e-4):\n    \"\"\"\n    Run the complete training, evaluation, and submission generation pipeline.\n\n    Parameters:\n    - train_ratio: Ratio of data to use for training vs validation.\n    - batch_size: Batch size for DataLoaders.\n    - num_epochs: Number of training epochs.\n    - learning_rate: Initial learning rate for the optimizer.\n    \"\"\"\n    global labels_dict  # Allow modification if needed, though preprocess happens earlier\n    global device, accelerator_type  # Use globally defined device/type\n    global train_slices, test_slices, SAMPLE_SUBMISSION  # Use global paths\n\n    try:\n        print(\"\\n--- Starting Pipeline ---\")\n        print(\"Setting up data loaders...\")\n\n        # Split training data into train/val sets\n        if not train_slices:\n            print(\"Error: No training slices found. Cannot proceed.\")\n            return None, None, None\n\n        train_paths, val_paths = train_test_split(\n            train_slices,\n            test_size=1 - train_ratio,\n            random_state=42\n        )\n        print(f\"Training images: {len(train_paths)}, Validation images: {len(val_paths)}\")\n\n        # --- Create Datasets ---\n        train_dataset = BacterialMotorDataset(\n            image_paths=train_paths,\n            labels_dict=labels_dict,\n            transform=train_transform,\n            target_size=(target_height, target_width)  # Pass target size\n        )\n\n        val_dataset = BacterialMotorDataset(\n            image_paths=val_paths,\n            labels_dict=labels_dict,\n            transform=val_transform,\n            target_size=(target_height, target_width)  # Pass target size\n        )\n\n        # If test slices exist, create test dataset\n        test_dataset = None\n        if test_slices:\n            test_dataset = BacterialMotorDataset(\n                image_paths=test_slices,\n                labels_dict=None,  # No labels for test set\n                transform=val_transform,\n                is_test=True,\n                target_size=(target_height, target_width)  # Pass target size\n            )\n            print(f\"Test images: {len(test_slices)}\")\n        else:\n            print(\n                \"Warning: No test slices found. Submission file will be based on predictions for validation set if needed, or empty.\")\n\n        # --- Create DataLoaders ---\n        num_workers = os.cpu_count() // 2 if os.cpu_count() else 2  # Use reasonable number of workers\n        print(f\"Using {num_workers} workers for DataLoaders.\")\n\n        train_loader = DataLoader(\n            train_dataset,\n            batch_size=batch_size,\n            shuffle=True,\n            num_workers=num_workers,\n            collate_fn=custom_collate_fn,\n            pin_memory=(accelerator_type == 'gpu'),  # Pin memory only if using GPU\n            drop_last=True  # Drop last incomplete batch during training\n        )\n\n        val_loader = DataLoader(\n            val_dataset,\n            batch_size=batch_size * 2,  # Use larger batch size for validation if memory allows\n            shuffle=False,\n            num_workers=num_workers,\n            collate_fn=custom_collate_fn,\n            pin_memory=(accelerator_type == 'gpu')\n        )\n\n        # If test slices exist, create test dataloader\n        test_loader = None\n        if test_dataset is not None:\n            test_loader = DataLoader(\n                test_dataset,\n                batch_size=batch_size * 2,  # Larger batch size for inference\n                shuffle=False,\n                num_workers=num_workers,\n                collate_fn=custom_collate_fn,\n                pin_memory=(accelerator_type == 'gpu')\n            )\n\n        # --- Setup Model, Loss, Optimizer, Scheduler ---\n        print(\"Setting up model, loss, optimizer, and scheduler...\")\n\n        # Model initialization (assuming UNet is defined above)\n        # n_channels=1 for grayscale images, n_classes=1 for heatmap output\n        model = UNet(n_channels=1, n_classes=1, bilinear=True, init_features=32).to(device)\n\n        # Loss function: BCEWithLogitsLoss is suitable for heatmap regression\n        # It combines sigmoid and BCELoss, which is numerically more stable\n        criterion = nn.BCEWithLogitsLoss()\n\n        # Optimizer\n        optimizer = Adam(model.parameters(), lr=learning_rate)\n\n        # Learning rate scheduler\n        # Reduces learning rate when validation loss stops improving\n        scheduler = ReduceLROnPlateau(optimizer, mode='min', factor=0.5, patience=5, verbose=True)\n\n        # --- Train Model ---\n        print(\"Starting training...\")\n        final_model, training_history = train_model(\n            model,\n            train_loader,\n            val_loader,\n            criterion,\n            optimizer,\n            scheduler,\n            num_epochs,\n            accelerator_type,\n            device\n        )\n\n        # --- Evaluate Model on Validation Set (if training was successful) ---\n        validation_metrics = {}\n        if final_model is not None and val_loader is not None:\n            print(\"\\n--- Evaluating on Validation Set ---\")\n            _, validation_metrics = predict_and_evaluate(\n                final_model,\n                val_loader,\n                device,\n                detection_threshold=0.5,  # Adjust as needed\n                distance_threshold=20,  # Adjust as needed\n                is_test_set=False\n            )\n        else:\n            print(\"\\nSkipping validation evaluation as model/loader is not available.\")\n\n        # --- Generate Submission File (if a test_loader exists) ---\n        predictions_for_submission = []\n        if test_loader is not None and final_model is not None:\n            print(\"\\n--- Generating Predictions for Submission ---\")\n            predictions_for_submission, _ = predict_and_evaluate(\n                final_model,\n                test_loader,\n                device,\n                detection_threshold=0.5,  # Use same threshold as validation or tune separately\n                distance_threshold=20,\n                # Distance threshold is not used for submission, but required by function signature\n                is_test_set=True\n            )\n\n            submission_df = prepare_submission(predictions_for_submission)\n\n            # Save submission file\n            output_submission_path = Path(\"./submission.csv\")\n            submission_df.to_csv(output_submission_path, index=False)\n            print(f\"\\nSubmission file saved to {output_submission_path} with {len(submission_df)} predictions.\")\n            print(\"Submission file head:\")\n            print(submission_df.head())\n        else:\n            print(\"\\nSkipping submission generation as test data or trained model is not available.\")\n            # Create an empty submission.csv if no predictions were generated or test data was missing\n            empty_submission_df = prepare_submission([])  # Pass empty list\n            output_submission_path = Path(\"./submission.csv\")\n            empty_submission_df.to_csv(output_submission_path, index=False)\n            print(f\"\\nCreated an empty submission file at {output_submission_path}.\")\n\n\n    except Exception as e:\n        print(f\"\\n--- Pipeline encountered a critical error: {e} ---\")\n        # import traceback\n        # traceback.print_exc() # For full stack trace in output\n\n    return final_model, training_history, validation_metrics\n\n\n# Run the pipeline\nif __name__ == '__main__':\n    # Adjust parameters as needed\n    HYPERPARAMS = {\n        'train_ratio': 0.9,  # Use more data for training if dataset is large\n        'batch_size': 32,  # Adjust based on GPU memory (16 or 32 often reasonable)\n        'num_epochs': 15,   # Adjust based on convergence observed in plots (start lower)\n        # 'num_epochs': 1,  # Keep small for testing\n        'learning_rate': 3e-4  # Common starting point for Adam\n    }\n\n    # Print system info\n    print(\"--- System Information ---\")\n    print(f\"PyTorch version: {torch.__version__}\")\n    # setup_device() call already prints GPU info if available\n    print(f\"Accelerator: {accelerator_type}, Device: {device}\")\n    print(f\"Hyperparameters: {HYPERPARAMS}\")\n\n    # Run pipeline\n    final_model, training_history, validation_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    )\n\n    if final_model:\n        print(\"\\nPipeline completed successfully.\")\n        if validation_metrics:\n            print(f\"Final Validation Metrics: {validation_metrics}\")\n    else:\n        print(\"\\nPipeline execution failed.\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-05-10T08:12:14.096482Z","iopub.execute_input":"2025-05-10T08:12:14.097250Z","execution_failed":"2025-05-10T12:45:28.879Z"}},"outputs":[],"execution_count":null}]}