{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":91249,"databundleVersionId":11294684,"sourceType":"competition"}],"dockerImageVersionId":30919,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Channel Attention Block Module for 2D Object Detection and Localization (CBAM2D)\n\n## Introduction\n\nThis report outlines a Python script implementing an enhanced object detection pipeline for locating bacterial flagellar motors in 2D tomographic slices. The code integrates a U-Net-inspired backbone with a CenterNet detection framework, augmented by a custom Convolutional Block Attention Module (CBAM) with self-attention. It processes a dataset of JPEG images and CSV labels, trains a model using a custom loss function, and generates predictions for a test set, producing a submission file and visualizations. Leveraging PyTorch, the script optimizes for GPU execution and includes data augmentation, evaluation metrics (F2 score), and 3D non-maximum suppression (NMS) for refined outputs.\n\n## Related Work\n\nThis network is inspired by Centernet (Duan 2019), with a Convolutional Block Attention Module (Woo 2018), and Self-Attention (Vaswani 2017). The backbone of the network is a U-net (Ronneberger 2015), which can be derived from a Convolutional Neural Network (LeCun 1989). \n\nCenterNet, introduced by Duan et al. in the 2019 paper \"CenterNet: Keypoint Triplets for Object Detection\" (Duan 2019), is an anchor-free object detection framework that represents objects as center points rather than bounding boxes, using keypoint estimation to predict object centers, sizes, and offsets. It simplifies detection by eliminating the need for anchor boxes, making it efficient and adaptable. The network in question draws inspiration from this approach, likely adopting CenterNet’s philosophy of focusing on keypoint-based representations or its heatmap-based prediction strategy for tasks like object detection or segmentation. \n\nTo enhance this, it incorporates the Convolutional Block Attention Module (CBAM) from Woo et al.’s 2018 paper \"CBAM: Convolutional Block Attention Module\" (Woo 2018). CBAM introduces a dual attention mechanism—channel attention (to emphasize important feature channels) and spatial attention (to focus on relevant spatial regions)—which refines feature maps by adaptively weighting them. By integrating CBAM, the network likely boosts its ability to prioritize critical features, improving detection or segmentation accuracy over a pure CenterNet-inspired design. \n\nSelf-attention is a mechanism in neural networks that allows a model to weigh the importance of different parts of an input sequence or feature map relative to each other, enabling it to capture long-range dependencies and contextual relationships without relying on sequential processing. Introduced prominently in the Transformer architecture (Vaswani 2017), it operates by computing three vectors—query (Q), key (K), and value (V)—from the input using learned linear transformations. The attention scores are derived by taking the dot product of the query and key vectors, scaled and normalized via a softmax function, which determines how much focus each element should receive. These scores are then used to weight the value vectors, producing an output that is a context-aware combination of the input features. In the provided code, the SelfAttention class adapts this concept to 2D convolutional feature maps, reshaping them into vectors and applying the mechanism spatially, enhancing the network’s ability to emphasize relevant regions (e.g., flagellar motor locations) while suppressing noise, thus improving feature refinement within the CBAM module.\n\nThe backbone of this network is U-Net, a fully convolutional architecture introduced by Ronneberger et al. in the 2015 paper \"U-Net: Convolutional Networks for Biomedical Image Segmentation\" (Ronneberger 2015), originally designed for precise pixel-wise segmentation in biomedical imaging. U-Net features a symmetric U-shaped structure with a contracting path (downsampling to capture context) and an expansive path (upsampling to recover spatial details), connected by skip connections that preserve fine-grained information. In this network, U-Net serves as the foundational feature extractor, likely processing input images to generate rich feature maps that the CenterNet-inspired head and CBAM modules can refine for specific tasks. U-Net itself is a derivative of the Convolutional Neural Network (CNN) concept pioneered by Yann LeCun et al. in the 1989 paper \"Backpropagation Applied to Handwritten Zip Code Recognition,\" which introduced the idea of using convolutional layers to extract hierarchical features from images. U-Net builds on this by adapting the CNN framework for segmentation, replacing fully connected layers with upsampling and skip connections, making it a specialized evolution of LeCun’s original CNN design.\n\nThis network appears to be a hybrid architecture blending multiple influential ideas. It uses U-Net (Ronneberger 2015) as its backbone, rooted in the CNN paradigm (LeCun 1989), to extract multi-scale features through its encoder-decoder structure. It takes inspiration from CenterNet (Duan 2019) for its task-specific head, possibly employing a keypoint-based or heatmap-based approach for detection or segmentation, diverging from traditional anchor-based methods. The addition of CBAM (Woo 2018) enhances this framework by introducing attention mechanisms that refine U-Net’s feature maps, allowing the network to focus on the most relevant channels and spatial regions. This combination suggests a design optimized for tasks requiring both precise localization (from U-Net), efficient object representation (from CenterNet), and feature enhancement (from CBAM), potentially targeting applications like biomedical segmentation, object detection, or keypoint estimation with improved performance over individual components.\n\n## Dataset\n\nThe dataset comprises tomographic images from the \"BYU Locating Bacterial Flagellar Motors 2025\" Kaggle competition, split into training and test sets. The training set includes JPEG images organized by tomogram ID (e.g., train/tomo_id/slice_XXXX.jpg) and a CSV file (train_labels.csv) with motor coordinates (Motor axis 0, Motor axis 1, Motor axis 2), array shapes, and voxel spacing. Each tomogram contains multiple 2D slices (e.g., 512x512 pixels), with motor annotations specifying 3D locations. The test set mirrors this structure but lacks labels, requiring prediction of motor positions. The custom FlagellarDataset class preprocesses images to 128x128 resolution, generates Gaussian heatmaps for training targets, and applies augmentations (flips, rotations) to enhance robustness.\n\n## Method\n\nThe method involves a modular pipeline implemented in PyTorch:\n\n1. Model Architecture (OptimizedCenterNet2D):\n  - Backbone: A U-Net-like encoder-decoder with four convolutional blocks (32, 64, 128, 256 channels) for downsampling, followed by upsampling layers with skip connections.\n  - Attention: An EnhancedCABM2D module at the bottleneck (256 channels), combining channel attention (global pooling, FC layers), spatial attention (dilated convolutions), and self-attention (query-key-value mechanism) to refine features.\n  - Heads: Three outputs—heatmap (sigmoid-activated), size (2D), and offset (2D)—predict motor centers, bounding box dimensions, and sub-pixel adjustments, respectively.\n\n2. Data Processing (FlagellarDataset):\n  - Loads and resizes images, normalizes them, and constructs training targets (heatmaps, sizes, offsets) from motor coordinates. Augmentation is applied to training data.\n\n3. Training (train_model):\n  - Uses a custom CenterNetLoss combining focal loss (for heatmap), Dice loss (for overlap), and MSE losses (for size/offset), optimized with AdamW and gradient accumulation. Mixed precision (AMP) enhances efficiency on GPU.\n\nEvaluates using F2 score with a 1000 Ångström threshold, employing early stopping based on validation loss.\n\n4. Inference (generate_submission):\n  - Processes test tomograms in batches, applies 3D NMS to filter detections, and outputs motor coordinates in a CSV file.\n\n5. Visualization (plot_slices, plot_test_predictions):\n  - Displays predicted vs. ground-truth heatmaps and centers for training examples, and plots 3D motor distributions for test predictions.\n\nThe script executes on GPU if available, with dynamic batch sizing and FP16 support for efficiency.\n\n## Results \n| Metric | Score |  \n|--------|-------|\n|Train FB-score| 0.4424 |\n|Validation FB-score| 0.3874 | \n|LB Score| 0.351 |\n\n## Conclusion\n\nThis script implements a robust solution for detecting flagellar motors, blending U-Net’s segmentation prowess, CenterNet’s keypoint detection, and an enhanced CBAM for feature refinement. It effectively handles the dataset’s 2D slice structure, producing precise 3D motor predictions via a heatmap-based approach. The use of custom loss functions, attention mechanisms, and post-processing (NMS) ensures high accuracy and interpretability, as evidenced by F2 scores and visualizations. The pipeline is optimized for Kaggle’s competition environment, delivering a submission-ready CSV and insightful plots in approximately 100 epochs or less, depending on early stopping.\n\n## References:\n1. Woo, S., Park, J., Lee, J. Y., & Kweon, I. S. (2018). Cbam: Convolutional block attention module. In Proceedings of the European conference on computer vision (ECCV) (pp. 3-19). https://arxiv.org/abs/1807.06521\n2. Vaswani, A., Shazeer, N., Parmar, N., Uszkoreit, J., Jones, L., Gomez, A. N., ... & Polosukhin, I. (2017). Attention is all you need. Advances in neural information processing systems, 30. https://proceedings.neurips.cc/paper/2017/hash/3f5ee243547dee91fbd053c1c4a845aa-Abstract.html\n3. Duan, K., Bai, S., Xie, L., Qi, H., Huang, Q., & Tian, Q. (2019). Centernet: Keypoint triplets for object detection. In Proceedings of the IEEE/CVF international conference on computer vision (pp. 6569-6578). https://openaccess.thecvf.com/content_ICCV_2019/html/Duan_CenterNet_Keypoint_Triplets_for_Object_Detection_ICCV_2019_paper.html\n4. Ronneberger, O., Fischer, P., & Brox, T. (2015). U-net: Convolutional networks for biomedical image segmentation. In Medical image computing and computer-assisted intervention–MICCAI 2015: 18th international conference, Munich, Germany, October 5-9, 2015, proceedings, part III 18 (pp. 234-241). Springer international publishing. https://link.springer.com/chapter/10.1007/978-3-319-24574-4_28\n5. LeCun, Y., Boser, B., Denker, J., Henderson, D., Howard, R., Hubbard, W., & Jackel, L. (1989). Handwritten digit recognition with a back-propagation network. Advances in neural information processing systems, 2. https://proceedings.neurips.cc/paper/1989/hash/53c3bce66e43be4f209556518c2fcb54-Abstract.html","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import DataLoader, Dataset, random_split\nfrom pathlib import Path\nfrom tqdm import tqdm\nimport os\nfrom PIL import Image\nimport torchvision.transforms as T\nimport matplotlib.pyplot as plt\nimport matplotlib.patches as patches\nfrom concurrent.futures import ThreadPoolExecutor\nimport time\nfrom sklearn.model_selection import train_test_split\nimport multiprocessing as mp\nfrom mpl_toolkits.mplot3d import Axes3D\n\ntry:\n    mp.set_start_method('spawn', force=True)\nexcept RuntimeError:\n    pass\n\nclass MultiHeadSelfAttention(nn.Module):\n    def __init__(self, in_channels, num_heads=8):\n        super(MultiHeadSelfAttention, self).__init__()\n        assert in_channels % num_heads == 0, \"in_channels must be divisible by num_heads\"\n        \n        self.in_channels = in_channels\n        self.num_heads = num_heads\n        self.head_dim = in_channels // num_heads\n        \n        self.query = nn.Conv2d(in_channels, in_channels, kernel_size=1)\n        self.key = nn.Conv2d(in_channels, in_channels, kernel_size=1)\n        self.value = nn.Conv2d(in_channels, in_channels, kernel_size=1)\n        self.out_proj = nn.Conv2d(in_channels, in_channels, kernel_size=1)\n        self.gamma = nn.Parameter(torch.zeros(1))\n        self.softmax = nn.Softmax(dim=-1)\n\n    def forward(self, x):\n        batch, channels, height, width = x.size()\n        n_pixels = height * width\n        \n        proj_query = self.query(x).view(batch, self.num_heads, self.head_dim, n_pixels).permute(0, 1, 3, 2)\n        proj_key = self.key(x).view(batch, self.num_heads, self.head_dim, n_pixels)\n        proj_value = self.value(x).view(batch, self.num_heads, self.head_dim, n_pixels).permute(0, 1, 3, 2)\n        \n        # Compute attention scores with clipping\n        energy = torch.matmul(proj_query, proj_key) / (self.head_dim ** 0.5)\n        energy = torch.clamp(energy, min=-10, max=10)  # Prevent extreme values\n        attention = self.softmax(energy)\n        \n        out = torch.matmul(attention, proj_value)        \n        out = out.permute(0, 1, 3, 2).contiguous().view(batch, self.in_channels, height, width)\n        out = self.out_proj(out)\n        \n        return self.gamma * out + x\n\n# Replace the SelfAttention in EnhancedCABM2D with MultiHeadSelfAttention\nclass EnhancedCABM2D(nn.Module):\n    def __init__(self, in_channels, reduction=16):\n        super(EnhancedCABM2D, self).__init__()\n        self.global_pool = nn.AdaptiveAvgPool2d(1)\n        self.fc1 = nn.Conv2d(in_channels, in_channels // reduction, kernel_size=1)\n        self.fc2 = nn.Conv2d(in_channels // reduction, in_channels, kernel_size=1)\n        self.sigmoid = nn.Sigmoid()\n        self.conv_spatial1 = nn.Conv2d(in_channels, 1, kernel_size=7, padding=3, dilation=1, bias=False)\n        self.conv_spatial2 = nn.Conv2d(in_channels, 1, kernel_size=7, padding=9, dilation=3, bias=False)\n        self.conv_refine = nn.Conv2d(in_channels, in_channels, kernel_size=3, padding=1)\n        self.bn = nn.BatchNorm2d(in_channels)\n        self.relu = nn.ReLU(inplace=True)\n        self.self_attention = MultiHeadSelfAttention(in_channels, num_heads=8)  # Updated to MultiHeadSelfAttention\n\n    def forward(self, x):\n        channel_avg = self.global_pool(x)\n        channel_att = self.fc1(channel_avg)\n        channel_att = self.relu(channel_att)\n        channel_att = self.fc2(channel_att)\n        channel_att = self.sigmoid(channel_att)\n        x_channel = x * channel_att\n        spatial_att1 = self.conv_spatial1(x_channel)\n        spatial_att2 = self.conv_spatial2(x_channel)\n        spatial_att = self.sigmoid(spatial_att1 + spatial_att2)\n        x_spatial = x_channel * spatial_att\n        x_self_att = self.self_attention(x_spatial)\n        x_refined = self.conv_refine(x_self_att)\n        x_refined = self.bn(x_refined)\n        x_refined = self.relu(x_refined)\n        return x + x_refined\n\nclass OptimizedCenterNet2D(nn.Module):\n    def __init__(self, in_channels=1, out_channels=1):\n        super(OptimizedCenterNet2D, self).__init__()\n        self.enc1 = nn.Sequential(\n            nn.Conv2d(in_channels, 32, kernel_size=3, padding=1),\n            nn.BatchNorm2d(32),\n            nn.ReLU(),\n            nn.Dropout(0.3)\n        )\n        self.enc2 = nn.Sequential(\n            nn.Conv2d(32, 64, kernel_size=3, padding=1, stride=2),\n            nn.BatchNorm2d(64),\n            nn.ReLU(),\n            nn.Dropout(0.3)\n        )\n        self.enc3 = nn.Sequential(\n            nn.Conv2d(64, 128, kernel_size=3, padding=1, stride=2),\n            nn.BatchNorm2d(128),\n            nn.ReLU(),\n            nn.Dropout(0.3)\n        )\n        self.enc4 = nn.Sequential(\n            nn.Conv2d(128, 256, kernel_size=3, padding=1, stride=2),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Dropout(0.3)\n        )\n        self.attention = EnhancedCABM2D(in_channels=256)\n        self.dec4 = nn.Sequential(\n            nn.ConvTranspose2d(256, 128, kernel_size=3, stride=2, padding=1, output_padding=1),\n            nn.BatchNorm2d(128),\n            nn.ReLU()\n        )\n        self.dec3 = nn.Sequential(\n            nn.ConvTranspose2d(128, 64, kernel_size=3, stride=2, padding=1, output_padding=1),\n            nn.BatchNorm2d(64),\n            nn.ReLU()\n        )\n        self.dec2 = nn.Sequential(\n            nn.ConvTranspose2d(64, 32, kernel_size=3, stride=2, padding=1, output_padding=1),\n            nn.BatchNorm2d(32),\n            nn.ReLU()\n        )\n        self.heatmap_head = nn.Sequential(\n            nn.Conv2d(32, out_channels, kernel_size=1),\n            nn.Sigmoid()\n        )\n        self.size_head = nn.Conv2d(32, 2, kernel_size=1)\n        self.offset_head = nn.Conv2d(32, 2, kernel_size=1)\n\n    def forward(self, x):\n        e1 = self.enc1(x)\n        e2 = self.enc2(e1)\n        e3 = self.enc3(e2)\n        e4 = self.enc4(e3)\n        e4_att = self.attention(e4)\n        d4 = self.dec4(e4_att) + e3\n        d3 = self.dec3(d4) + e2\n        d2 = self.dec2(d3) + e1\n        heatmap = self.heatmap_head(d2)\n        size = self.size_head(d2)\n        offset = self.offset_head(d2)\n        return heatmap, size, offset\n\nclass FlagellarDataset(Dataset):\n    def __init__(self, csv_file=None, root_dir=None, new_size=(256, 256), trust_region=4, is_test=False):\n        self.root_dir = Path(root_dir)\n        self.new_size = new_size\n        self.trust_region = trust_region\n        self.is_test = is_test\n        self.data = []\n        self.spatial_augment = T.Compose([\n            T.RandomHorizontalFlip(p=0.5),\n            T.RandomRotation(degrees=10),  # Reduced from 30\n            T.RandomAffine(degrees=0, scale=(0.95, 1.05))  # Narrowed from (0.8, 1.2)\n        ])\n        self.image_augment = T.Compose([\n            T.RandomApply([T.ColorJitter(brightness=0.1, contrast=0.1)], p=0.2)  # Reduced intensity, dropped blur\n        ])\n\n        if csv_file and os.path.exists(csv_file) and not is_test:\n            labels = pd.read_csv(csv_file)\n            self.tomo_ids = labels['tomo_id'].unique().tolist()\n            self.motors_map = labels.groupby('tomo_id').apply(\n                lambda g: np.array(g[['Motor axis 0', 'Motor axis 1', 'Motor axis 2']].values[0], dtype=np.float32),\n                include_groups=False\n            ).to_dict()\n            self.size_map = labels.groupby('tomo_id').agg({\n                'Array shape (axis 1)': 'first',\n                'Array shape (axis 2)': 'first',\n                'Array shape (axis 0)': 'first',\n                'Voxel spacing': 'first'\n            }).to_dict('index')\n\n            for tomo_id in tqdm(self.tomo_ids, desc=\"Loading data\"):\n                motor = self.motors_map[tomo_id]\n                files = sorted(self.root_dir.joinpath(tomo_id).glob('*.jpg'))\n                num_slices = len(files)\n\n                if np.all(motor == -1):\n                    center_z = num_slices // 2\n                else:\n                    center_z = int(motor[0])\n\n                z_start = max(0, center_z - self.trust_region)\n                z_end = min(num_slices, center_z + self.trust_region + 1)\n\n                for z in range(z_start, z_end):\n                    img = Image.open(files[z]).convert('L')\n                    img = T.functional.resize(img, new_size)\n                    img = T.functional.to_tensor(img)\n                    img = (img - img.mean()) / (img.std() + 1e-8)\n\n                    motor_norm = torch.tensor([\n                        motor[1] / self.size_map[tomo_id]['Array shape (axis 1)'],\n                        motor[2] / self.size_map[tomo_id]['Array shape (axis 2)']\n                    ], dtype=torch.float32) if np.all(motor != -1) else torch.tensor([-1, -1], dtype=torch.float32)\n\n                    heatmap = self.generate_gaussian_heatmap(motor_norm, self.new_size) if np.all(motor != -1) else torch.zeros(self.new_size)\n                    if np.all(motor != -1):\n                        voxel_spacing = self.size_map[tomo_id]['Voxel spacing']\n                        orig_height = self.size_map[tomo_id]['Array shape (axis 1)']\n                        orig_width = self.size_map[tomo_id]['Array shape (axis 2)']\n                        target_size_pixels_y = 1000.0 / voxel_spacing\n                        target_size_pixels_x = 1000.0 / voxel_spacing\n                        size_target = torch.tensor([\n                            target_size_pixels_y / orig_height,\n                            target_size_pixels_x / orig_width\n                        ], dtype=torch.float32)\n                    else:\n                        size_target = torch.zeros(2, dtype=torch.float32)\n                    offset_target = torch.zeros(2, dtype=torch.float32)\n\n                    if not self.is_test and np.all(motor != -1):\n                        img_aug = self.image_augment(img)\n                        stacked = torch.stack([img_aug[0], heatmap], dim=0)\n                        stacked_aug = self.spatial_augment(stacked)\n                        img = stacked_aug[0].unsqueeze(0)\n                        heatmap = stacked_aug[1]\n                        size_map = torch.full(self.new_size, size_target[0], dtype=torch.float32)\n                        size_map2 = torch.full(self.new_size, size_target[1], dtype=torch.float32)\n                        offset_map = torch.full(self.new_size, offset_target[0], dtype=torch.float32)\n                        offset_map2 = torch.full(self.new_size, offset_target[1], dtype=torch.float32)\n\n                        peak_idx = heatmap.view(-1).argmax()\n                        y_new = peak_idx // self.new_size[1]\n                        x_new = peak_idx % self.new_size[1]\n                        motor_norm = torch.tensor([y_new / self.new_size[0], x_new / self.new_size[1]], dtype=torch.float32)\n\n                    self.data.append({\n                        'tomo_id': tomo_id,\n                        'slice': img,\n                        'heatmap': heatmap,\n                        'size': size_target,\n                        'offset': offset_target,\n                        'center': motor_norm,\n                        'z': z,\n                        'orig_shape': torch.tensor([self.size_map[tomo_id]['Array shape (axis 1)'],\n                                                   self.size_map[tomo_id]['Array shape (axis 2)'],\n                                                   num_slices], dtype=torch.float32),\n                        'voxel_spacing': self.size_map[tomo_id]['Voxel spacing'],\n                        'motor': torch.tensor(motor, dtype=torch.float32)\n                    })\n        else:\n            self.tomo_ids = [d.name for d in self.root_dir.iterdir() if d.is_dir()]\n            for tomo_id in tqdm(self.tomo_ids, desc=\"Loading test data\"):\n                files = sorted(self.root_dir.joinpath(tomo_id).glob('*.jpg'))\n                num_slices = len(files)\n                orig_shape = torch.tensor([Image.open(files[0]).size[1], Image.open(files[0]).size[0], num_slices], dtype=torch.float32)\n                for z, file in enumerate(files):\n                    self.data.append({\n                        'tomo_id': tomo_id,\n                        'slice_path': str(file),\n                        'z': z,\n                        'orig_shape': orig_shape\n                    })\n\n    def generate_gaussian_heatmap(self, center, xy_size, sigma=4.0):\n        heatmap = torch.zeros(xy_size)\n        yc, xc = center\n        yc = yc * xy_size[0]\n        xc = xc * xy_size[1]\n        y_coords, x_coords = torch.meshgrid(\n            torch.arange(xy_size[0], dtype=torch.float32),\n            torch.arange(xy_size[1], dtype=torch.float32),\n            indexing='ij'\n        )\n        dist = ((y_coords - yc) ** 2 + (x_coords - xc) ** 2) / (2.0 * sigma ** 2)\n        heatmap = torch.exp(-dist)\n        return heatmap\n\n    def __len__(self):\n        return len(self.data)\n\n    def __getitem__(self, idx):\n        item = self.data[idx]\n        if self.is_test:\n            return {\n                'tomo_id': item['tomo_id'],\n                'slice_path': item['slice_path'],\n                'z': item['z'],\n                'orig_shape': item['orig_shape']\n            }\n        return {\n            'tomo_id': item['tomo_id'],\n            'slice': item['slice'],\n            'heatmap': item['heatmap'].unsqueeze(0),\n            'size': item['size'],\n            'offset': item['offset'],\n            'center': item['center'],\n            'z': item['z'],\n            'orig_shape': item['orig_shape'],\n            'voxel_spacing': item['voxel_spacing'],\n            'motor': item['motor']\n        }\n\ndef custom_collate_fn(batch):\n    if 'slice' in batch[0]:\n        return {\n            'tomo_id': [item['tomo_id'] for item in batch],\n            'slice': torch.stack([item['slice'] for item in batch]),\n            'heatmap': torch.stack([item['heatmap'] for item in batch]),\n            'size': torch.stack([item['size'] for item in batch]),\n            'offset': torch.stack([item['offset'] for item in batch]),\n            'center': torch.stack([item['center'] for item in batch]),\n            'z': torch.tensor([item['z'] for item in batch], dtype=torch.long),\n            'orig_shape': torch.stack([item['orig_shape'] for item in batch]),\n            'voxel_spacing': torch.tensor([item['voxel_spacing'] for item in batch], dtype=torch.float32),\n            'motor': torch.stack([item['motor'] for item in batch])\n        }\n    else:\n        return {\n            'tomo_id': [item['tomo_id'] for item in batch],\n            'slice_path': [item['slice_path'] for item in batch],\n            'z': torch.tensor([item['z'] for item in batch], dtype=torch.long),\n            'orig_shape': torch.stack([item['orig_shape'] for item in batch])\n        }\n\nclass CenterNetLoss(nn.Module):\n    def __init__(self, gamma=2.0, size_weight=0.5, offset_weight=0.1):\n        super(CenterNetLoss, self).__init__()\n        self.gamma = gamma  # Focusing parameter\n        self.size_weight = size_weight\n        self.offset_weight = offset_weight\n\n    def gaussian_focal_loss(self, pred_heatmap, target_heatmap):\n        # pred_heatmap: Predicted probabilities (after sigmoid)\n        # target_heatmap: Ground truth Gaussian heatmap (0 to 1)\n        \n        # Positive and negative terms weighted by target heatmap values\n        pos_loss = -target_heatmap * (1 - pred_heatmap) ** self.gamma * torch.log(pred_heatmap + 1e-6)\n        neg_loss = -(1 - target_heatmap) * pred_heatmap ** self.gamma * torch.log(1 - pred_heatmap + 1e-6)\n        \n        # Sum over all pixels and normalize by batch size\n        loss = (pos_loss + neg_loss).sum() / pred_heatmap.size(0)\n        return loss\n\n    def forward(self, pred_heatmap, pred_size, pred_offset, target_heatmap, target_size, target_offset):\n        # Ensure pred_heatmap is in probability space (since heatmap_head has Sigmoid)\n        gfl = self.gaussian_focal_loss(pred_heatmap, target_heatmap)\n\n        # Size and offset losses (unchanged)\n        batch_size = pred_heatmap.size(0)\n        pred_size_at_centers = torch.zeros(batch_size, 2, device=pred_heatmap.device)\n        pred_offset_at_centers = torch.zeros(batch_size, 2, device=pred_heatmap.device)\n        for i in range(batch_size):\n            heatmap = pred_heatmap[i].squeeze()\n            peak_idx = heatmap.view(-1).argmax()\n            y = peak_idx // 256\n            x = peak_idx % 256\n            pred_size_at_centers[i] = pred_size[i, :, y, x]\n            pred_offset_at_centers[i] = pred_offset[i, :, y, x]\n\n        size_loss = F.mse_loss(pred_size_at_centers, target_size, reduction='mean') * self.size_weight\n        offset_loss = F.mse_loss(pred_offset_at_centers, target_offset, reduction='mean') * self.offset_weight\n\n        return gfl + size_loss + offset_loss\n\ndef extract_centroid(heatmap, size, offset, xy_size=(256, 256), threshold=0.1):  # Updated for 256x256\n    heatmap = heatmap.squeeze()\n    if heatmap.max() < threshold:\n        return torch.tensor([-1, -1], dtype=torch.float32, device=heatmap.device), torch.zeros(2, device=heatmap.device), torch.zeros(2, device=heatmap.device)\n    \n    peak_value, peak_idx = heatmap.view(-1).topk(1)\n    if peak_value < threshold:\n        return torch.tensor([-1, -1], dtype=torch.float32, device=heatmap.device), torch.zeros(2, device=heatmap.device), torch.zeros(2, device=heatmap.device)\n    \n    y = peak_idx // xy_size[1]\n    x = peak_idx % xy_size[1]\n    \n    y_norm = torch.clamp(y.float() / xy_size[0], 0, 1)\n    x_norm = torch.clamp(x.float() / xy_size[1], 0, 1)\n    \n    center = torch.tensor([y_norm, x_norm], dtype=torch.float32, device=heatmap.device)\n    pred_size = size[:, y, x].squeeze()\n    pred_offset = offset[:, y, x].squeeze()\n    center = center + pred_offset\n    return center, pred_size, pred_offset\n\ndef denormalize_predictions(pred_center, pred_size, z, orig_shape):\n    pred_center_denorm = torch.zeros(3, dtype=torch.float32, device=pred_center.device)\n    pred_center_denorm[0] = z\n    pred_center_denorm[1] = pred_center[0] * orig_shape[0]\n    pred_center_denorm[2] = pred_center[1] * orig_shape[1]\n    pred_size_denorm = pred_size * torch.tensor([orig_shape[0], orig_shape[1]], dtype=torch.float32, device=pred_size.device)\n    return pred_center_denorm, pred_size_denorm\n\ndef calculate_fbeta_score(pred_centers, true_centers, voxel_spacings, threshold_angstroms=1000, beta=2.0):\n    TP, TN, FP, FN = 0, 0, 0, 0\n    for pred_center, true_center, voxel_spacing in zip(pred_centers, true_centers, voxel_spacings):\n        voxel_spacing = voxel_spacing.item()\n        pred_array = pred_center.detach().cpu().numpy()\n        true_array = true_center.detach().cpu().numpy()\n\n        if np.all(true_array == -1):\n            if np.all(pred_array == -1):\n                TN += 1\n            else:\n                FP += 1\n            continue\n        if np.all(pred_array == -1):\n            FN += 1\n            continue\n\n        distance = np.linalg.norm((true_array - pred_array) * voxel_spacing)\n        if distance <= threshold_angstroms:\n            TP += 1\n        else:\n            FN += 1\n\n    if TP + FP + FN == 0:\n        fbeta = 0.0\n    else:\n        beta2 = beta ** 2\n        fbeta = (1 + beta2) * TP / ((1 + beta2) * TP + beta2 * FN + FP)\n    return fbeta, TP, TN, FP, FN\n\ndef plot_slices(model, example, device):\n    model.eval()\n    with torch.no_grad():\n        slice_data = example['slice'].unsqueeze(0).to(device)\n        heatmap_true = example['heatmap'].unsqueeze(0)\n        size_true = example['size']\n        offset_true = example['offset']\n        center = example['center']\n        orig_shape = example['orig_shape'].to(device)\n        tomo_id = example['tomo_id']\n        z = example['z']\n\n        heatmap_pred, size_pred, offset_pred = model(slice_data)\n        heatmap_pred, size_pred, offset_pred = heatmap_pred[0], size_pred[0], offset_pred[0]\n        heatmap_true = heatmap_true[0]\n\n        pred_center_norm, pred_size_norm, pred_offset_norm = extract_centroid(heatmap_pred, size_pred, offset_pred, threshold=0.3)  # Higher threshold\n        pred_center, pred_size = denormalize_predictions(pred_center_norm, pred_size_norm, z, orig_shape[:2])\n        center_denorm, _ = denormalize_predictions(center, size_true, z, orig_shape[:2]) if torch.all(center >= 0) else (torch.tensor([-1, -1, -1], device=device), torch.zeros(2, device=device))\n\n        slice_data = slice_data[0, 0].cpu().numpy()\n        heatmap_pred_slice = heatmap_pred[0].cpu().numpy()\n        heatmap_true_slice = heatmap_true[0].cpu().numpy()\n\n        distance = np.linalg.norm(center_denorm.cpu().numpy() - pred_center.cpu().numpy()) if torch.all(center >= 0) else float('inf')\n        distance_text = f\"Distance: {distance:.2f}\" if distance != float('inf') else \"No GT Motor\"\n\n        fig, axes = plt.subplots(1, 3, figsize=(15, 5))\n\n        ax = axes[0]\n        ax.imshow(slice_data, cmap='gray')\n        ax.set_title(f\"Tomo: {tomo_id}\\nZ: {z}\\n{distance_text}\")\n        ax.axis('off')\n\n        scale_y = orig_shape[0].item() / 256  # Updated for 256x256\n        scale_x = orig_shape[1].item() / 256\n\n        if torch.all(center >= 0):\n            true_y_display = center_denorm[1].item() / scale_y\n            true_x_display = center_denorm[2].item() / scale_x\n            ax.plot(true_x_display, true_y_display, 'go', label='GT Center')\n            rect_gt = patches.Rectangle(\n                (max(0, true_x_display - size_true[1].item() * scale_x / 2), max(0, true_y_display - size_true[0].item() * scale_y / 2)),\n                size_true[1].item() * scale_x, size_true[0].item() * scale_y, linewidth=2, edgecolor='g', facecolor='none', label='GT Box'\n            )\n            ax.add_patch(rect_gt)\n\n        if not torch.all(pred_center == -1):\n            pred_y_display = pred_center[1].item() / scale_y\n            pred_x_display = pred_center[2].item() / scale_x\n            ax.plot(pred_x_display, pred_y_display, 'ro', label='Pred Center')\n            rect_pred = patches.Rectangle(\n                (max(0, pred_x_display - pred_size[1].item() / 2), max(0, pred_y_display - pred_size[0].item() / 2)),\n                pred_size[1].item(), pred_size[0].item(), linewidth=2, edgecolor='r', facecolor='none', label='Pred Box'\n            )\n            ax.add_patch(rect_pred)\n\n        ax.legend()\n\n        axes[1].imshow(heatmap_pred_slice, cmap='hot', vmin=0, vmax=1)\n        axes[1].set_title(\"Predicted Heatmap\")\n        axes[1].axis('off')\n\n        axes[2].imshow(heatmap_true_slice, cmap='hot', vmin=0, vmax=1)\n        axes[2].set_title(\"Ground Truth Heatmap\")\n        axes[2].axis('off')\n\n        plt.tight_layout()\n        plt.show()\n\ndef train_model(model, train_loader, val_loader, num_epochs=150, device='cuda', accum_steps=2, max_grad_norm=1.0):\n    criterion = CenterNetLoss()\n    optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=1e-4)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=num_epochs, eta_min=1e-6)\n    model.to(device)\n    scaler = torch.cuda.amp.GradScaler()\n\n    best_val_fbeta = 0.0\n    patience = 50\n    epochs_no_improve = 0\n    best_model_path = 'best_model.pth'\n    example_to_plot = None\n\n    for epoch in range(num_epochs):\n        model.train()\n        train_loss = 0.0\n        train_preds, train_trues, train_voxel_spacings = [], [], []\n\n        optimizer.zero_grad()\n        for i, batch in enumerate(tqdm(train_loader, desc=f\"Epoch {epoch+1}/{num_epochs} - Train\")):\n            slices = batch['slice'].to(device)\n            heatmaps = batch['heatmap'].to(device)\n            sizes = batch['size'].to(device)\n            offsets = batch['offset'].to(device)\n            centers = batch['center'].to(device)\n            orig_shapes = batch['orig_shape'].to(device)\n            voxel_spacings = batch['voxel_spacing'].to(device)\n            motors = batch['motor'].to(device)\n            zs = batch['z'].to(device)\n\n            with torch.cuda.amp.autocast():\n                pred_heatmaps, pred_sizes, pred_offsets = model(slices)\n                loss = criterion(pred_heatmaps, pred_sizes, pred_offsets, heatmaps, sizes, offsets)\n                loss = loss / accum_steps\n            scaler.scale(loss).backward()\n            train_loss += loss.item() * slices.size(0) * accum_steps\n\n            if (i + 1) % accum_steps == 0:\n                # Apply gradient clipping before unscaling and stepping\n                scaler.unscale_(optimizer)\n                torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=max_grad_norm)\n                scaler.step(optimizer)\n                scaler.update()\n                optimizer.zero_grad()\n\n            pred_centers_norm, pred_sizes_norm, _ = zip(*[extract_centroid(pred_heatmaps[i], pred_sizes[i], pred_offsets[i], threshold=0.3) \n                                                          for i in range(slices.size(0))])\n            pred_centers, _ = zip(*[denormalize_predictions(pred_centers_norm[i], pred_sizes_norm[i], zs[i], orig_shapes[i, :2])\n                                    for i in range(slices.size(0))])\n            train_preds.extend(pred_centers)\n            train_trues.extend(motors)\n            train_voxel_spacings.extend(voxel_spacings)\n\n            if example_to_plot is None:\n                example_to_plot = {\n                    'slice': slices[0].cpu(),\n                    'heatmap': heatmaps[0].cpu(),\n                    'size': sizes[0].cpu(),\n                    'offset': offsets[0].cpu(),\n                    'center': centers[0].cpu(),\n                    'orig_shape': orig_shapes[0].cpu(),\n                    'tomo_id': batch['tomo_id'][0],\n                    'z': zs[0].cpu()\n                }\n\n        if (i + 1) % accum_steps != 0:\n            scaler.unscale_(optimizer)\n            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=max_grad_norm)\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad()\n\n        model.eval()\n        val_loss = 0.0\n        val_preds, val_trues, val_voxel_spacings = [], [], []\n        with torch.no_grad():\n            for batch in tqdm(val_loader, desc=f\"Epoch {epoch+1}/{num_epochs} - Val\"):\n                slices = batch['slice'].to(device)\n                heatmaps = batch['heatmap'].to(device)\n                sizes = batch['size'].to(device)\n                offsets = batch['offset'].to(device)\n                centers = batch['center'].to(device)\n                orig_shapes = batch['orig_shape'].to(device)\n                voxel_spacings = batch['voxel_spacing'].to(device)\n                motors = batch['motor'].to(device)\n                zs = batch['z'].to(device)\n\n                with torch.cuda.amp.autocast():\n                    pred_heatmaps, pred_sizes, pred_offsets = model(slices)\n                    loss = criterion(pred_heatmaps, pred_sizes, pred_offsets, heatmaps, sizes, offsets)\n                val_loss += loss.item() * slices.size(0)\n\n                pred_centers_norm, pred_sizes_norm, _ = zip(*[extract_centroid(pred_heatmaps[i], pred_sizes[i], pred_offsets[i], threshold=0.3) \n                                                              for i in range(slices.size(0))])\n                pred_centers, _ = zip(*[denormalize_predictions(pred_centers_norm[i], pred_sizes_norm[i], zs[i], orig_shapes[i, :2])\n                                        for i in range(slices.size(0))])\n                val_preds.extend(pred_centers)\n                val_trues.extend(motors)\n                val_voxel_spacings.extend(voxel_spacings)\n\n        avg_train_loss = train_loss / len(train_loader.dataset)\n        avg_val_loss = val_loss / len(val_loader.dataset)\n        train_fbeta, train_TP, train_TN, train_FP, train_FN = calculate_fbeta_score(train_preds, train_trues, train_voxel_spacings)\n        val_fbeta, val_TP, val_TN, val_FP, val_FN = calculate_fbeta_score(val_preds, val_trues, val_voxel_spacings)\n        scheduler.step()\n\n        print(f\"Epoch {epoch+1}/{num_epochs}\")\n        print(f\"Train Loss: {avg_train_loss:.6f}, Train F2: {train_fbeta:.4f}, TP: {train_TP}, TN: {train_TN}, FP: {train_FP}, FN: {train_FN}\")\n        print(f\"Val Loss: {avg_val_loss:.6f}, Val F2: {val_fbeta:.4f}, TP: {val_TP}, TN: {val_TN}, FP: {val_FP}, FN: {val_FN}\")\n\n        if example_to_plot is not None:\n            plot_slices(model, example_to_plot, device)\n            example_to_plot = None\n\n        if val_fbeta > best_val_fbeta:\n            best_val_fbeta = val_fbeta\n            epochs_no_improve = 0\n            torch.save(model.state_dict(), best_model_path)\n            print(f\"Saved best model with Val F2: {best_val_fbeta:.4f}\")\n        else:\n            epochs_no_improve += 1\n            print(f\"No improvement in Val F2. Epochs without improvement: {epochs_no_improve}/{patience}\")\n            if epochs_no_improve >= patience:\n                print(f\"Early stopping at epoch {epoch+1}\")\n                model.load_state_dict(torch.load(best_model_path))\n                print(f\"Loaded best model with Val F2: {best_val_fbeta:.4f}\")\n                break\n\n    if os.path.exists(best_model_path):\n        model.load_state_dict(torch.load(best_model_path))\n        print(f\"Training completed. Best Val F2: {best_val_fbeta:.4f}\")\n\n    return best_val_fbeta\n\ndef preprocess_batch(slice_paths, new_size=(256, 256)):  # Updated for 256x256\n    images = []\n    for path in slice_paths:\n        img = Image.open(path).convert('L')\n        img = T.functional.resize(img, new_size)\n        img = T.functional.to_tensor(img)\n        img = (img - img.mean()) / (img.std() + 1e-8)\n        images.append(img)\n    return torch.stack(images)\n\ndef process_tomogram(tomo_id, model, test_dataset, device, index=0, total=1, confidence_threshold=0.3, nms_threshold=0.1):  # Tighter thresholds\n    print(f\"Processing tomogram {tomo_id} ({index}/{total})\")\n    tomo_dir = os.path.join('/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/test', tomo_id)\n    slice_files = sorted([os.path.join(tomo_dir, f) for f in os.listdir(tomo_dir) if f.endswith('.jpg')])\n    num_slices = len(slice_files)\n\n    tomo_data = next(item for item in test_dataset.data if item['tomo_id'] == tomo_id)\n    orig_shape = tomo_data['orig_shape']\n\n    all_detections = []\n    batch_size = 4 if device.startswith('cuda') else os.cpu_count() * 2\n\n    if device.startswith('cuda'):\n        gpu_mem = torch.cuda.get_device_properties(0).total_memory / 1e9\n        free_mem = gpu_mem - torch.cuda.memory_allocated(0) / 1e9\n        batch_size = max(4, min(16, int(free_mem / 2)))\n\n    model.eval()\n    with torch.no_grad():\n        for batch_start in range(0, num_slices, batch_size):\n            batch_end = min(batch_start + batch_size, num_slices)\n            batch_paths = slice_files[batch_start:batch_end]\n            batch_slices = preprocess_batch(batch_paths).to(device)\n\n            with torch.cuda.amp.autocast():\n                pred_heatmaps, pred_sizes, pred_offsets = model(batch_slices)\n\n            for i, (heatmap, size, offset) in enumerate(zip(pred_heatmaps, pred_sizes, pred_offsets)):\n                confidence = heatmap.max().item()\n                if confidence >= confidence_threshold:\n                    pred_center_norm, pred_size_norm, _ = extract_centroid(heatmap, size, offset, threshold=confidence_threshold)\n                    z = batch_start + i\n                    pred_center, pred_size = denormalize_predictions(pred_center_norm, pred_size_norm, z, orig_shape[:2])\n                    all_detections.append({\n                        'z': z,\n                        'y': pred_center[1].item(),\n                        'x': pred_center[2].item(),\n                        'confidence': confidence,\n                        'width': pred_size[1].item(),\n                        'height': pred_size[0].item()\n                    })\n\n    final_detections = perform_3d_nms(all_detections, nms_threshold)\n\n    if not final_detections:\n        return {'tomo_id': tomo_id, 'Motor axis 0': -1, 'Motor axis 1': -1, 'Motor axis 2': -1}\n\n    final_detections.sort(key=lambda x: x['confidence'], reverse=True)\n    best_detection = final_detections[0]\n\n    return {\n        'tomo_id': tomo_id,\n        'Motor axis 0': round(best_detection['z']),\n        'Motor axis 1': round(best_detection['y']),\n        'Motor axis 2': round(best_detection['x'])\n    }\n\ndef perform_3d_nms(detections, distance_threshold=0.1):  # Tighter NMS\n    if not detections:\n        return []\n\n    detections = sorted(detections, key=lambda x: x['confidence'], reverse=True)\n    final_detections = []\n\n    def distance_3d(d1, d2):\n        return np.sqrt((d1['z'] - d2['z'])**2 + (d1['y'] - d2['y'])**2 + (d1['x'] - d2['x'])**2)\n\n    trust_region = 4\n    threshold = trust_region * distance_threshold\n\n    while detections:\n        best_detection = detections.pop(0)\n        final_detections.append(best_detection)\n        detections = [d for d in detections if distance_3d(d, best_detection) > threshold]\n\n    return final_detections\n\ndef process_tomogram_wrapper(args):\n    tomo_id, model, test_dataset, device, index, total = args\n    return process_tomogram(tomo_id, model, test_dataset, device, index, total)\n\ndef generate_submission(test_dataset, model, device):\n    total_tomos = len(test_dataset.tomo_ids)\n    model.to(device)\n    if device.startswith('cuda'):\n        try:\n            model.half()\n            print(\"Using FP16 for inference\")\n        except:\n            print(\"FP16 not supported\")\n\n    results = []\n    with ThreadPoolExecutor(max_workers=1) as executor:\n        args_list = [(tomo_id, model, test_dataset, device, i + 1, total_tomos) \n                     for i, tomo_id in enumerate(test_dataset.tomo_ids)]\n        results = list(executor.map(process_tomogram_wrapper, args_list))\n\n    submission_df = pd.DataFrame(results, columns=['tomo_id', 'Motor axis 0', 'Motor axis 1', 'Motor axis 2'])\n    submission_df.to_csv('/kaggle/working/submission.csv', index=False)\n    motors_found = sum(1 for r in results if r['Motor axis 0'] != -1)\n    print(f\"Submission saved. Motors detected: {motors_found}/{total_tomos}\")\n    return submission_df\n\ndef plot_test_predictions(submission_df, test_dataset):\n    valid_preds = submission_df[submission_df['Motor axis 0'] != -1]\n    shape_map = {item['tomo_id']: item['orig_shape'] for item in test_dataset.data}\n\n    for idx, row in valid_preds.iterrows():\n        tomo_id = row['tomo_id']\n        z = row['Motor axis 0']\n        y = row['Motor axis 1']\n        x = row['Motor axis 2']\n\n        slice_path = os.path.join('/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/test', tomo_id, f'slice_{z:04d}.jpg')\n        if not os.path.exists(slice_path):\n            print(f\"Slice {slice_path} not found, skipping.\")\n            continue\n\n        img = Image.open(slice_path).convert('L')\n        img_array = np.array(img)\n\n        orig_shape = shape_map[tomo_id]\n        orig_height, orig_width = orig_shape[0].item(), orig_shape[1].item()\n\n        plt.figure(figsize=(8, 8))\n        plt.imshow(img_array, cmap='gray')\n        plt.plot(x, y, 'ro', label='Predicted Motor', markersize=10)\n        plt.title(f\"Tomogram: {tomo_id}, Z: {z}, Shape: {int(orig_height)}x{int(orig_width)}\")\n        plt.legend()\n        plt.axis('off')\n        plt.show()\n\n    fig = plt.figure(figsize=(10, 8))\n    ax = fig.add_subplot(111, projection='3d')\n    \n    zs = valid_preds['Motor axis 0']\n    ys = valid_preds['Motor axis 1']\n    xs = valid_preds['Motor axis 2']\n\n    ax.scatter(xs, ys, zs, c='r', marker='o', label='Predicted Motors')\n    ax.set_xlabel('X (Motor axis 2)')\n    ax.set_ylabel('Y (Motor axis 1)')\n    ax.set_zlabel('Z (Motor axis 0)')\n    ax.set_title('3D Distribution of Predicted Motors in Test Set')\n    ax.legend()\n    plt.show()\n\ndef main():\n    device = 'cuda:0' if torch.cuda.is_available() else 'cpu'\n    if device.startswith('cuda'):\n        torch.backends.cudnn.benchmark = True\n        torch.backends.cuda.matmul.allow_tf32 = True\n        torch.backends.cudnn.allow_tf32 = True\n        print(f\"Using GPU: {torch.cuda.get_device_name(0)}\")\n    else:\n        print(\"Using CPU\")\n\n    model = OptimizedCenterNet2D().to(device)\n\n    trust_region = 4\n    full_train_dataset = FlagellarDataset(\n        csv_file='/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/train_labels.csv',\n        root_dir='/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/train',\n        new_size=(256, 256),\n        trust_region=trust_region\n    )\n\n    # Get unique tomo_ids and their corresponding indices\n    tomo_id_to_indices = {}\n    for idx, item in enumerate(full_train_dataset.data):\n        tomo_id = item['tomo_id']\n        if tomo_id not in tomo_id_to_indices:\n            tomo_id_to_indices[tomo_id] = []\n        tomo_id_to_indices[tomo_id].append(idx)\n\n    # Split tomo_ids into train and validation sets\n    unique_tomo_ids = list(tomo_id_to_indices.keys())\n    train_tomo_ids, val_tomo_ids = train_test_split(\n        unique_tomo_ids, test_size=0.2, random_state=42\n    )\n\n    # Map tomo_ids back to dataset indices\n    train_idx = []\n    val_idx = []\n    for tomo_id in train_tomo_ids:\n        train_idx.extend(tomo_id_to_indices[tomo_id])\n    for tomo_id in val_tomo_ids:\n        val_idx.extend(tomo_id_to_indices[tomo_id])\n\n    # Create train and validation subsets\n    train_dataset = torch.utils.data.Subset(full_train_dataset, train_idx)\n    val_dataset = torch.utils.data.Subset(full_train_dataset, val_idx)\n\n    # Verify no overlap in tomo_ids\n    train_tomo_set = set(train_dataset.dataset.data[i]['tomo_id'] for i in train_idx)\n    val_tomo_set = set(val_dataset.dataset.data[i]['tomo_id'] for i in val_idx)\n    overlap = train_tomo_set.intersection(val_tomo_set)\n    assert len(overlap) == 0, f\"Overlap detected in tomo_ids: {overlap}\"\n    print(f\"Train tomo_ids: {len(train_tomo_set)}, Val tomo_ids: {len(val_tomo_set)}\")\n\n    # Create data loaders\n    train_loader = DataLoader(\n        train_dataset, batch_size=8, shuffle=True,\n        collate_fn=custom_collate_fn, num_workers=0, pin_memory=True\n    )\n    val_loader = DataLoader(\n        val_dataset, batch_size=8, shuffle=False,\n        collate_fn=custom_collate_fn, num_workers=0, pin_memory=True\n    )\n\n    test_dataset = FlagellarDataset(\n        csv_file=None,\n        root_dir='/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/test',\n        new_size=(256, 256),\n        trust_region=trust_region,\n        is_test=True\n    )\n\n    train_model(model, train_loader, val_loader, num_epochs=150, device=device)\n    submission_df = generate_submission(test_dataset, model, device)\n    plot_test_predictions(submission_df, test_dataset)\n    \nif __name__ == \"__main__\":\n    start_time = time.time()\n    main()\n    print(f\"Total execution time: {(time.time() - start_time)/60:.2f} minutes\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T07:09:53.381628Z","iopub.execute_input":"2025-03-25T07:09:53.382010Z","iopub.status.idle":"2025-03-25T09:21:23.718262Z","shell.execute_reply.started":"2025-03-25T07:09:53.381977Z","shell.execute_reply":"2025-03-25T09:21:23.717399Z"}},"outputs":[],"execution_count":null}]}