{"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":"## Flagellar Motor Detection in Bacteria Tomograms\n### A Deep Learning Approach for BYU - Locating Bacterial Flagellar Motors 2025\n\n### Overview\nThis notebook presents a solution for automatically detecting flagellar motors in 3D tomograms of bacteria. This approach:\n- Develops a dual-task 3D CNN that both detects and localizes motors\n- Uses a custom loss function optimized for the competition metrics","metadata":{}},{"cell_type":"code","source":"## Imports\n\nimport os\nimport glob\nimport random\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nfrom tqdm.notebook import tqdm\nfrom sklearn.model_selection import train_test_split\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-06T08:57:55.130995Z","iopub.execute_input":"2025-03-06T08:57:55.131273Z","iopub.status.idle":"2025-03-06T08:57:59.425489Z","shell.execute_reply.started":"2025-03-06T08:57:55.131250Z","shell.execute_reply":"2025-03-06T08:57:59.424551Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define global constants\nDATA_DIR = '/kaggle/input/byu-locating-bacterial-flagellar-motors-2025'\nTRAIN_CSV = os.path.join(DATA_DIR, 'train_labels.csv')\nTRAIN_DIR = os.path.join(DATA_DIR, 'train')\nTEST_DIR = os.path.join(DATA_DIR, 'test')\nOUTPUT_DIR = './'\nMODEL_DIR = './models'\n\n# Create output directories\nos.makedirs(OUTPUT_DIR, exist_ok=True) \nos.makedirs(MODEL_DIR, exist_ok=True)\n\n# Set device\nDEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {DEVICE}\")\n\n# Set seeds for reproducibility\nRANDOM_SEED = 42\nrandom.seed(RANDOM_SEED)\nnp.random.seed(RANDOM_SEED)\ntorch.manual_seed(RANDOM_SEED)\nif torch.cuda.is_available():\n    torch.cuda.manual_seed(RANDOM_SEED)\n    torch.backends.cudnn.deterministic = True","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-06T08:57:59.426739Z","iopub.execute_input":"2025-03-06T08:57:59.427209Z","iopub.status.idle":"2025-03-06T08:57:59.489401Z","shell.execute_reply.started":"2025-03-06T08:57:59.427183Z","shell.execute_reply":"2025-03-06T08:57:59.488484Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 1. Exploratory Data Analysis","metadata":{}},{"cell_type":"code","source":"\n# Load the training labels\ntrain_labels = pd.read_csv(TRAIN_CSV)\n\n# Display basic information\nprint(\"Training dataset shape:\", train_labels.shape)\nprint(\"\\nColumns in the dataset:\")\ndisplay(train_labels.columns)\n\n# Basic statistics\nprint(\"\\nBasic statistics:\")\ndisplay(train_labels.describe())\n\n# Count unique tomograms\nunique_tomo_count = train_labels['tomo_id'].nunique()\nprint(f\"\\nNumber of unique tomograms: {unique_tomo_count}\")\n\n# Check the distribution of motors per tomogram\nmotors_per_tomo = train_labels.groupby('tomo_id')['Number of motors'].first().value_counts().sort_index()\nprint(\"\\nMotors per tomogram distribution:\")\nprint(motors_per_tomo)\n\n# Display a few sample rows\nprint(\"\\nSample rows:\")\ndisplay(train_labels.head())\n\n# Check for missing values\nprint(\"\\nMissing values per column:\")\ndisplay(train_labels.isnull().sum())\n\n# Check the range of tomogram sizes\nprint(\"\\nTomogram size ranges:\")\nprint(\"Z-axis (slices):\", train_labels['Array shape (axis 0)'].min(), \"to\", train_labels['Array shape (axis 0)'].max())\nprint(\"X-axis (width):\", train_labels['Array shape (axis 1)'].min(), \"to\", train_labels['Array shape (axis 1)'].max()) \nprint(\"Y-axis (height):\", train_labels['Array shape (axis 2)'].min(), \"to\", train_labels['Array shape (axis 2)'].max())\n\n# Check voxel spacing distribution\nprint(\"\\nVoxel spacing distribution:\")\ndisplay(train_labels['Voxel spacing'].value_counts().sort_index())\n\n# Visualize a sample tomogram\n# Get one tomogram ID\nsample_tomo_id = train_labels['tomo_id'].iloc[0]\nprint(f\"\\nVisualizing sample tomogram: {sample_tomo_id}\")\n\n# Check the folder structure\nsample_folder = os.path.join(TRAIN_DIR, sample_tomo_id)\nif os.path.exists(sample_folder):\n    # List files in the folder\n    files = sorted(glob.glob(os.path.join(sample_folder, '*.jpg')))\n    print(f\"Number of slice files: {len(files)}\")\n    \n    if files:\n        # Load one slice to check dimensions\n        sample_slice = Image.open(files[0])\n        print(f\"Sample slice dimensions: {sample_slice.size}\")\n        \n        # Display a few slices\n        fig, axes = plt.subplots(1, 3, figsize=(15, 5))\n        \n        # Pick slices from beginning, middle, and end\n        indices = [0, len(files)//2, len(files)-1]\n        for i, idx in enumerate(indices):\n            img = Image.open(files[idx])\n            axes[i].imshow(img, cmap='gray')\n            axes[i].set_title(f\"Slice {idx}\")\n            axes[i].axis('off')\n        \n        plt.tight_layout()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-06T08:57:59.491167Z","iopub.execute_input":"2025-03-06T08:57:59.491399Z","iopub.status.idle":"2025-03-06T08:58:01.157248Z","shell.execute_reply.started":"2025-03-06T08:57:59.491380Z","shell.execute_reply":"2025-03-06T08:58:01.156170Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 2. Dataset Class\n","metadata":{}},{"cell_type":"code","source":"import os\nimport glob\nimport random\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\nfrom scipy.ndimage import rotate, zoom\nimport torch\nfrom torch.utils.data import Dataset\n\nclass TomogramDataset(Dataset):\n    \"\"\"\n    Dataset for loading 3D tomograms from stacks of 2D JPG slices with augmentation.\n    \"\"\"\n    def __init__(self, csv_file, root_dir, train=True, max_slices=64, target_size=(128, 128), \n                 augment_prob=0.5, rotation_range=15, contrast_range=(0.8, 1.2), \n                 brightness_range=(-0.1, 0.1), noise_level=0.02, flip_prob=0.3,\n                 zoom_range=(0.9, 1.1)):\n        \"\"\"\n        Initialize the dataset.\n        \n        Args:\n            csv_file (str): Path to the CSV file with annotations.\n            root_dir (str): Directory with all the tomogram slice directories.\n            train (bool): Whether this is for training (enables augmentations).\n            max_slices (int): Maximum number of slices to use.\n            target_size (tuple): Target size for each 2D slice.\n            augment_prob (float): Probability of applying augmentation.\n            rotation_range (float): Maximum rotation angle in degrees.\n            contrast_range (tuple): Range for contrast adjustment.\n            brightness_range (tuple): Range for brightness adjustment.\n            noise_level (float): Maximum level of Gaussian noise to add.\n            flip_prob (float): Probability of flipping.\n            zoom_range (tuple): Range for random zoom.\n        \"\"\"\n        self.labels_df = pd.read_csv(csv_file)\n        self.root_dir = root_dir\n        self.train = train\n        self.max_slices = max_slices\n        self.target_size = target_size\n        \n        # Augmentation parameters\n        self.augment_prob = augment_prob\n        self.rotation_range = rotation_range\n        self.contrast_range = contrast_range\n        self.brightness_range = brightness_range\n        self.noise_level = noise_level\n        self.flip_prob = flip_prob\n        self.zoom_range = zoom_range\n        \n        # Process tomogram metadata\n        self.process_metadata()\n        \n        # Cache file paths and shapes\n        self.cache_file_paths()\n    \n    def process_metadata(self):\n        # Get unique tomograms\n        tomo_ids = self.labels_df['tomo_id'].unique()\n        self.tomo_df = pd.DataFrame({'tomo_id': tomo_ids})\n        \n        # For each tomogram, get its properties\n        for tomo_id in tomo_ids:\n            tomo_rows = self.labels_df[self.labels_df['tomo_id'] == tomo_id]\n            \n            # Get motor count\n            num_motors = tomo_rows['Number of motors'].iloc[0]\n            self.tomo_df.loc[self.tomo_df['tomo_id'] == tomo_id, 'Number of motors'] = num_motors\n            \n            # Get array shape and voxel spacing\n            self.tomo_df.loc[self.tomo_df['tomo_id'] == tomo_id, 'Array shape (axis 0)'] = tomo_rows['Array shape (axis 0)'].iloc[0]\n            self.tomo_df.loc[self.tomo_df['tomo_id'] == tomo_id, 'Array shape (axis 1)'] = tomo_rows['Array shape (axis 1)'].iloc[0]\n            self.tomo_df.loc[self.tomo_df['tomo_id'] == tomo_id, 'Array shape (axis 2)'] = tomo_rows['Array shape (axis 2)'].iloc[0]\n            self.tomo_df.loc[self.tomo_df['tomo_id'] == tomo_id, 'Voxel spacing'] = tomo_rows['Voxel spacing'].iloc[0]\n            \n            # Get motor axes (use first motor for training)\n            if num_motors > 0:\n                motor_row = tomo_rows[tomo_rows['Motor axis 0'] != -1].iloc[0]\n                self.tomo_df.loc[self.tomo_df['tomo_id'] == tomo_id, 'Motor axis 0'] = motor_row['Motor axis 0']\n                self.tomo_df.loc[self.tomo_df['tomo_id'] == tomo_id, 'Motor axis 1'] = motor_row['Motor axis 1']\n                self.tomo_df.loc[self.tomo_df['tomo_id'] == tomo_id, 'Motor axis 2'] = motor_row['Motor axis 2']\n            else:\n                # No motor\n                self.tomo_df.loc[self.tomo_df['tomo_id'] == tomo_id, 'Motor axis 0'] = -1\n                self.tomo_df.loc[self.tomo_df['tomo_id'] == tomo_id, 'Motor axis 1'] = -1\n                self.tomo_df.loc[self.tomo_df['tomo_id'] == tomo_id, 'Motor axis 2'] = -1\n    \n    def cache_file_paths(self):\n        self.slice_files = {}\n        \n        for idx, row in self.tomo_df.iterrows():\n            tomo_id = row['tomo_id']\n            tomo_dir = os.path.join(self.root_dir, tomo_id)\n            \n            # Get all slice files and sort them\n            files = sorted(glob.glob(os.path.join(tomo_dir, '*.jpg')))\n            self.slice_files[tomo_id] = files\n    \n    def __len__(self):\n        return len(self.tomo_df)\n    \n    def apply_augmentation(self, volume):\n        \"\"\"Apply various 3D augmentations to the volume.\"\"\"\n        # Only apply augmentation during training with probability\n        if not self.train or random.random() > self.augment_prob:\n            return volume\n            \n        # Random rotation (3D)\n        if random.random() < 0.7:  # 70% chance of rotation\n            angle_x = random.uniform(-self.rotation_range, self.rotation_range)\n            angle_y = random.uniform(-self.rotation_range, self.rotation_range)\n            angle_z = random.uniform(-self.rotation_range, self.rotation_range)\n            \n            # Apply rotation around each axis\n            if angle_x != 0:\n                volume = rotate(volume, angle_x, axes=(1, 2), reshape=False, order=1, mode='constant', cval=0)\n            if angle_y != 0:\n                volume = rotate(volume, angle_y, axes=(0, 2), reshape=False, order=1, mode='constant', cval=0)\n            if angle_z != 0:\n                volume = rotate(volume, angle_z, axes=(0, 1), reshape=False, order=1, mode='constant', cval=0)\n        \n        # Random zoom\n        if random.random() < 0.5:  # 50% chance of zoom\n            zoom_factor = random.uniform(self.zoom_range[0], self.zoom_range[1])\n            if zoom_factor != 1:\n                # Calculate padding or cropping needed\n                orig_shape = volume.shape\n                zoomed = zoom(volume, (1, zoom_factor, zoom_factor), order=1, mode='constant', cval=0)\n                \n                # If zoomed in (zoom_factor > 1), need to crop\n                if zoom_factor > 1:\n                    z, y, x = zoomed.shape\n                    start_y = (y - orig_shape[1]) // 2\n                    start_x = (x - orig_shape[2]) // 2\n                    volume = zoomed[:, start_y:start_y+orig_shape[1], start_x:start_x+orig_shape[2]]\n                # If zoomed out (zoom_factor < 1), need to pad\n                else:\n                    z, y, x = zoomed.shape\n                    pad_y = (orig_shape[1] - y) // 2\n                    pad_x = (orig_shape[2] - x) // 2\n                    volume = np.pad(zoomed, ((0, 0), (pad_y, orig_shape[1]-y-pad_y), (pad_x, orig_shape[2]-x-pad_x)), \n                                   mode='constant')\n        \n        # Random flips - using copy to ensure contiguous array\n        if random.random() < self.flip_prob:\n            volume = np.flip(volume, axis=1).copy()  # Flip horizontally\n            \n        if random.random() < self.flip_prob:\n            volume = np.flip(volume, axis=2).copy()  # Flip vertically\n            \n        # Random contrast\n        if random.random() < 0.6:  # 60% chance of contrast adjustment\n            contrast_factor = random.uniform(self.contrast_range[0], self.contrast_range[1])\n            mean = volume.mean()\n            volume = (volume - mean) * contrast_factor + mean\n            volume = np.clip(volume, 0, 1)\n        \n        # Random brightness\n        if random.random() < 0.6:  # 60% chance of brightness adjustment\n            brightness_factor = random.uniform(self.brightness_range[0], self.brightness_range[1])\n            volume = volume + brightness_factor\n            volume = np.clip(volume, 0, 1)\n        \n        # Add Gaussian noise\n        if random.random() < 0.4:  # 40% chance of adding noise\n            noise = np.random.normal(0, self.noise_level, volume.shape)\n            volume = volume + noise\n            volume = np.clip(volume, 0, 1)\n        \n        # Random dropout (simulating missing data)\n        if random.random() < 0.3:  # 30% chance of dropout\n            mask = np.random.rand(*volume.shape) > 0.05  # Dropout 5% of voxels\n            volume = volume * mask\n            \n        # Random intensity shifts for specific regions (simulating artifacts)\n        if random.random() < 0.2:  # 20% chance of intensity artifacts\n            num_regions = random.randint(1, 3)\n            for _ in range(num_regions):\n                z_size = random.randint(1, max(2, volume.shape[0] // 10))\n                y_size = random.randint(5, max(10, volume.shape[1] // 5))\n                x_size = random.randint(5, max(10, volume.shape[2] // 5))\n                \n                z_start = random.randint(0, volume.shape[0] - z_size)\n                y_start = random.randint(0, volume.shape[1] - y_size)\n                x_start = random.randint(0, volume.shape[2] - x_size)\n                \n                intensity_shift = random.uniform(-0.2, 0.2)\n                \n                region = volume[z_start:z_start+z_size, y_start:y_start+y_size, x_start:x_start+x_size]\n                region = region + intensity_shift\n                volume[z_start:z_start+z_size, y_start:y_start+y_size, x_start:x_start+x_size] = np.clip(region, 0, 1)\n        \n        # Ensure the volume is C-contiguous before returning\n        return np.ascontiguousarray(volume)\n    \n    def load_volume(self, tomo_id):\n        files = self.slice_files[tomo_id]\n        \n        # Determine which slices to load\n        if self.max_slices is not None and len(files) > self.max_slices:\n            # Subsample evenly\n            indices = np.linspace(0, len(files)-1, self.max_slices, dtype=int)\n            files_to_load = [files[i] for i in indices]\n        else:\n            files_to_load = files\n        \n        # Load slices\n        slices = []\n        for file_path in files_to_load:\n            img = Image.open(file_path).convert('L')  # Convert to grayscale\n            img = img.resize(self.target_size, Image.BILINEAR)\n            slices.append(np.array(img))\n        \n        # Stack slices to form volume\n        volume = np.stack(slices)\n        \n        # Pad if needed\n        if self.max_slices is not None and volume.shape[0] < self.max_slices:\n            pad_width = self.max_slices - volume.shape[0]\n            pad_before = pad_width // 2\n            pad_after = pad_width - pad_before\n            volume = np.pad(volume, ((pad_before, pad_after), (0, 0), (0, 0)), mode='constant')\n        \n        # Normalize to [0, 1]\n        volume = volume.astype(np.float32) / 255.0\n        \n        return volume\n    \n    def __getitem__(self, idx):\n        row = self.tomo_df.iloc[idx]\n        tomo_id = row['tomo_id']\n        \n        # Load volume\n        volume = self.load_volume(tomo_id)\n        \n        # Get labels\n        motor_axes = row[['Motor axis 0', 'Motor axis 1', 'Motor axis 2']].values.astype(np.float32)\n        has_motor = not (motor_axes == -1).all()\n        \n        # Get original shape for coordinate normalization\n        array_shape = np.array([\n            row['Array shape (axis 0)'],\n            row['Array shape (axis 1)'],\n            row['Array shape (axis 2)']\n        ], dtype=np.float32)\n        \n        # Process coordinates\n        if not has_motor:\n            motor_axes = np.zeros(3, dtype=np.float32)\n        else:\n            # Apply coordinate jitter in training\n            if self.train and random.random() < 0.5:\n                # Add small random jitter to coordinates (within 5% of dimension)\n                jitter_z = np.random.uniform(-0.05, 0.05) * array_shape[0]\n                jitter_x = np.random.uniform(-0.05, 0.05) * array_shape[1]\n                jitter_y = np.random.uniform(-0.05, 0.05) * array_shape[2]\n                \n                motor_axes[0] += jitter_z\n                motor_axes[1] += jitter_x\n                motor_axes[2] += jitter_y\n                \n                # Ensure coordinates are still within bounds\n                motor_axes[0] = max(0, min(motor_axes[0], array_shape[0] - 1))\n                motor_axes[1] = max(0, min(motor_axes[1], array_shape[1] - 1))\n                motor_axes[2] = max(0, min(motor_axes[2], array_shape[2] - 1))\n            \n            # Normalize coordinates to [0, 1]\n            motor_axes[0] = motor_axes[0] / array_shape[0]\n            motor_axes[1] = motor_axes[1] / array_shape[1]\n            motor_axes[2] = motor_axes[2] / array_shape[2]\n        \n        # Apply augmentation to volume\n        volume = self.apply_augmentation(volume)\n        \n        # Ensure the volume is C-contiguous\n        if not volume.flags.c_contiguous:\n            volume = np.ascontiguousarray(volume)\n        \n        # Convert to tensor\n        volume = torch.from_numpy(volume).unsqueeze(0).float()  # Add channel dimension\n        motor_axes = torch.from_numpy(motor_axes)\n        has_motor = torch.tensor([float(has_motor)])\n        \n        return {\n            'tomo_id': tomo_id,\n            'volume': volume,\n            'has_motor': has_motor,\n            'motor_axes': motor_axes,\n            'original_shape': torch.tensor([\n                row['Array shape (axis 0)'], \n                row['Array shape (axis 1)'], \n                row['Array shape (axis 2)']\n            ], dtype=torch.float32),\n            'voxel_spacing': torch.tensor([row['Voxel spacing']], dtype=torch.float32)\n        }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-06T08:58:01.158861Z","iopub.execute_input":"2025-03-06T08:58:01.159184Z","iopub.status.idle":"2025-03-06T08:58:01.195810Z","shell.execute_reply.started":"2025-03-06T08:58:01.159157Z","shell.execute_reply":"2025-03-06T08:58:01.195015Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n\nclass TestTomogramDataset(Dataset):\n    \"\"\"Dataset for loading test tomograms with no labels.\"\"\"\n    def __init__(self, root_dir, max_slices=64, target_size=(128, 128)):\n        self.root_dir = root_dir\n        self.max_slices = max_slices\n        self.target_size = target_size\n        \n        # Get all tomogram directories\n        self.tomo_dirs = sorted([d for d in os.listdir(root_dir) if os.path.isdir(os.path.join(root_dir, d))])\n        \n        # Cache file paths\n        self.slice_files = {}\n        for tomo_id in self.tomo_dirs:\n            tomo_dir = os.path.join(self.root_dir, tomo_id)\n            files = sorted(glob.glob(os.path.join(tomo_dir, '*.jpg')))\n            self.slice_files[tomo_id] = files\n    \n    def __len__(self):\n        return len(self.tomo_dirs)\n    \n    def load_volume(self, tomo_id):\n        files = self.slice_files[tomo_id]\n        \n        # Get array shape\n        z_shape = len(files)\n        if z_shape > 0:\n            img = Image.open(files[0])\n            x_shape, y_shape = img.size\n        else:\n            raise ValueError(f\"No slices found for tomogram {tomo_id}\")\n        \n        # Store original shape\n        original_shape = np.array([z_shape, x_shape, y_shape])\n        \n        # Determine which slices to load\n        if self.max_slices is not None and z_shape > self.max_slices:\n            # Subsample evenly\n            indices = np.linspace(0, z_shape-1, self.max_slices, dtype=int)\n            files_to_load = [files[i] for i in indices]\n        else:\n            files_to_load = files\n        \n        # Load slices\n        slices = []\n        for file_path in files_to_load:\n            img = Image.open(file_path).convert('L')  # Convert to grayscale\n            img = img.resize(self.target_size, Image.BILINEAR)\n            slices.append(np.array(img))\n        \n        # Stack slices to form volume\n        volume = np.stack(slices)\n        \n        # Pad if needed\n        if self.max_slices is not None and volume.shape[0] < self.max_slices:\n            pad_width = self.max_slices - volume.shape[0]\n            pad_before = pad_width // 2\n            pad_after = pad_width - pad_before\n            volume = np.pad(volume, ((pad_before, pad_after), (0, 0), (0, 0)), mode='constant')\n        \n        # Normalize to [0, 1]\n        volume = volume.astype(np.float32) / 255.0\n        \n        return volume, original_shape\n    \n    def __getitem__(self, idx):\n        tomo_id = self.tomo_dirs[idx]\n        \n        # Load volume\n        volume, original_shape = self.load_volume(tomo_id)\n        \n        # Convert to tensor\n        volume = torch.from_numpy(volume).unsqueeze(0)  # Add channel dimension\n        \n        return {\n            'tomo_id': tomo_id,\n            'volume': volume,\n            'original_shape': torch.tensor(original_shape, dtype=torch.float32)\n        }\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-06T08:58:01.196677Z","iopub.execute_input":"2025-03-06T08:58:01.197004Z","iopub.status.idle":"2025-03-06T08:58:01.223395Z","shell.execute_reply.started":"2025-03-06T08:58:01.196963Z","shell.execute_reply":"2025-03-06T08:58:01.222745Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3. Model Architecture\n","metadata":{}},{"cell_type":"code","source":"class SEBlock3D(nn.Module):\n    def __init__(self, channel, reduction=16):\n        super().__init__()\n        self.avg_pool = nn.AdaptiveAvgPool3d(1)\n        self.fc = nn.Sequential(\n            nn.Linear(channel, channel // reduction, bias=False),\n            nn.ReLU(inplace=True),\n            nn.Linear(channel // reduction, channel, bias=False),\n            nn.Sigmoid()\n        )\n        self.spatial_conv = nn.Conv3d(2, 1, kernel_size=7, padding=3, bias=False)\n        self.spatial_sigmoid = nn.Sigmoid()\n\n    def forward(self, x):\n        b, c, d, h, w = x.size()\n        y = self.avg_pool(x).view(b, c)\n        y = self.fc(y).view(b, c, 1, 1, 1)\n        x = x * y\n        y = torch.cat([torch.max(x, 1, keepdim=True)[0], torch.mean(x, 1, keepdim=True)], dim=1)\n        y = self.spatial_conv(y)\n        y = self.spatial_sigmoid(y)\n        return x * y\n\nclass EnhancedAttentionBlock(nn.Module):\n    def __init__(self, embed_dim, num_heads, dropout=0.1):\n        super().__init__()\n        self.mha = nn.MultiheadAttention(embed_dim=embed_dim, num_heads=num_heads, dropout=dropout, batch_first=True)\n        self.ln1 = nn.LayerNorm(embed_dim)  # Layer norm before attention\n        self.ffn = nn.Sequential(\n            nn.Linear(embed_dim, embed_dim * 4),  # Expansion\n            nn.ReLU(),\n            nn.Dropout(dropout),\n            nn.Linear(embed_dim * 4, embed_dim),  # Reduction\n            nn.Dropout(dropout)\n        )\n        self.ln2 = nn.LayerNorm(embed_dim)  # Layer norm before FFN\n        self.dropout = nn.Dropout(dropout)\n\n    def forward(self, x):\n        # Input x: (batch, seq_len, embed_dim)\n        identity = x\n        x = self.ln1(x)\n        attn_output, _ = self.mha(x, x, x)\n        x = identity + self.dropout(attn_output)  # Residual connection\n        identity = x\n        x = self.ln2(x)\n        ffn_output = self.ffn(x)\n        x = identity + self.dropout(ffn_output)  # Residual connection\n        return x\n\nclass FlagellarMotorNet(nn.Module):\n    def __init__(self, input_channels=1, base_filters=16, num_heads=4, dropout=0.1):\n        super().__init__()\n        self.conv1 = nn.Conv3d(input_channels, base_filters, kernel_size=3, stride=1, padding=1)\n        self.bn1 = nn.BatchNorm3d(base_filters)\n        self.se1 = SEBlock3D(base_filters)\n        self.pool1 = nn.MaxPool3d(kernel_size=2)\n        \n        self.conv2 = nn.Conv3d(base_filters, base_filters*2, kernel_size=3, stride=1, padding=1)\n        self.bn2 = nn.BatchNorm3d(base_filters*2)\n        self.se2 = SEBlock3D(base_filters*2)\n        self.pool2 = nn.MaxPool3d(kernel_size=2)\n        self.res_conv2 = nn.Conv3d(base_filters, base_filters*2, kernel_size=1, stride=1)\n        \n        self.conv3 = nn.Conv3d(base_filters*2, base_filters*4, kernel_size=3, stride=1, padding=1)\n        self.bn3 = nn.BatchNorm3d(base_filters*4)\n        self.se3 = SEBlock3D(base_filters*4)\n        self.pool3 = nn.MaxPool3d(kernel_size=2)\n        self.res_conv3 = nn.Conv3d(base_filters*2, base_filters*4, kernel_size=1, stride=1)\n        \n        self.conv4 = nn.Conv3d(base_filters*4, base_filters*8, kernel_size=3, stride=1, padding=1)\n        self.bn4 = nn.BatchNorm3d(base_filters*8)\n        self.se4 = SEBlock3D(base_filters*8)\n        self.pool4 = nn.MaxPool3d(kernel_size=2)\n        self.res_conv4 = nn.Conv3d(base_filters*4, base_filters*8, kernel_size=1, stride=1)\n        \n        self.attn_channels = base_filters * 8\n        self.num_heads = num_heads\n        assert self.attn_channels % num_heads == 0, \"attn_channels must be divisible by num_heads\"\n        self.attn_block = EnhancedAttentionBlock(embed_dim=self.attn_channels, num_heads=num_heads, dropout=dropout)\n        \n        self.spatial_size = 4 * 8 * 8  # Assuming input size reduces to 4x8x8 after pooling\n        self.fc_size = self.attn_channels * self.spatial_size\n        \n        self.fc_presence = nn.Sequential(\n            nn.Linear(self.fc_size, 64),\n            nn.ReLU(),\n            nn.Dropout(0.5),\n            nn.Linear(64, 1),\n            nn.Sigmoid()\n        )\n        self.fc_location = nn.Sequential(\n            nn.Linear(self.fc_size, 128),\n            nn.ReLU(),\n            nn.Dropout(0.5),\n            nn.Linear(128, 3),\n            nn.Sigmoid()\n        )\n    \n    def forward(self, x):\n        x = F.relu(self.bn1(self.conv1(x)))\n        x = self.se1(x)\n        x = self.pool1(x)\n        \n        identity = x\n        x = F.relu(self.bn2(self.conv2(x)))\n        x = self.se2(x)\n        x = self.pool2(x)\n        identity = self.pool2(self.res_conv2(identity))\n        x = x + identity\n        \n        identity = x\n        x = F.relu(self.bn3(self.conv3(x)))\n        x = self.se3(x)\n        x = self.pool3(x)\n        identity = self.pool3(self.res_conv3(identity))\n        x = x + identity\n        \n        identity = x\n        x = F.relu(self.bn4(self.conv4(x)))\n        x = self.se4(x)\n        x = self.pool4(x)\n        identity = self.pool4(self.res_conv4(identity))\n        x = x + identity\n        \n        b, c, d, h, w = x.size()\n        x = x.view(b, c, -1).permute(0, 2, 1)  # (batch, seq_len, embed_dim)\n        x = self.attn_block(x)\n        x = x.permute(0, 2, 1).view(b, c, d, h, w)\n        \n        x = x.reshape(b, -1)\n        presence = self.fc_presence(x)\n        location = self.fc_location(x)\n        return presence, location\n        \n\nclass ResBlock3D(nn.Module):\n    def __init__(self, in_channels, out_channels, stride=1, bottleneck_factor=4):\n        super().__init__()\n        bottleneck_channels = out_channels // bottleneck_factor\n        self.conv1 = nn.Conv3d(in_channels, bottleneck_channels, kernel_size=1, bias=False)\n        self.bn1 = nn.BatchNorm3d(bottleneck_channels)\n        self.conv2 = nn.Conv3d(bottleneck_channels, bottleneck_channels, kernel_size=3, stride=stride, padding=1, bias=False)\n        self.bn2 = nn.BatchNorm3d(bottleneck_channels)\n        self.conv3 = nn.Conv3d(bottleneck_channels, out_channels, kernel_size=1, bias=False)\n        self.bn3 = nn.BatchNorm3d(out_channels)\n        self.ln = nn.LayerNorm(out_channels)  # Normalize across channels\n        \n        self.shortcut = nn.Sequential()\n        if stride != 1 or in_channels != out_channels:\n            self.shortcut = nn.Sequential(\n                nn.Conv3d(in_channels, out_channels, kernel_size=1, stride=stride, bias=False),\n                nn.BatchNorm3d(out_channels)\n            )\n    \n    def forward(self, x):\n        identity = x\n        out = F.relu(self.bn1(self.conv1(x)))\n        out = F.relu(self.bn2(self.conv2(out)))\n        out = self.bn3(self.conv3(out))\n        shortcut = self.shortcut(identity)\n        \n        # Reshape for LayerNorm: move channels to the last dimension\n        out = out + shortcut  # Shape: [batch, channels, depth, height, width]\n        b, c, d, h, w = out.size()\n        out = out.permute(0, 2, 3, 4, 1).contiguous()  # Shape: [batch, depth, height, width, channels]\n        out = self.ln(out)  # Apply LayerNorm across channels\n        out = out.permute(0, 4, 1, 2, 3).contiguous()  # Shape: [batch, channels, depth, height, width]\n        \n        out = F.relu(out)\n        return out\n\n\nclass ResNet3D(nn.Module):\n    def __init__(self, input_channels=1, base_filters=16):\n        super().__init__()\n        self.conv1 = nn.Conv3d(input_channels, base_filters, kernel_size=7, stride=2, padding=3, bias=False)\n        self.bn1 = nn.BatchNorm3d(base_filters)\n        self.pool = nn.MaxPool3d(kernel_size=3, stride=2, padding=1)\n        \n        self.layer1 = self._make_layer(base_filters, base_filters, 2, stride=1)\n        self.layer2 = self._make_layer(base_filters, base_filters*2, 2, stride=2)\n        self.layer3 = self._make_layer(base_filters*2, base_filters*4, 2, stride=2)\n        self.layer4 = self._make_layer(base_filters*4, base_filters*8, 2, stride=2)\n        \n        self.avg_pool = nn.AdaptiveAvgPool3d((1, 1, 1))\n        self.fc_presence = nn.Sequential(\n            nn.Linear(base_filters*8, 64),\n            nn.ReLU(),\n            nn.Dropout(0.5),\n            nn.Linear(64, 1),\n            nn.Sigmoid()\n        )\n        self.fc_location = nn.Sequential(\n            nn.Linear(base_filters*8, 128),\n            nn.ReLU(),\n            nn.Dropout(0.5),\n            nn.Linear(128, 3),\n            nn.Sigmoid()\n        )\n    \n    def _make_layer(self, in_channels, out_channels, blocks, stride):\n        layers = [ResBlock3D(in_channels, out_channels, stride)]\n        for _ in range(1, blocks):\n            layers.append(ResBlock3D(out_channels, out_channels, stride=1))\n        return nn.Sequential(*layers)\n    \n    def forward(self, x):\n        x = F.relu(self.bn1(self.conv1(x)))\n        x = self.pool(x)\n        x = self.layer1(x)\n        x = self.layer2(x)\n        x = self.layer3(x)\n        x = self.layer4(x)\n        x = self.avg_pool(x)\n        x = x.view(x.size(0), -1)\n        presence = self.fc_presence(x)\n        location = self.fc_location(x)\n        return presence, location\n# EnsembleNet\nclass EnsembleNet(nn.Module):\n    def __init__(self, input_channels=1, base_filters=16, num_heads=4):\n        super().__init__()\n        self.model1 = FlagellarMotorNet(input_channels=input_channels, base_filters=base_filters, num_heads=num_heads)\n        self.model2 = ResNet3D(input_channels=input_channels, base_filters=base_filters)\n    \n    def forward(self, x):\n        p1, l1 = self.model1(x)\n        p2, l2 = self.model2(x)\n        presence = (p1 + p2) / 2\n        location = (l1 + l2) / 2\n        return presence, location","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-06T08:58:01.224189Z","iopub.execute_input":"2025-03-06T08:58:01.224398Z","iopub.status.idle":"2025-03-06T08:58:01.250662Z","shell.execute_reply.started":"2025-03-06T08:58:01.224373Z","shell.execute_reply":"2025-03-06T08:58:01.250028Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 4. Training Function\n","metadata":{}},{"cell_type":"code","source":"\n# Custom loss function\nclass FlagellarMotorLoss(nn.Module):\n    \"\"\"\n    Custom loss function for flagellar motor detection and localization.\n    Combines binary cross-entropy for motor presence with MSE for location.\n    \"\"\"\n    def __init__(self, presence_weight=1.0, location_weight=3.0):\n        super(FlagellarMotorLoss, self).__init__()\n        self.presence_weight = presence_weight\n        self.location_weight = location_weight\n        self.bce_loss = nn.BCELoss()\n        self.mse_loss = nn.MSELoss()\n    \n    def forward(self, presence_pred, location_pred, presence_true, location_true):\n        # Presence loss (binary cross-entropy)\n        presence_loss = self.bce_loss(presence_pred, presence_true)\n        \n        # Location loss (only computed for tomograms with motors)\n        if torch.sum(presence_true) > 0:\n            # Select only the samples with motors\n            has_motor = presence_true.squeeze() > 0.5\n            if has_motor.sum() > 0:\n                location_pred_with_motor = location_pred[has_motor]\n                location_true_with_motor = location_true[has_motor]\n                \n                location_loss = self.mse_loss(location_pred_with_motor, location_true_with_motor)\n                \n                # Calculate Euclidean distance (for monitoring)\n                euclidean_dist = torch.sqrt(torch.sum((location_pred_with_motor - location_true_with_motor) ** 2, dim=1))\n                avg_euclidean_dist = euclidean_dist.mean()\n            else:\n                location_loss = torch.tensor(0.0, device=presence_loss.device)\n                avg_euclidean_dist = torch.tensor(0.0, device=presence_loss.device)\n        else:\n            location_loss = torch.tensor(0.0, device=presence_loss.device)\n            avg_euclidean_dist = torch.tensor(0.0, device=presence_loss.device)\n        \n        # Total loss\n        total_loss = self.presence_weight * presence_loss + self.location_weight * location_loss\n        \n        return total_loss, presence_loss, location_loss, avg_euclidean_dist\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-06T08:58:01.251278Z","iopub.execute_input":"2025-03-06T08:58:01.251506Z","iopub.status.idle":"2025-03-06T08:58:01.266071Z","shell.execute_reply.started":"2025-03-06T08:58:01.251485Z","shell.execute_reply":"2025-03-06T08:58:01.265238Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n\ndef train_epoch(model, dataloader, optimizer, criterion, device):\n    \"\"\"Train for one epoch\"\"\"\n    model.train()\n    epoch_loss = 0\n    epoch_presence_loss = 0\n    epoch_location_loss = 0\n    epoch_euclidean_dist = 0\n    \n    progress_bar = tqdm(dataloader, desc=\"Training\")\n    \n    for batch in progress_bar:\n        # Move data to device\n        volume = batch['volume'].to(device)\n        has_motor = batch['has_motor'].to(device)\n        motor_axes = batch['motor_axes'].to(device)\n        \n        # Forward pass\n        presence_pred, location_pred = model(volume)\n        \n        # Calculate loss\n        loss, presence_loss, location_loss, euclidean_dist = criterion(\n            presence_pred, location_pred, has_motor, motor_axes\n        )\n        \n        # Backward pass and optimize\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n        \n        # Update metrics\n        epoch_loss += loss.item()\n        epoch_presence_loss += presence_loss.item()\n        epoch_location_loss += location_loss.item()\n        epoch_euclidean_dist += euclidean_dist.item()\n        \n        # Update progress bar\n        progress_bar.set_postfix({\n            'loss': loss.item(),\n            'p_loss': presence_loss.item(),\n            'l_loss': location_loss.item(),\n            'eucl_dist': euclidean_dist.item()\n        })\n    \n    # Calculate average metrics\n    num_batches = len(dataloader)\n    avg_loss = epoch_loss / num_batches\n    avg_presence_loss = epoch_presence_loss / num_batches\n    avg_location_loss = epoch_location_loss / num_batches\n    avg_euclidean_dist = epoch_euclidean_dist / num_batches\n    \n    return {\n        'loss': avg_loss,\n        'presence_loss': avg_presence_loss,\n        'location_loss': avg_location_loss,\n        'euclidean_dist': avg_euclidean_dist\n    }\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-06T08:58:01.268351Z","iopub.execute_input":"2025-03-06T08:58:01.268573Z","iopub.status.idle":"2025-03-06T08:58:01.283685Z","shell.execute_reply.started":"2025-03-06T08:58:01.268555Z","shell.execute_reply":"2025-03-06T08:58:01.283065Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Validation Function\n\ndef validate(model, dataloader, criterion, device, threshold=0.5):\n    \"\"\"Validate the model\"\"\"\n    model.eval()\n    epoch_loss = 0\n    epoch_presence_loss = 0\n    epoch_location_loss = 0\n    epoch_euclidean_dist = 0\n    \n    # Track predictions for F-beta score\n    true_positives = 0\n    false_positives = 0\n    false_negatives = 0\n    \n    progress_bar = tqdm(dataloader, desc=\"Validation\")\n    \n    with torch.no_grad():\n        for batch in progress_bar:\n            # Move data to device\n            volume = batch['volume'].to(device)\n            has_motor = batch['has_motor'].to(device)\n            motor_axes = batch['motor_axes'].to(device)\n            original_shape = batch['original_shape'].to(device)\n            voxel_spacing = batch['voxel_spacing'].to(device)\n            \n            # Forward pass\n            presence_pred, location_pred = model(volume)\n            \n            # Calculate loss\n            loss, presence_loss, location_loss, euclidean_dist = criterion(\n                presence_pred, location_pred, has_motor, motor_axes\n            )\n            \n            # Update metrics\n            epoch_loss += loss.item()\n            epoch_presence_loss += presence_loss.item()\n            epoch_location_loss += location_loss.item()\n            epoch_euclidean_dist += euclidean_dist.item()\n            \n            # Calculate F-beta metrics\n            for i in range(len(presence_pred)):\n                # Check if model predicts a motor\n                pred_has_motor = presence_pred[i].item() > threshold\n                true_has_motor = has_motor[i].item() > 0.5\n                \n                if pred_has_motor and true_has_motor:\n                    # Convert normalized coordinates back to original space\n                    pred_coords = location_pred[i].cpu().numpy()\n                    true_coords = motor_axes[i].cpu().numpy()\n                    shape = original_shape[i].cpu().numpy()\n                    spacing = voxel_spacing[i].item()\n                    \n                    # Denormalize coordinates\n                    pred_coords_orig = np.array([\n                        pred_coords[0] * shape[0],\n                        pred_coords[1] * shape[1],\n                        pred_coords[2] * shape[2]\n                    ])\n                    \n                    true_coords_orig = np.array([\n                        true_coords[0] * shape[0],\n                        true_coords[1] * shape[1],\n                        true_coords[2] * shape[2]\n                    ])\n                    \n                    # Calculate Euclidean distance in Angstroms\n                    dist = np.sqrt(np.sum((pred_coords_orig - true_coords_orig) ** 2)) * spacing\n                    \n                    # Check if prediction is within threshold (1000 Angstroms)\n                    if dist <= 1000:\n                        true_positives += 1\n                    else:\n                        false_positives += 1\n                        false_negatives += 1\n                elif pred_has_motor and not true_has_motor:\n                    false_positives += 1\n                elif not pred_has_motor and true_has_motor:\n                    false_negatives += 1\n            \n            # Update progress bar\n            progress_bar.set_postfix({\n                'loss': loss.item(),\n                'p_loss': presence_loss.item(),\n                'l_loss': location_loss.item(),\n                'eucl_dist': euclidean_dist.item()\n            })\n    \n    # Calculate average metrics\n    num_batches = len(dataloader)\n    avg_loss = epoch_loss / num_batches\n    avg_presence_loss = epoch_presence_loss / num_batches\n    avg_location_loss = epoch_location_loss / num_batches\n    avg_euclidean_dist = epoch_euclidean_dist / num_batches\n    \n    # Calculate F-beta score (beta=2)\n    beta = 2\n    if true_positives + false_positives > 0:\n        precision = true_positives / (true_positives + false_positives)\n    else:\n        precision = 0\n    \n    if true_positives + false_negatives > 0:\n        recall = true_positives / (true_positives + false_negatives)\n    else:\n        recall = 0\n    \n    if precision + recall > 0:\n        f_beta = (1 + beta**2) * precision * recall / ((beta**2 * precision) + recall)\n    else:\n        f_beta = 0\n    \n    return {\n        'loss': avg_loss,\n        'presence_loss': avg_presence_loss,\n        'location_loss': avg_location_loss,\n        'euclidean_dist': avg_euclidean_dist,\n        'f_beta': f_beta,\n        'precision': precision,\n        'recall': recall,\n        'true_positives': true_positives,\n        'false_positives': false_positives,\n        'false_negatives': false_negatives\n    }\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-06T08:58:01.284724Z","iopub.execute_input":"2025-03-06T08:58:01.284912Z","iopub.status.idle":"2025-03-06T08:58:01.300239Z","shell.execute_reply.started":"2025-03-06T08:58:01.284895Z","shell.execute_reply":"2025-03-06T08:58:01.299582Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## Prediction function\ndef predict(model, dataloader, device, threshold=0.5):\n    \"\"\"Generate predictions for test set\"\"\"\n    model.eval()\n    predictions = []\n    \n    with torch.no_grad():\n        for batch in tqdm(dataloader, desc=\"Predicting\"):\n            # Move data to device\n            volume = batch['volume'].to(device)\n            tomo_ids = batch['tomo_id']\n            original_shape = batch['original_shape'].to(device)\n            \n            # Forward pass\n            presence_pred, location_pred = model(volume)\n            \n            # Process predictions\n            for i in range(len(presence_pred)):\n                tomo_id = tomo_ids[i]\n                pred_has_motor = presence_pred[i].item() > threshold\n                \n                if pred_has_motor:\n                    # Convert normalized coordinates back to original space\n                    pred_coords = location_pred[i].cpu().numpy()\n                    shape = original_shape[i].cpu().numpy()\n                    \n                    # Denormalize coordinates\n                    pred_coords_orig = np.array([\n                        pred_coords[0] * shape[0],\n                        pred_coords[1] * shape[1],\n                        pred_coords[2] * shape[2]\n                    ])\n                    \n                    predictions.append({\n                        'tomo_id': tomo_id,\n                        'Motor axis 0': pred_coords_orig[0],\n                        'Motor axis 1': pred_coords_orig[1],\n                        'Motor axis 2': pred_coords_orig[2]\n                    })\n                else:\n                    predictions.append({\n                        'tomo_id': tomo_id,\n                        'Motor axis 0': -1,\n                        'Motor axis 1': -1,\n                        'Motor axis 2': -1\n                    })\n    \n    return pd.DataFrame(predictions)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-06T08:58:01.301075Z","iopub.execute_input":"2025-03-06T08:58:01.301346Z","iopub.status.idle":"2025-03-06T08:58:01.319105Z","shell.execute_reply.started":"2025-03-06T08:58:01.301318Z","shell.execute_reply":"2025-03-06T08:58:01.318476Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 5. Training Loop\n","metadata":{}},{"cell_type":"code","source":"def train_model():\n    \"\"\"Train the model and save checkpoints\"\"\"\n    # Configuration\n    config = {\n        'batch_size': 16,\n        'num_workers': 2,\n        'max_slices': 64,\n        'target_size': (128, 128),\n        'learning_rate': 0.0005,\n        'weight_decay': 0.0001,\n        'epochs': 20, \n        'presence_weight': 1.0,\n        'location_weight': 3.0,\n        'threshold': 0.5,\n        'validation_split': 0.2\n    }\n    \n    # Load and preprocess data\n    train_df = pd.read_csv(TRAIN_CSV)\n    \n    # Get unique tomograms\n    tomo_ids = train_df['tomo_id'].unique()\n    \n    # Split tomograms into train and validation sets\n    train_tomo_ids, val_tomo_ids = train_test_split(\n        tomo_ids, \n        test_size=config['validation_split'], \n        random_state=RANDOM_SEED,\n        stratify=train_df.drop_duplicates('tomo_id')['Number of motors'] > 0  # Stratify by motor presence\n    )\n    \n    # Filter train_df to get only the relevant tomograms\n    train_set_df = train_df[train_df['tomo_id'].isin(train_tomo_ids)]\n    val_set_df = train_df[train_df['tomo_id'].isin(val_tomo_ids)]\n    \n    # Create temporary CSVs for the datasets\n    train_csv = os.path.join(OUTPUT_DIR, 'train_set.csv')\n    val_csv = os.path.join(OUTPUT_DIR, 'val_set.csv')\n    \n    train_set_df.to_csv(train_csv, index=False)\n    val_set_df.to_csv(val_csv, index=False)\n    \n    # Create datasets\n    train_dataset = TomogramDataset(\n        csv_file=train_csv,\n        root_dir=TRAIN_DIR,\n        train=True,\n        max_slices=config['max_slices'],\n        target_size=config['target_size']\n    )\n    \n    val_dataset = TomogramDataset(\n        csv_file=val_csv,\n        root_dir=TRAIN_DIR,\n        train=False,\n        max_slices=config['max_slices'],\n        target_size=config['target_size']\n    )\n    \n    # Create data loaders\n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=config['batch_size'],\n        shuffle=True,\n        num_workers=config['num_workers'],\n        pin_memory=True\n    )\n    \n    val_loader = DataLoader(\n        val_dataset,\n        batch_size=config['batch_size'],\n        shuffle=False,\n        num_workers=config['num_workers'],\n        pin_memory=True\n    )\n    \n    # Print dataset sizes\n    print(f\"Training dataset size: {len(train_dataset)}\")\n    print(f\"Validation dataset size: {len(val_dataset)}\")\n    \n    # Initialize model\n    model = EnsembleNet(\n        input_channels=1,\n        base_filters=16\n    ).to(DEVICE)\n    \n    # Initialize optimizer\n    optimizer = optim.Adam(\n        model.parameters(),\n        lr=config['learning_rate'],\n        weight_decay=config['weight_decay']\n    )\n    \n    # Initialize scheduler\n    scheduler = ReduceLROnPlateau(\n        optimizer,\n        mode='min',\n        factor=0.5,\n        patience=5,\n        verbose=True\n    )\n    \n    # Initialize loss function\n    criterion = FlagellarMotorLoss(\n        presence_weight=config['presence_weight'],\n        location_weight=config['location_weight']\n    )\n    \n    # Initialize best metrics\n    best_val_loss = float('inf')\n    best_f_beta = 0\n    \n    # Training loop\n    for epoch in range(config['epochs']):\n        print(f\"\\nEpoch {epoch+1}/{config['epochs']}\")\n        \n        # Train\n        train_metrics = train_epoch(model, train_loader, optimizer, criterion, DEVICE)\n        \n        # Validate\n        val_metrics = validate(model, val_loader, criterion, DEVICE, threshold=config['threshold'])\n        \n        # Update scheduler\n        scheduler.step(val_metrics['loss'])\n        \n        # Print metrics\n        print(f\"Train Loss: {train_metrics['loss']:.4f}, Val Loss: {val_metrics['loss']:.4f}\")\n        print(f\"Val F-beta (β=2): {val_metrics['f_beta']:.4f}, Precision: {val_metrics['precision']:.4f}, Recall: {val_metrics['recall']:.4f}\")\n        print(f\"Val Euclidean Dist: {val_metrics['euclidean_dist']:.4f}\")\n        \n        # Save best model (by loss)\n        if val_metrics['loss'] < best_val_loss:\n            best_val_loss = val_metrics['loss']\n            \n            # Save model\n            torch.save({\n                'epoch': epoch,\n                'model_state_dict': model.state_dict(),\n                'optimizer_state_dict': optimizer.state_dict(),\n                'scheduler_state_dict': scheduler.state_dict(),\n                'val_metrics': val_metrics,\n                'config': config\n            }, os.path.join(MODEL_DIR, 'best_model_loss.pth'))\n            \n            print(f\"Saved best model by loss: {best_val_loss:.4f}\")\n        \n        # Save best model (by F-beta)\n        if val_metrics['f_beta'] > best_f_beta:\n            best_f_beta = val_metrics['f_beta']\n            \n            # Save model\n            torch.save({\n                'epoch': epoch,\n                'model_state_dict': model.state_dict(),\n                'optimizer_state_dict': optimizer.state_dict(),\n                'scheduler_state_dict': scheduler.state_dict(),\n                'val_metrics': val_metrics,\n                'config': config\n            }, os.path.join(MODEL_DIR, 'best_model_fbeta.pth'))\n            \n            print(f\"Saved best model by F-beta: {best_f_beta:.4f}\")\n    \n    # Clean up temporary files\n    os.remove(train_csv)\n    os.remove(val_csv)\n    \n    print(\"\\nTraining completed!\")\n    return model\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-06T08:58:01.319874Z","iopub.execute_input":"2025-03-06T08:58:01.320147Z","iopub.status.idle":"2025-03-06T08:58:01.339220Z","shell.execute_reply.started":"2025-03-06T08:58:01.320116Z","shell.execute_reply":"2025-03-06T08:58:01.338477Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 6. Generate Predictions\n","metadata":{}},{"cell_type":"code","source":"def generate_predictions(model_path=None):\n    \"\"\"Generate predictions for test set\"\"\"\n    # Configuration\n    config = {\n        'batch_size': 4,\n        'num_workers': 2,\n        'max_slices': 64,\n        'target_size': (128, 128),\n        'threshold': 0.5\n    }\n    \n    # Use specified model path or default\n    if model_path is None:\n        model_path = os.path.join(MODEL_DIR, 'best_model_fbeta.pth')\n    \n    # Create test dataset\n    test_dataset = TestTomogramDataset(\n        root_dir=TEST_DIR,\n        max_slices=config['max_slices'],\n        target_size=config['target_size']\n    )\n    \n    # Create test dataloader\n    test_loader = DataLoader(\n        test_dataset,\n        batch_size=config['batch_size'],\n        shuffle=False,\n        num_workers=config['num_workers'],\n        pin_memory=True\n    )\n    \n    print(f\"Test dataset size: {len(test_dataset)}\")\n    \n    # Load model or create a new one if not found\n    if os.path.exists(model_path):\n        checkpoint = torch.load(model_path, map_location=DEVICE)\n        \n        # Initialize model\n        model = EnsembleNet(\n            input_channels=1,\n            base_filters=16\n        ).to(DEVICE)\n        \n        # Load weights\n        model.load_state_dict(checkpoint['model_state_dict'])\n        \n        print(f\"Loaded model from {model_path}\")\n        print(f\"Model was trained for {checkpoint['epoch']+1} epochs\")\n        print(f\"Validation metrics at checkpoint: F-beta = {checkpoint['val_metrics']['f_beta']:.4f}\")\n    else:\n        print(f\"Model not found at {model_path}, creating new model\")\n        model = EnsembleNet(\n            input_channels=1,\n            base_filters=16\n        ).to(DEVICE)\n    \n    # Generate predictions\n    predictions_df = predict(model, test_loader, DEVICE, threshold=config['threshold'])\n    \n    # Save predictions\n    output_file = os.path.join(OUTPUT_DIR, 'submission.csv')\n    predictions_df.to_csv(output_file, index=False)\n    \n    # Print statistics\n    motor_count = (predictions_df[['Motor axis 0', 'Motor axis 1', 'Motor axis 2']] != -1).all(axis=1).sum()\n    print(f\"Created submission file with {len(predictions_df)} predictions\")\n    print(f\"Number of motors predicted: {motor_count}\")\n    print(f\"Percentage of motors predicted: {motor_count / len(predictions_df) * 100:.2f}%\")\n    \n    return predictions_df\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-06T08:58:01.340058Z","iopub.execute_input":"2025-03-06T08:58:01.340325Z","iopub.status.idle":"2025-03-06T08:58:01.358963Z","shell.execute_reply.started":"2025-03-06T08:58:01.340306Z","shell.execute_reply":"2025-03-06T08:58:01.358133Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 7. Run the Train Pipeline\n","metadata":{}},{"cell_type":"code","source":"# Train the model\nmodel = train_model()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-06T08:58:01.359792Z","iopub.execute_input":"2025-03-06T08:58:01.360085Z","execution_failed":"2025-03-06T09:01:38.231Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Generate predictions \npredictions_df = generate_predictions()","metadata":{"trusted":true,"execution":{"execution_failed":"2025-03-06T09:01:38.232Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 8. Analysis and Visualization\n","metadata":{}},{"cell_type":"code","source":"def visualize_predictions(predictions_df, sample_count=3):\n    \"\"\"Visualize a few sample predictions\"\"\"\n    # Select samples with and without motors\n    motors_present = predictions_df[predictions_df['Motor axis 0'] != -1].sample(min(sample_count, len(predictions_df[predictions_df['Motor axis 0'] != -1])))\n    motors_absent = predictions_df[predictions_df['Motor axis 0'] == -1].sample(min(sample_count, len(predictions_df[predictions_df['Motor axis 0'] == -1])))\n    \n    # Combine the samples\n    samples = pd.concat([motors_present, motors_absent])\n    \n    # Display predictions\n    print(\"Sample predictions:\")\n    for _, row in samples.iterrows():\n        tomo_id = row['tomo_id']\n        if row['Motor axis 0'] == -1:\n            print(f\"Tomogram {tomo_id}: No motor detected\")\n        else:\n            coords = (row['Motor axis 0'], row['Motor axis 1'], row['Motor axis 2'])\n            print(f\"Tomogram {tomo_id}: Motor detected at coordinates {coords}\")\n\n    # You could add code here to visualize specific tomogram slices with overlaid predictions\n    # This would require loading the tomograms and plotting slices near the predicted motor location\n","metadata":{"trusted":true,"execution":{"execution_failed":"2025-03-06T09:01:38.232Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 9. Conclusion and Next Steps\n\nI've built an efficient data loading pipeline for variable-sized tomograms and designed a 3D CNN with dual-output heads for detection and localization. I also implemented a custom loss function optimized for the F-beta metric and set up proper training and validation procedures. Next, I'll focus on ensembling multiple models, expanding data augmentation, experimenting with 3D U-Net, fine-tuning hyperparameters, and applying cross-validation for a more robust evaluation.\n","metadata":{}},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}