{"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":"none","dataSources":[{"sourceId":91249,"databundleVersionId":11294684,"sourceType":"competition"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport pandas as pd\nimport numpy as np\nfrom PIL import Image\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as transforms\nimport glob\nimport math\nimport random\nfrom scipy.ndimage import gaussian_filter\n\n# Define constants and hyperparameters (will need tuning)\nPATCH_SIZE = (96, 96) # Height, Width for each slice in a patch\nSTACK_SIZE = 7 # Number of slices in a patch (should be odd)\nHALF_STACK = STACK_SIZE // 2\nINPUT_SHAPE = (STACK_SIZE, PATCH_SIZE[0], PATCH_SIZE[1])\n\n# Training hyperparameters\nBATCH_SIZE = 32\nLEARNING_RATE = 0.001\nNUM_EPOCHS = 10\nPOSITIVE_NEGATIVE_RATIO = 1.0 # Ratio of positive to negative samples in training\n# Factor to weigh regression loss against classification loss\nREGRESSION_LOSS_WEIGHT = 1.0\n\n# Inference hyperparameters\n# Sliding window stride in (Z, Y, X)\nSLIDING_WINDOW_STRIDE = (HALF_STACK + 1, PATCH_SIZE[0] // 2, PATCH_SIZE[1] // 2)\nCLASSIFICATION_THRESHOLD = 0.5 # Threshold for motor presence\n# NMS distance threshold in voxel units - needs tuning based on voxel spacing and motor size\n# A starting point could be related to the expected size of a motor in voxels\nNMS_DISTANCE_THRESHOLD_VOXELS = 20 # Placeholder - needs tuning\n# Threshold for considering a negative sample location far from a motor during training\nDISTANCE_THRESHOLD_VOXELS_NEG = 50 # Placeholder - needs tuning\n\n# Data directories\nTRAIN_DIR = '/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/train'\nTEST_DIR = '/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/test'\nTRAIN_LABELS = '/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/train_labels.csv'\nSAMPLE_SUBMISSION = '/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/sample_submission.csv'\n\n# Custom Dataset\nclass FlagellaDataset(Dataset):\n    def __init__(self, data_dir, labels_file, patch_size, stack_size, positive_negative_ratio=1.0, is_train=True, distance_threshold_voxels_neg=50):\n        self.data_dir = data_dir\n        self.labels_file = labels_file\n        self.patch_size = patch_size\n        self.stack_size = stack_size\n        self.half_stack = stack_size // 2\n        self.is_train = is_train\n        self.positive_negative_ratio = positive_negative_ratio\n        self.distance_threshold_voxels_neg = distance_threshold_voxels_neg\n\n        self.tomograms = {}\n        # Store column mapping as an instance attribute\n        self.col_mapping = {}\n        self.load_tomogram_info()\n\n\n        if self.is_train:\n            self.labels = pd.read_csv(labels_file)\n             # Ensure column mapping is done before accessing labels in get_positive_samples\n            if not self.col_mapping: # This should be populated by load_tomogram_info for is_train=True\n                 print(\"Error: Column mapping not initialized for training data after load_tomogram_info.\")\n                 # This indicates a failure in load_tomogram_info that wasn't caught by the initial check\n                 raise RuntimeError(\"Failed to initialize column mapping.\")\n\n\n            self.positive_samples = self.get_positive_samples()\n            self.negative_samples = self.get_negative_samples()\n            self.samples = self.balance_samples()\n        else:\n            self.test_tomograms = glob.glob(os.path.join(self.data_dir, '*'))\n            self.test_tomograms = [os.path.basename(t) for t in self.test_tomograms] # Get just the directory names\n            self.samples = [{'tomo_id': tomo_id, 'is_motor': -1} for tomo_id in self.test_tomograms] # Dummy samples for test iteration\n\n\n    def load_tomogram_info(self):\n        # This function needs to load tomogram shape and voxel spacing for both train and test sets.\n        # For training, this info is in train_labels.csv.\n        # For test, this info is not in a readily available CSV in the provided structure but is likely accessible\n        # during the competition run (e.g., by inspecting the image files).\n        # We will simulate loading this info for test by reading the first image slice and assuming a placeholder\n        # voxel spacing. **In a real competition, the method to get test tomogram info needs to be confirmed.**\n\n        if self.is_train:\n            labels_df = pd.read_csv(self.labels_file)\n\n            # --- Start of direct column check and user instruction ---\n            # Define expected column names and their *potential* actual names in the CSV.\n            # Initially, assume actual names are the same as expected names (value is None).\n            # If a KeyError occurs, the user needs to update the values in this dictionary\n            # with the exact column names found in their CSV.\n            required_cols = {\n                'tomo_id': None, # Update value with actual column name from CSV if different\n                'Motor axis 0': None, # Update value with actual column name from CSV if different\n                'Motor axis 1': None, # Update value with actual column name from CSV if different\n                'Motor axis 2': None, # Update value with actual column name from CSV if different\n                'Array shape axis 0': None, # Update value with actual column name from CSV if different\n                'Array shape axis 1': None, # Update value with actual column name from CSV if different\n                'Array shape axis 2': None, # Update value with actual column name from CSV if different\n                'Voxel spacing': None # Update value with actual column name from CSV if different\n            }\n            actual_cols = labels_df.columns.tolist()\n\n            print(\"--- Checking columns in train_labels.csv ---\")\n            print(\"Actual columns found:\", actual_cols)\n            print(\"Expected columns (based on competition data description):\")\n            for col in required_cols.keys():\n                print(f\"- {col}\")\n\n            self.col_mapping = {}\n            missing_cols = []\n\n            for expected_col, actual_name_override in required_cols.items():\n                # Use the override name if provided, otherwise use the expected name\n                name_to_check = expected_col if actual_name_override is None else actual_name_override\n\n                if name_to_check in actual_cols:\n                    self.col_mapping[expected_col] = name_to_check\n                else:\n                    missing_cols.append(expected_col)\n\n            if missing_cols:\n                print(\"\\nError: The following required columns were NOT found in train_labels.csv:\")\n                for col in missing_cols:\n                    print(f\"- {col}\")\n                print(\"\\nACTION REQUIRED:\")\n                print(\"The column names in your train_labels.csv file do not exactly match the names the code is looking for.\")\n                print(\"Please compare the 'Actual columns found' list above with the 'Expected columns' list.\")\n                print(\"If the names are different (e.g., different capitalization, spaces, or underscores),\")\n                # Corrected print statement - removed undefined variables X and Y\n                print(f\"you MUST update the `required_cols` dictionary in the FlagellaDataset class's `load_tomogram_info` method.\")\n                print(f\"For each missing column from the 'Expected columns' list (e.g., 'Array shape axis 0'),\")\n                print(f\"find its corresponding exact name in the 'Actual columns found' list (e.g., 'Array_shape_axis_0').\")\n                print(f\"Then, in the `required_cols` dictionary, change the line that looks like:\")\n                print(f\"    'Expected Column Name': None,\")\n                print(f\"to use the actual column name found in your CSV. For example, like this:\")\n                print(f\"    'Expected Column Name': 'Actual Column Name from CSV',\")\n                print(\"For example, if the expected column 'Array shape axis 0' is actually named 'Array_shape_axis_0' in your file,\")\n                print(\"change the line in `required_cols` from `'Array shape axis 0': None,` to `'Array shape axis 0': 'Array_shape_axis_0',`.\")\n                print(\"Do this for all missing columns listed above.\")\n                print(\"Once you have updated the `required_cols` dictionary with the correct names from your CSV, run the code again.\")\n                raise KeyError(\"Missing required columns in train_labels.csv. Please update column names in the code as instructed above.\")\n\n            print(\"\\n--- All required columns found and mapped successfully. ---\")\n            for expected, actual in self.col_mapping.items():\n                 print(f\"  '{expected}' -> '{actual}'\")\n\n            # --- End of direct column check and user instruction ---\n\n\n            for _, row in labels_df.iterrows():\n                tomo_id = row[self.col_mapping['tomo_id']]\n                if tomo_id not in self.tomograms:\n                    tomo_path = os.path.join(self.data_dir, tomo_id)\n                    # Assuming all motors in a tomogram have the same shape and spacing\n                    self.tomograms[tomo_id] = {\n                        'path': tomo_path,\n                        'shape': (int(row[self.col_mapping['Array shape axis 0']]), int(row[self.col_mapping['Array shape axis 1']]), int(row[self.col_mapping['Array shape axis 2']])), # Z, Y, X\n                        'voxel_spacing': float(row[self.col_mapping['Voxel spacing']])\n                    }\n        else:\n             test_tomograms_paths = glob.glob(os.path.join(self.data_dir, '*'))\n             for tomo_path in test_tomograms_paths:\n                 tomo_id = os.path.basename(tomo_path)\n                 slice_files = sorted(glob.glob(os.path.join(tomo_path, '*.jpeg')))\n                 if slice_files:\n                     img = Image.open(slice_files[0])\n                     width, height = img.size # PIL gives (width, height)\n                     depth = len(slice_files)\n                     # Placeholder voxel spacing for test. **Crucial to get actual test voxel spacing.**\n                     # Assuming a default of 1.0 if not available.\n                     # In a real scenario, try to infer from metadata or competition guidelines.\n                     voxel_spacing = 1.0 # Default placeholder\n                     # If there's a way to get test voxel spacing per tomogram, add it here.\n                     self.tomograms[tomo_id] = {\n                         'path': tomo_path,\n                         'shape': (depth, height, width), # Z, Y, X\n                         'voxel_spacing': voxel_spacing # Placeholder\n                     }\n\n\n    def get_positive_samples(self):\n        positive_samples = []\n        # Use the mapped column names\n        # Ensure col_mapping is populated\n        if not self.col_mapping:\n             print(\"Error: Column mapping not available when getting positive samples.\")\n             raise AttributeError(\"Column mapping not initialized. Check load_tomogram_info.\")\n\n        for index, row in self.labels.iterrows():\n            tomo_id = row[self.col_mapping['tomo_id']]\n            motor_coords = (float(row[self.col_mapping['Motor axis 0']]), float(row[self.col_mapping['Motor axis 1']]), float(row[self.col_mapping['Motor axis 2']])) # Z, Y, X\n            positive_samples.append({'tomo_id': tomo_id, 'motor_coords': motor_coords, 'is_motor': 1})\n        return positive_samples\n\n    def get_negative_samples(self):\n        negative_samples = []\n        # Ensure col_mapping is populated for accessing 'tomo_id' from self.labels if needed,\n        # though the main tomogram info comes from self.tomograms which is populated using the mapping.\n        # If self.labels is not loaded (e.g., is_train is False), this part won't be reached in this method's logic.\n        # If self.labels is loaded but mapping failed, an error would have been raised in __init__.\n        # So we can assume self.labels and self.col_mapping are available here if is_train is True.\n        tomograms_with_motors = set(self.labels[self.col_mapping['tomo_id']].unique()) if hasattr(self, 'labels') and self.col_mapping else set()\n\n\n        all_tomograms = set(self.tomograms.keys())\n        tomograms_without_motors = list(all_tomograms - tomograms_with_motors)\n\n        # Add negative samples from tomograms without motors\n        for tomo_id in tomograms_without_motors:\n             tomo_info = self.tomograms[tomo_id]\n             depth, height, width = tomo_info['shape']\n             # Added a check to ensure dimensions are positive before calculating num_samples\n             if depth is not None and height is not None and width is not None and depth > 0 and height > 0 and width > 0:\n                 num_samples = max(10, (depth * height * width) // (self.patch_size[0] * self.patch_size[1] * self.stack_size) // 5) # Heuristic\n                 for _ in range(num_samples):\n                     z = random.randint(0, depth - 1)\n                     y = random.randint(0, height - 1)\n                     x = random.randint(0, width - 1)\n                     negative_samples.append({'tomo_id': tomo_id, 'motor_coords': (z, y, x), 'is_motor': 0})\n             else:\n                 print(f\"Warning: Tomogram {tomo_id} has invalid dimensions: {tomo_info.get('shape', 'N/A')}. Skipping negative sampling from this tomogram.\")\n\n\n        # Add negative samples from regions far from motors in tomograms with motors\n        if hasattr(self, 'labels') and self.col_mapping: # Only do this if labels were loaded successfully\n            for tomo_id in tomograms_with_motors:\n                tomo_info = self.tomograms[tomo_id]\n                depth, height, width = tomo_info['shape']\n                # Ensure motor coordinate columns are in col_mapping before accessing labels\n                motor_cols = [self.col_mapping.get('Motor axis 0'), self.col_mapping.get('Motor axis 1'), self.col_mapping.get('Motor axis 2')]\n                if None in motor_cols:\n                     print(f\"Warning: Motor coordinate columns not fully mapped for tomogram {tomo_id}. Skipping negative sampling from regions near motors for this tomogram.\")\n                     continue # Skip if motor columns weren't mapped\n\n                motor_locations = self.labels[self.labels[self.col_mapping['tomo_id']] == tomo_id][motor_cols].values # Z, Y, X\n\n                # Sample locations and check distance to any motor\n                # Added a check to ensure dimensions are positive before calculating num_samples\n                if depth is not None and height is not None and width is not None and depth > 0 and height > 0 and width > 0:\n                    num_samples = max(10, (depth * height * width) // (self.patch_size[0] * self.patch_size[1] * self.stack_size) // 10) # Heuristic\n                    for _ in range(num_samples):\n                        z = random.randint(0, depth - 1)\n                        y = random.randint(0, height - 1)\n                        x = random.randint(0, width - 1)\n\n                        # Check if this random location is far from any motor in voxel space\n                        min_distance = float('inf')\n                        for motor_z, motor_y, motor_x in motor_locations:\n                             distance = np.sqrt((z - motor_z)**2 + (y - motor_y)**2 + (x - motor_x)**2)\n                             min_distance = min(min_distance, distance)\n\n                        if min_distance > self.distance_threshold_voxels_neg:\n                             negative_samples.append({'tomo_id': tomo_id, 'motor_coords': (z, y, x), 'is_motor': 0})\n                else:\n                     print(f\"Warning: Tomogram {tomo_id} has invalid dimensions: {tomo_info.get('shape', 'N/A')}. Skipping negative sampling from this tomogram.\")\n\n\n        # Limit the number of negative samples to manage dataset size\n        # Only limit if there are positive samples, otherwise keep all negatives found\n        if hasattr(self, 'positive_samples') and len(self.positive_samples) > 0:\n             max_negative_samples = int(len(self.positive_samples) * self.positive_negative_ratio * 2) # Allow up to twice the target ratio\n             if len(negative_samples) > max_negative_samples:\n                  negative_samples = random.sample(negative_samples, max_negative_samples)\n\n\n        return negative_samples\n\n    def balance_samples(self):\n        # Ensure positive_samples and negative_samples are populated before balancing\n        num_positive = len(getattr(self, 'positive_samples', []))\n        num_negative = len(getattr(self, 'negative_samples', []))\n\n        if num_positive == 0:\n             print(\"Warning: No positive samples found for balancing. Using available negative samples.\")\n             return getattr(self, 'negative_samples', [])\n\n        if num_positive * self.positive_negative_ratio < num_negative:\n            # Undersample negative samples\n            sampled_negative_indices = random.sample(range(num_negative), int(num_positive * self.positive_negative_ratio))\n            balanced_negative_samples = [self.negative_samples[i] for i in sampled_negative_indices]\n            print(f\"Undersampled negative samples from {num_negative} to {len(balanced_negative_samples)}\")\n        else:\n            # Use all negative samples\n            balanced_negative_samples = self.negative_samples\n            print(f\"Using all {num_negative} negative samples.\")\n\n        # Combine positive and balanced negative samples\n        all_samples = self.positive_samples + balanced_negative_samples\n        random.shuffle(all_samples)\n        print(f\"Total training samples: {len(all_samples)}\")\n        return all_samples\n\n\n    def load_slice(self, tomo_id, slice_idx):\n        # Load a single slice (JPEG image)\n        if tomo_id not in self.tomograms:\n             print(f\"Error: Tomogram info not available for {tomo_id} in load_slice.\")\n             # Return a dummy black slice if tomogram info is missing\n             # This case should ideally not happen if load_tomogram_info ran correctly\n             return np.zeros(PATCH_SIZE, dtype=np.uint8) # Return based on patch size as a fallback\n\n\n        tomo_path = self.tomograms[tomo_id]['path']\n        # Assuming slice file names are zero-padded and start from 0 or 1. Need to check data specifics.\n        # Assuming 0-indexed and 4-digit padding for now.\n        slice_filename = os.path.join(tomo_path, f'{slice_idx:04d}.jpeg')\n        try:\n            img = Image.open(slice_filename).convert('L') # Convert to grayscale\n            return np.array(img)\n        except FileNotFoundError:\n            # print(f\"Warning: Slice file not found: {slice_filename}\") # Avoid too many warnings during training\n            # Return an array of zeros if slice is missing (e.g., near boundaries when centering)\n            # This assumes all slices in a tomogram have the same H, W, which should be true.\n            # Use the tomogram's known height and width\n            depth, height, width = self.tomograms[tomo_id]['shape']\n            return np.zeros((height, width), dtype=np.uint8)\n\n\n    def get_patch(self, tomo_id, center_coords):\n        # Extract a patch of slices centered around the given coordinates (z, y, x)\n        if tomo_id not in self.tomograms:\n             print(f\"Error: Tomogram info not available for {tomo_id} in get_patch.\")\n             # Return a dummy black patch if tomogram info is missing\n             # Shape: [STACK_SIZE, PATCH_SIZE[0], PATCH_SIZE[1]]\n             return np.zeros((self.stack_size, self.patch_size[0], self.patch_size[1]), dtype=np.uint8)\n\n\n        depth, height, width = self.tomograms[tomo_id]['shape']\n        center_z, center_y, center_x = center_coords\n\n        # Determine the z-range of slices to load\n        z_start = int(center_z - self.half_stack)\n        z_end = int(center_z + self.half_stack + 1)\n\n        patch_slices = []\n        for z in range(z_start, z_end):\n            # Ensure z is within tomogram bounds before loading, otherwise use a black slice\n            if 0 <= z < depth:\n                 slice_data = self.load_slice(tomo_id, z)\n                 patch_slices.append(slice_data) # load_slice already handles missing files by returning zeros\n            else:\n                # If z is out of bounds, append a black slice\n                # This assumes all slices in a tomogram have the same H, W\n                 patch_slices.append(np.zeros((height, width), dtype=np.uint8))\n\n\n        # Stack the slices\n        patch_stack = np.stack(patch_slices, axis=0) # Shape: [STACK_SIZE, H, W]\n\n        # Extract the 2D patch from each slice centered around y, x\n        # Calculate the start and end indices for cropping in Y and X\n        patch_2d_start_y = int(center_y - self.patch_size[0] // 2)\n        patch_2d_end_y = patch_2d_start_y + self.patch_size[0]\n        patch_2d_start_x = int(center_x - self.patch_size[1] // 2)\n        patch_2d_end_x = patch_2d_start_x + self.patch_size[1]\n\n        # Calculate padding needed for Y and X dimensions\n        pad_y_before = max(0, -patch_2d_start_y)\n        pad_y_after = max(0, patch_2d_end_y - height)\n        pad_x_before = max(0, -patch_2d_start_x)\n        pad_x_after = max(0, patch_2d_end_x - width)\n\n        # Apply padding\n        padded_patch_stack = np.pad(patch_stack, (\n            (0, 0), # No padding in Z\n            (pad_y_before, pad_y_after), # Padding in Y\n            (pad_x_before, pad_x_after)  # Padding in X\n        ), mode='constant', constant_values=0)\n\n        # Calculate the crop indices on the padded stack\n        # The start indices on the padded stack are the amount of padding added before\n        crop_start_y_padded = pad_y_before\n        crop_end_y_padded = crop_start_y_padded + self.patch_size[0]\n        crop_start_x_padded = pad_x_before\n        crop_end_x_padded = crop_start_x_padded + self.patch_size[1]\n\n\n        # Crop the padded patch to the desired size\n        cropped_patch = padded_patch_stack[:, crop_start_y_padded:crop_end_y_padded, crop_start_x_padded:crop_start_x_padded + self.patch_size[1]] # Corrected x-cropping\n\n\n        return cropped_patch\n\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        sample = self.samples[idx]\n        tomo_id = sample['tomo_id']\n        is_motor = sample['is_motor']\n\n        if is_motor == 1: # Positive sample\n            motor_coords = sample['motor_coords'] # Z, Y, X in tomogram voxel space\n            center_coords = motor_coords # Center patch at motor location\n            patch = self.get_patch(tomo_id, center_coords)\n            label = 1.0 # Use float for BCEWithLogitsLoss\n            # Regression target: motor coordinates relative to the patch's top-left-front corner, in voxel units\n            # The patch extraction is centered at center_coords (z, y, x).\n            # The top-left-front corner of the extracted patch volume in the original tomogram\n            # is approximately (center_z - half_stack, center_y - patch_size[0]//2, center_x - patch_size[1]//2).\n            # We need the motor_coords (z_motor, y_motor, x_motor) relative to this corner.\n            patch_origin_in_tomo = (\n                center_coords[0] - self.half_stack,\n                center_coords[1] - self.patch_size[0] // 2,\n                center_coords[2] - self.patch_size[1] // 2\n            )\n            regression_target = torch.tensor([\n                motor_coords[0] - patch_origin_in_tomo[0],\n                motor_coords[1] - patch_origin_in_tomo[1],\n                motor_coords[2] - patch_origin_in_tomo[2]\n            ], dtype=torch.float32)\n\n\n        elif is_motor == 0: # Negative sample\n            center_coords = sample['motor_coords'] # Random coordinates in tomogram voxel space\n            patch = self.get_patch(tomo_id, center_coords)\n            label = 0.0 # Use float for BCEWithLogitsLoss\n            regression_target = torch.zeros(3, dtype=torch.float32) # Dummy target for negative samples\n\n        else: # Test sample (is_motor == -1) - This part is not used with DataLoader for training/validation\n             # It's a placeholder for the sample structure, actual patch loading for test is in predict function\n             # Returning dummy values here to avoid errors if this is somehow accessed.\n             patch = np.zeros((self.stack_size, self.patch_size[0], self.patch_size[1]), dtype=np.float32)\n             patch_tensor = torch.tensor(patch, dtype=torch.float32)\n             label = -1.0 # Indicate this is a test sample\n             regression_target = torch.zeros(3, dtype=torch.float32) # Dummy target\n\n\n        # Normalize patch to [0, 1]\n        patch = patch.astype(np.float32) / 255.0\n        patch_tensor = torch.tensor(patch, dtype=torch.float32)\n        # The model expects input shape [BATCH_SIZE, STACK_SIZE, H, W]\n        # The patch_tensor is currently [STACK_SIZE, H, W]\n\n        return patch_tensor, torch.tensor(label, dtype=torch.float32), regression_target\n\n\n# Define the 2.5D CNN Model\nclass FlagellaCNN(nn.Module):\n    def __init__(self, stack_size):\n        super(FlagellaCNN, self).__init__()\n        # Using STACK_SIZE as input channels for 2D convolutions\n        self.features = nn.Sequential(\n            nn.Conv2d(stack_size, 32, kernel_size=3, padding=1),\n            nn.ReLU(),\n            nn.MaxPool2d(kernel_size=2, stride=2),\n            nn.Conv2d(32, 64, kernel_size=3, padding=1),\n            nn.ReLU(),\n            nn.MaxPool2d(kernel_size=2, stride=2),\n            nn.Conv2d(64, 128, kernel_size=3, padding=1),\n            nn.ReLU(),\n            nn.MaxPool2d(kernel_size=2, stride=2)\n        )\n\n        # Calculate the size of the flattened features\n        # Assuming input size (STACK_SIZE, PATCH_SIZE[0], PATCH_SIZE[1])\n        # After 3 MaxPool2d with kernel 2, stride 2, the spatial dimensions are divided by 2^3 = 8\n        # Patch spatial size is PATCH_SIZE[0] x PATCH_SIZE[1]\n        flattened_size = 128 * (PATCH_SIZE[0] // 8) * (PATCH_SIZE[1] // 8)\n\n\n        self.classifier = nn.Sequential(\n            nn.Linear(flattened_size, 64),\n            nn.ReLU(),\n            nn.Dropout(0.5),\n            nn.Linear(64, 1) # Binary classification output (motor or no motor)\n        )\n\n        self.regressor = nn.Sequential(\n            nn.Linear(flattened_size, 64),\n            nn.ReLU(),\n            nn.Dropout(0.5),\n            nn.Linear(64, 3) # Regression output (z, y, x coordinates relative to patch)\n        )\n\n    def forward(self, x):\n        # Input shape: [BATCH_SIZE, STACK_SIZE, H, W]\n        x = self.features(x)\n        x = torch.flatten(x, 1) # Flatten for the fully connected layers\n\n        classification_output = self.classifier(x)\n        regression_output = self.regressor(x)\n\n        return classification_output, regression_output\n\n# Combined Loss Function\nclass CombinedLoss(nn.Module):\n    def __init__(self, regression_loss_weight=1.0):\n        super(CombinedLoss, self).__init__()\n        self.classification_criterion = nn.BCEWithLogitsLoss() # Uses Sigmoid internally\n        self.regression_criterion = nn.MSELoss() # Mean Squared Error for regression\n        self.regression_loss_weight = regression_loss_weight\n\n    def forward(self, classification_preds, regression_preds, classification_labels, regression_targets):\n        classification_loss = self.classification_criterion(classification_preds.squeeze(-1), classification_labels) # Squeeze to match shape\n\n        # Apply regression loss only to positive samples\n        positive_mask = (classification_labels == 1)\n        if torch.sum(positive_mask) > 0:\n            regression_loss = self.regression_criterion(regression_preds[positive_mask], regression_targets[positive_mask])\n            total_loss = classification_loss + self.regression_loss_weight * regression_loss\n        else:\n            regression_loss = torch.tensor(0.0).to(classification_preds.device)\n            total_loss = classification_loss\n\n        return total_loss, classification_loss, regression_loss\n\n# Training Function\ndef train_model(model, dataloader, criterion, optimizer, num_epochs):\n    model.train()\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    model.to(device)\n\n    for epoch in range(num_epochs):\n        running_loss = 0.0\n        running_cls_loss = 0.0\n        running_reg_loss = 0.0\n        for patches, labels, regression_targets in dataloader:\n            # Filter out dummy test samples if any somehow got into the loader\n            train_indices = (labels != -1)\n            if torch.sum(train_indices) == 0:\n                 continue # Skip batch if it only contains dummy test samples\n\n            patches = patches[train_indices].to(device)\n            labels = labels[train_indices].to(device)\n            regression_targets = regression_targets[train_indices].to(device)\n\n            optimizer.zero_grad()\n\n            classification_preds, regression_preds = model(patches)\n\n            total_loss, classification_loss, regression_loss = criterion(classification_preds, regression_preds, labels, regression_targets)\n\n            total_loss.backward()\n            optimizer.step()\n\n            running_loss += total_loss.item()\n            running_cls_loss += classification_loss.item()\n            running_reg_loss += regression_loss.item()\n\n        # Calculate epoch loss based on actual number of training batches\n        num_train_batches = len(dataloader)\n        if num_train_batches > 0:\n            epoch_loss = running_loss / num_train_batches\n            epoch_cls_loss = running_cls_loss / num_train_batches\n            epoch_reg_loss = running_reg_loss / num_train_batches\n            print(f\"Epoch {epoch+1}/{num_epochs}, Loss: {epoch_loss:.4f}, Cls Loss: {epoch_cls_loss:.4f}, Reg Loss: {epoch_reg_loss:.4f}\")\n        else:\n             print(f\"Epoch {epoch+1}/{num_epochs}, No training batches processed.\")\n\n\n    print(\"Training finished.\")\n\n# Inference Function\ndef predict(model, test_data_dir, tomogram_info, patch_size, stack_size, sliding_window_stride, classification_threshold, nms_distance_threshold_voxels):\n    model.eval()\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    model.to(device)\n\n    predictions = [] # To store predictions for all tomograms\n\n    test_tomograms = glob.glob(os.path.join(test_data_dir, '*'))\n    test_tomograms = [os.path.basename(t) for t in test_tomograms]\n\n    # Create a dummy dataset instance to use its methods like load_slice and get_patch\n    # This is a workaround as these methods are tied to the instance's attributes like tomograms info\n    # In a more structured approach, these might be utility functions or part of a data loader for inference.\n    # We need tomogram_info available, which is passed to this predict function.\n    # Let's pass the tomogram_info explicitly to a helper class or directly use it here.\n    # Creating a dummy dataset instance is simpler for now.\n    # It requires patch_size and stack_size, but other init parameters don't strictly matter for load_slice/get_patch\n    # Set is_train=False to avoid attempting to load/process labels file\n    dummy_dataset = FlagellaDataset(data_dir=test_data_dir, labels_file=TRAIN_LABELS, patch_size=patch_size, stack_size=stack_size, is_train=False)\n    # Update the tomograms info in the dummy dataset with the passed tomogram_info for test data\n    dummy_dataset.tomograms.update(tomogram_info)\n    # Ensure col_mapping is empty or not used for test data within the dataset methods if they are called.\n    # Since load_tomogram_info was called with is_train=False, col_mapping should not be populated with train mapping.\n    # We don't need col_mapping for test data inference in this logic.\n\n\n    for tomo_id in test_tomograms:\n        print(f\"Processing tomogram: {tomo_id}\")\n        tomo_path = os.path.join(test_data_dir, tomo_id)\n        # Get tomogram shape and voxel spacing - **Need actual values for test set**\n        # For this simulation, use the placeholder from tomogram_info\n        if tomo_id not in dummy_dataset.tomograms: # Use dummy_dataset's tomograms info\n             print(f\"Warning: Tomogram info not found for {tomo_id}. Skipping.\")\n             predictions.append({'tomo_id': tomo_id, 'Motor axis 0': -1, 'Motor axis 1': -1, 'Motor axis 2': -1})\n             continue\n\n        tomo_info = dummy_dataset.tomograms[tomo_id]\n        depth, height, width = tomo_info['shape']\n        # Voxel spacing is needed for evaluation, but model predicts in voxel space.\n        # Keep it here as a reminder that it might be needed for distance calculations if NMS was in Angstroms.\n        voxel_spacing = tomo_info['voxel_spacing'] # Placeholder\n\n        tomo_predictions = [] # To store confident predictions for the current tomogram (predicted_voxel_coords, confidence)\n\n        # Sliding window approach\n        # Ensure the window slides completely through the tomogram\n        # Handle cases where any dimension is zero or negative first\n        if depth is None or height is None or width is None or depth <= 0 or height <= 0 or width <= 0:\n             print(f\"Warning: Tomogram {tomo_id} has invalid dimensions: {tomo_info.get('shape', 'N/A')}. Skipping inference.\")\n             predictions.append({'tomo_id': tomo_id, 'Motor axis 0': -1, 'Motor axis 1': -1, 'Motor axis 2': -1})\n             continue\n\n\n        z_steps = range(0, depth - stack_size + 1, sliding_window_stride[0])\n        if depth > 0 and depth < stack_size: # Handle cases where tomogram is smaller than patch size in Z\n             z_steps = [0] # Start at 0 if too small\n        elif depth <= 0: # Should be caught by the earlier check, but for safety\n             z_steps = [] # No steps\n\n        y_steps = range(0, height - patch_size[0] + 1, sliding_window_stride[1])\n        if height > 0 and height < patch_size[0]: # Handle cases where tomogram is smaller than patch size in Y\n             y_steps = [0] # Start at 0 if too small\n        elif height <= 0: # Should be caught by the earlier check\n             y_steps = [] # No steps\n\n        x_steps = range(0, width - patch_size[1] + 1, sliding_window_stride[2])\n        if width > 0 and width < patch_size[1]: # Handle cases where tomogram is smaller than patch size in X\n             x_steps = [0] # Start at 0 if too small\n        elif width <= 0: # Should be caught by the earlier check\n             x_steps = [] # No steps\n\n        # Re-check if steps are empty after considering small dimensions\n        if not z_steps or not y_steps or not x_steps:\n             print(f\"Warning: Tomogram {tomo_id} dimensions ({depth}, {height}, {width}) are incompatible with patch/stride size after boundary handling. Skipping inference.\")\n             predictions.append({'tomo_id': tomo_id, 'Motor axis 0': -1, 'Motor axis 1': -1, 'Motor axis 2': -1})\n             continue\n\n\n        for z in z_steps:\n            for y in y_steps:\n                for x in x_steps:\n                    # Extract patch using the dummy dataset's method\n                    # The center for get_patch should be in the middle of the current window\n                    patch_center_coords_in_tomo = (z + dummy_dataset.half_stack, y + dummy_dataset.patch_size[0]//2, x + dummy_dataset.patch_size[1]//2)\n                    # Pass dummy_dataset instance to get_patch to access its methods and attributes\n                    cropped_patch = dummy_dataset.get_patch(tomo_id, patch_center_coords_in_tomo)\n\n\n                    # Preprocess patch\n                    cropped_patch = cropped_patch.astype(np.float32) / 255.0\n                    patch_tensor = torch.tensor(cropped_patch, dtype=torch.float32).unsqueeze(0).to(device) # Add batch dimension\n\n                    # Get model predictions\n                    with torch.no_grad():\n                        classification_output, regression_output = model(patch_tensor)\n\n                    # Apply sigmoid to classification output to get confidence score\n                    confidence = torch.sigmoid(classification_output).item()\n\n                    if confidence > classification_threshold:\n                        # Motor detected, get predicted coordinates relative to patch origin\n                        predicted_coords_relative_to_patch_origin = regression_output[0].cpu().numpy() # Z, Y, X\n\n                        # Calculate the tomogram coordinates of the patch's top-left-front corner used for training target\n                        # This aligns with how the regression target was defined in the dataset's get_item\n                        # The patch origin in tomogram coordinates is (z, y, x) from the sliding window loop.\n                        patch_origin_coords_in_tomo = (z, y, x)\n\n\n                        predicted_motor_coords_voxel = (\n                            patch_origin_coords_in_tomo[0] + predicted_coords_relative_to_patch_origin[0],\n                            patch_origin_coords_in_tomo[1] + predicted_coords_relative_to_patch_origin[1],\n                            patch_origin_coords_in_tomo[2] + predicted_coords_relative_to_patch_origin[2]\n                        )\n\n\n                        tomo_predictions.append({\n                            'voxel_coords': predicted_motor_coords_voxel, # Predicted (Z, Y, X) in tomogram voxel space\n                            'confidence': confidence\n                        })\n\n        # Post-processing: Combine overlapping predictions using NMS or clustering\n        # For this example, a simple distance-based filtering and taking the most confident\n        # prediction among the remaining ones.\n        final_motor_coords_voxel = None\n        if tomo_predictions:\n            # Sort predictions by confidence\n            tomo_predictions.sort(key=lambda x: x['confidence'], reverse=True)\n\n            # Apply a simple distance-based filtering (not full NMS)\n            filtered_predictions = []\n            for pred in tomo_predictions:\n                is_overlapping = False\n                for final_pred in filtered_predictions:\n                    dist = np.sqrt(np.sum((np.array(pred['voxel_coords']) - np.array(final_pred['voxel_coords']))**2))\n                    # Using voxel_spacing here would be necessary if NMS was in Angstroms,\n                    # but we are doing NMS in voxel space as the predictions are in voxel space.\n                    # dist_angstrom = dist * voxel_spacing # Example if NMS was in Angstroms\n                    if dist < nms_distance_threshold_voxels:\n                        is_overlapping = True\n                        break\n                if not is_overlapping:\n                     filtered_predictions.append(pred)\n\n            # If after filtering, we still have predictions, take the one with highest confidence\n            if filtered_predictions:\n                 best_prediction = max(filtered_predictions, key=lambda x: x['confidence'])\n                 final_motor_coords_voxel = best_prediction['voxel_coords']\n\n\n        # Prepare for submission\n        if final_motor_coords_voxel is not None:\n            # The submission requires voxel coordinates (Motor axis 0, Motor axis 1, Motor axis 2)\n            # These correspond to Z, Y, X in the tomogram's voxel coordinate system.\n            predicted_z, predicted_y, predicted_x = final_motor_coords_voxel\n            predictions.append({'tomo_id': tomo_id, 'Motor axis 0': predicted_z, 'Motor axis 1': predicted_y, 'Motor axis 2': predicted_x})\n        else:\n            # No motor detected\n            predictions.append({'tomo_id': tomo_id, 'Motor axis 0': -1, 'Motor axis 1': -1, 'Motor axis 2': -1})\n\n    return predictions\n\n\n# Main execution block\nif __name__ == \"__main__\":\n    # Create dataset and dataloader\n    # The FlagellaDataset's __init__ now includes robust column mapping and error handling\n    # You might need to update the 'None' values in the `required_cols` dictionary within\n    # the FlagellaDataset class's `load_tomogram_info` method based on the error output.\n    train_dataset = FlagellaDataset(data_dir=TRAIN_DIR, labels_file=TRAIN_LABELS, patch_size=PATCH_SIZE, stack_size=STACK_SIZE,\n                                    positive_negative_ratio=POSITIVE_NEGATIVE_RATIO, is_train=True,\n                                    distance_threshold_voxels_neg=DISTANCE_THRESHOLD_VOXELS_NEG)\n    train_dataloader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True, num_workers=2) # Adjust num_workers as needed\n\n    # Initialize model, criterion, and optimizer\n    model = FlagellaCNN(stack_size=STACK_SIZE)\n    criterion = CombinedLoss(regression_loss_weight=REGRESSION_LOSS_WEIGHT)\n    optimizer = optim.Adam(model.parameters(), lr=LEARNING_RATE)\n\n    # Train the model\n    print(\"Starting training...\")\n    train_model(model, train_dataloader, criterion, optimizer, NUM_EPOCHS)\n    print(\"Training completed.\")\n\n    # Create test tomogram info dictionary for inference (simulating loading this info)\n    # In a real competition, this would be loaded based on available test data metadata\n    test_tomogram_info = {}\n    test_tomograms_paths = glob.glob(os.path.join(TEST_DIR, '*'))\n    for tomo_path in test_tomograms_paths:\n        tomo_id = os.path.basename(tomo_path)\n        slice_files = sorted(glob.glob(os.path.join(tomo_path, '*.jpeg')))\n        if slice_files:\n            img = Image.open(slice_files[0])\n            width, height = img.size # PIL gives (width, height)\n            depth = len(slice_files)\n            # Placeholder voxel spacing for test. **Needs actual test data info.**\n            # Assuming a default of 1.0 if not available.\n            # In a real scenario, try to infer from metadata or competition guidelines.\n            voxel_spacing = 1.0 # Default placeholder\n            # If there's a way to get test voxel spacing per tomogram, add it here.\n            test_tomogram_info[tomo_id] = {\n                'path': tomo_path,\n                'shape': (depth, height, width), # Z, Y, X\n                'voxel_spacing': voxel_spacing # Placeholder\n            }\n\n\n    # Perform inference on the test set\n    print(\"Starting inference...\")\n    # Pass the collected test_tomogram_info to the predict function\n    test_predictions = predict(model, TEST_DIR, test_tomogram_info, PATCH_SIZE, STACK_SIZE, SLIDING_WINDOW_STRIDE, CLASSIFICATION_THRESHOLD, NMS_DISTANCE_THRESHOLD_VOXELS)\n    print(\"Inference completed.\")\n\n    # Generate submission file\n    submission_df = pd.DataFrame(test_predictions)\n    submission_df = submission_df[['tomo_id', 'Motor axis 0', 'Motor axis 1', 'Motor axis 2']] # Ensure correct column order\n    submission_df.to_csv('submission.csv', index=False)\n\n    print(\"Submission file 'submission.csv' created successfully.\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-05-05T04:19:04.24828Z","iopub.execute_input":"2025-05-05T04:19:04.249214Z","iopub.status.idle":"2025-05-05T04:19:04.34884Z","shell.execute_reply.started":"2025-05-05T04:19:04.249179Z","shell.execute_reply":"2025-05-05T04:19:04.347942Z"}},"outputs":[{"name":"stdout","text":"--- Checking columns in train_labels.csv ---\nActual columns found: ['row_id', 'tomo_id', 'Motor axis 0', 'Motor axis 1', 'Motor axis 2', 'Array shape (axis 0)', 'Array shape (axis 1)', 'Array shape (axis 2)', 'Voxel spacing', 'Number of motors']\nExpected columns (based on competition data description):\n- tomo_id\n- Motor axis 0\n- Motor axis 1\n- Motor axis 2\n- Array shape axis 0\n- Array shape axis 1\n- Array shape axis 2\n- Voxel spacing\n\nError: The following required columns were NOT found in train_labels.csv:\n- Array shape axis 0\n- Array shape axis 1\n- Array shape axis 2\n\nACTION REQUIRED:\nThe column names in your train_labels.csv file do not exactly match the names the code is looking for.\nPlease compare the 'Actual columns found' list above with the 'Expected columns' list.\nIf the names are different (e.g., different capitalization, spaces, or underscores),\nyou MUST update the `required_cols` dictionary in the FlagellaDataset class's `load_tomogram_info` method.\nFor each missing column from the 'Expected columns' list (e.g., 'Array shape axis 0'),\nfind its corresponding exact name in the 'Actual columns found' list (e.g., 'Array_shape_axis_0').\nThen, in the `required_cols` dictionary, change the line that looks like:\n    'Expected Column Name': None,\nto use the actual column name found in your CSV. For example, like this:\n    'Expected Column Name': 'Actual Column Name from CSV',\nFor example, if the expected column 'Array shape axis 0' is actually named 'Array_shape_axis_0' in your file,\nchange the line in `required_cols` from `'Array shape axis 0': None,` to `'Array shape axis 0': 'Array_shape_axis_0',`.\nDo this for all missing columns listed above.\nOnce you have updated the `required_cols` dictionary with the correct names from your CSV, run the code again.\n","output_type":"stream"},{"traceback":["\u001b[0;31m---------------------------------------------------------------------------\u001b[0m","\u001b[0;31mKeyError\u001b[0m                                  Traceback (most recent call last)","\u001b[0;32m/tmp/ipykernel_1970/743568946.py\u001b[0m in \u001b[0;36m<cell line: 0>\u001b[0;34m()\u001b[0m\n\u001b[1;32m    729\u001b[0m     \u001b[0;31m# You might need to update the 'None' values in the `required_cols` dictionary within\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    730\u001b[0m     \u001b[0;31m# the FlagellaDataset class's `load_tomogram_info` method based on the error output.\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 731\u001b[0;31m     train_dataset = FlagellaDataset(data_dir=TRAIN_DIR, labels_file=TRAIN_LABELS, patch_size=PATCH_SIZE, stack_size=STACK_SIZE,\n\u001b[0m\u001b[1;32m    732\u001b[0m                                     \u001b[0mpositive_negative_ratio\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0mPOSITIVE_NEGATIVE_RATIO\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mis_train\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0;32mTrue\u001b[0m\u001b[0;34m,\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    733\u001b[0m                                     distance_threshold_voxels_neg=DISTANCE_THRESHOLD_VOXELS_NEG)\n","\u001b[0;32m/tmp/ipykernel_1970/743568946.py\u001b[0m in \u001b[0;36m__init__\u001b[0;34m(self, data_dir, labels_file, patch_size, stack_size, positive_negative_ratio, is_train, distance_threshold_voxels_neg)\u001b[0m\n\u001b[1;32m     58\u001b[0m         \u001b[0;31m# Store column mapping as an instance attribute\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m     59\u001b[0m         \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mcol_mapping\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0;34m{\u001b[0m\u001b[0;34m}\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m---> 60\u001b[0;31m         \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mload_tomogram_info\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m     61\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m     62\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/tmp/ipykernel_1970/743568946.py\u001b[0m in \u001b[0;36mload_tomogram_info\u001b[0;34m(self)\u001b[0m\n\u001b[1;32m    145\u001b[0m                 \u001b[0mprint\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m\"Do this for all missing columns listed above.\"\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    146\u001b[0m                 \u001b[0mprint\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m\"Once you have updated the `required_cols` dictionary with the correct names from your CSV, run the code again.\"\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 147\u001b[0;31m                 \u001b[0;32mraise\u001b[0m \u001b[0mKeyError\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m\"Missing required columns in train_labels.csv. Please update column names in the code as instructed above.\"\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m    148\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    149\u001b[0m             \u001b[0mprint\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m\"\\n--- All required columns found and mapped successfully. ---\"\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;31mKeyError\u001b[0m: 'Missing required columns in train_labels.csv. Please update column names in the code as instructed above.'"],"ename":"KeyError","evalue":"'Missing required columns in train_labels.csv. Please update column names in the code as instructed above.'","output_type":"error"}],"execution_count":7}]}