{"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-06T00:46:59.259051Z","iopub.execute_input":"2025-03-06T00:46:59.259492Z","iopub.status.idle":"2025-03-06T00:47:01.740034Z","shell.execute_reply.started":"2025-03-06T00:46:59.259448Z","shell.execute_reply":"2025-03-06T00:47:01.739363Z"}},"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-06T00:47:01.741049Z","iopub.execute_input":"2025-03-06T00:47:01.741434Z","iopub.status.idle":"2025-03-06T00:47:01.768620Z","shell.execute_reply.started":"2025-03-06T00:47:01.741402Z","shell.execute_reply":"2025-03-06T00:47:01.767787Z"}},"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-06T00:47:01.770051Z","iopub.execute_input":"2025-03-06T00:47:01.770287Z","iopub.status.idle":"2025-03-06T00:47:03.313565Z","shell.execute_reply.started":"2025-03-06T00:47:01.770254Z","shell.execute_reply":"2025-03-06T00:47:03.312629Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 2. Dataset Class\n","metadata":{}},{"cell_type":"code","source":"\nclass TomogramDataset(Dataset):\n    \"\"\"\n    Dataset for loading 3D tomograms from stacks of 2D JPG slices.\n    \"\"\"\n    def __init__(self, csv_file, root_dir, train=True, max_slices=64, target_size=(128, 128)):\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        # 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 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        # Process coordinates\n        if not has_motor:\n            motor_axes = np.zeros(3, dtype=np.float32)\n        else:\n            # Get original shape\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            # Apply data augmentation in training (random jitter to coordinates)\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        # Convert to tensor\n        volume = torch.from_numpy(volume).unsqueeze(0)  # 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        }\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-06T00:47:03.314749Z","iopub.execute_input":"2025-03-06T00:47:03.315008Z","iopub.status.idle":"2025-03-06T00:47:03.332107Z","shell.execute_reply.started":"2025-03-06T00:47:03.314983Z","shell.execute_reply":"2025-03-06T00:47:03.331296Z"}},"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-06T00:47:03.332947Z","iopub.execute_input":"2025-03-06T00:47:03.333171Z","iopub.status.idle":"2025-03-06T00:47:03.358403Z","shell.execute_reply.started":"2025-03-06T00:47:03.333151Z","shell.execute_reply":"2025-03-06T00:47:03.357549Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3. Model Architecture\n","metadata":{}},{"cell_type":"code","source":"\n\nclass FlagellarMotorNet(nn.Module):\n    \"\"\"\n    3D CNN model for detecting and localizing flagellar motors in tomograms.\n    \"\"\"\n    def __init__(self, input_channels=1, base_filters=16):\n        super(FlagellarMotorNet, self).__init__()\n        \n        # Feature extraction layers\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.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.pool2 = nn.MaxPool3d(kernel_size=2)\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.pool3 = nn.MaxPool3d(kernel_size=2)\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.pool4 = nn.MaxPool3d(kernel_size=2)\n        \n        # Calculate the size of the flattened features\n        # Assuming input size of [1, 64, 128, 128]\n        # After 4 pooling layers (each dividing by 2): [base_filters*8, 4, 8, 8]\n        self.fc_size = base_filters * 8 * 4 * 8 * 8\n        \n        # Classification head (motor presence)\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        \n        # Regression head (motor location)\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()  # Normalize coordinates to [0, 1]\n        )\n    \n    def forward(self, x):\n        # Feature extraction\n        x = F.relu(self.bn1(self.conv1(x)))\n        x = self.pool1(x)\n        \n        x = F.relu(self.bn2(self.conv2(x)))\n        x = self.pool2(x)\n        \n        x = F.relu(self.bn3(self.conv3(x)))\n        x = self.pool3(x)\n        \n        x = F.relu(self.bn4(self.conv4(x)))\n        x = self.pool4(x)\n        \n        # Flatten\n        x = x.view(x.size(0), -1)\n        \n        # Motor presence prediction\n        presence = self.fc_presence(x)\n        \n        # Motor location prediction\n        location = self.fc_location(x)\n        \n        return presence, location\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-06T00:47:03.359233Z","iopub.execute_input":"2025-03-06T00:47:03.359537Z","iopub.status.idle":"2025-03-06T00:47:03.377726Z","shell.execute_reply.started":"2025-03-06T00:47:03.359504Z","shell.execute_reply":"2025-03-06T00:47:03.377032Z"}},"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-06T00:47:03.378509Z","iopub.execute_input":"2025-03-06T00:47:03.378766Z","iopub.status.idle":"2025-03-06T00:47:03.396813Z","shell.execute_reply.started":"2025-03-06T00:47:03.378733Z","shell.execute_reply":"2025-03-06T00:47:03.396008Z"}},"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-06T00:47:03.398992Z","iopub.execute_input":"2025-03-06T00:47:03.399187Z","iopub.status.idle":"2025-03-06T00:47:03.413384Z","shell.execute_reply.started":"2025-03-06T00:47:03.399170Z","shell.execute_reply":"2025-03-06T00:47:03.412638Z"}},"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-06T00:47:03.414418Z","iopub.execute_input":"2025-03-06T00:47:03.414614Z","iopub.status.idle":"2025-03-06T00:47:03.428126Z","shell.execute_reply.started":"2025-03-06T00:47:03.414596Z","shell.execute_reply":"2025-03-06T00:47:03.427312Z"}},"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-06T00:47:03.429089Z","iopub.execute_input":"2025-03-06T00:47:03.429398Z","iopub.status.idle":"2025-03-06T00:47:03.449754Z","shell.execute_reply.started":"2025-03-06T00:47:03.429367Z","shell.execute_reply":"2025-03-06T00:47:03.449025Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 5. Training Loop\n","metadata":{}},{"cell_type":"code","source":"\n\ndef train_model():\n    \"\"\"Train the model and save checkpoints\"\"\"\n    # Configuration\n    config = {\n        'batch_size': 32,\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 = FlagellarMotorNet(\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-06T00:47:03.450555Z","iopub.execute_input":"2025-03-06T00:47:03.450751Z","iopub.status.idle":"2025-03-06T00:47:03.472498Z","shell.execute_reply.started":"2025-03-06T00:47:03.450734Z","shell.execute_reply":"2025-03-06T00:47:03.471569Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 6. Generate Predictions\n","metadata":{}},{"cell_type":"code","source":"\n\ndef 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 = FlagellarMotorNet(\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 = FlagellarMotorNet(\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-06T00:47:03.473178Z","iopub.execute_input":"2025-03-06T00:47:03.473412Z","iopub.status.idle":"2025-03-06T00:47:03.492676Z","shell.execute_reply.started":"2025-03-06T00:47:03.473379Z","shell.execute_reply":"2025-03-06T00:47:03.491822Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 7. Run the Train Pipeline\n","metadata":{}},{"cell_type":"code","source":"\n\n# Train the model\nmodel = train_model()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-06T00:47:03.493465Z","iopub.execute_input":"2025-03-06T00:47:03.493682Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Generate predictions \npredictions_df = generate_predictions()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 8. Analysis and Visualization\n","metadata":{}},{"cell_type":"code","source":"\n\ndef 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},"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}]}