{"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":"code","source":"import 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":{"execution":{"iopub.status.busy":"2025-03-29T12:03:42.825319Z","iopub.execute_input":"2025-03-29T12:03:42.825694Z","iopub.status.idle":"2025-03-29T12:03:45.444655Z","shell.execute_reply.started":"2025-03-29T12:03:42.825670Z","shell.execute_reply":"2025-03-29T12:03:45.443698Z"},"trusted":true},"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":{"execution":{"iopub.status.busy":"2025-03-29T12:03:45.451200Z","iopub.execute_input":"2025-03-29T12:03:45.451404Z","iopub.status.idle":"2025-03-29T12:03:45.482617Z","shell.execute_reply.started":"2025-03-29T12:03:45.451384Z","shell.execute_reply":"2025-03-29T12:03:45.481496Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 1. Data","metadata":{}},{"cell_type":"code","source":"class TomogramDataset(Dataset):\n    \"\"\"\n    Dataset for loading 3D tomograms from stacks of 2D JPG slices.\n    Handles both training and test data.\n    \"\"\"\n    def __init__(self, root_dir, max_slices=64, target_size=(128, 128), train=True, csv_file=None, exclude_no_motor=False):\n        self.root_dir = root_dir\n        self.max_slices = max_slices\n        self.target_size = target_size\n        self.train = train\n        self.exclude_no_motor = exclude_no_motor\n        \n        if train:\n            # Training mode - load labels from CSV\n            if csv_file is None:\n                raise ValueError(\"csv_file must be provided for training mode\")\n            self.labels_df = pd.read_csv(csv_file)\n            self.process_metadata()\n        else:\n            # Test mode - get tomogram directories directly\n            self.tomo_dirs = sorted([d for d in os.listdir(root_dir) \n                                   if os.path.isdir(os.path.join(root_dir, d))])\n        \n        # Cache file paths\n        self.cache_file_paths()\n    \n    def process_metadata(self):\n        \"\"\"\n        모든 motor를 개별 데이터 포인트로 처리하도록 수정\n        각 motor마다 고유한 ID 부여\n        exclude_no_motor가 True인 경우 motor가 없는 tomogram 제외\n        \"\"\"\n        processed_data = []\n        \n        # 각 tomogram에 대해\n        for tomo_id in self.labels_df['tomo_id'].unique():\n            tomo_rows = self.labels_df[self.labels_df['tomo_id'] == tomo_id]\n            \n            # 기본 tomogram 정보 가져오기\n            base_info = {\n                'original_tomo_id': tomo_id,  # 원본 tomo_id 보존\n                'Array shape (axis 0)': tomo_rows['Array shape (axis 0)'].iloc[0],\n                'Array shape (axis 1)': tomo_rows['Array shape (axis 1)'].iloc[0],\n                'Array shape (axis 2)': tomo_rows['Array shape (axis 2)'].iloc[0],\n                'Voxel spacing': tomo_rows['Voxel spacing'].iloc[0],\n                'Number of motors': tomo_rows['Number of motors'].iloc[0]\n            }\n            \n            num_motors = base_info['Number of motors']\n            \n            if num_motors == 0:\n                # motor가 없는 경우\n                if not self.exclude_no_motor:  # exclude_no_motor가 False일 때만 추가\n                    motor_info = base_info.copy()\n                    motor_info.update({\n                        'tomo_id': f\"{tomo_id}_no_motor\",  # 고유 ID 생성\n                        'Motor axis 0': -1,\n                        'Motor axis 1': -1,\n                        'Motor axis 2': -1\n                    })\n                    processed_data.append(motor_info)\n            else:\n                # motor가 있는 경우 - 각 motor를 개별 데이터 포인트로 추가\n                motor_rows = tomo_rows[tomo_rows['Motor axis 0'] != -1]\n                for motor_idx, motor_row in enumerate(motor_rows.iterrows()):\n                    motor_info = base_info.copy()\n                    _, row = motor_row\n                    motor_info.update({\n                        'tomo_id': f\"{tomo_id}_motor_{motor_idx}\",  # 고유 ID 생성\n                        'Motor axis 0': row['Motor axis 0'],\n                        'Motor axis 1': row['Motor axis 1'],\n                        'Motor axis 2': row['Motor axis 2']\n                    })\n                    processed_data.append(motor_info)\n        \n        # DataFrame으로 변환\n        self.tomo_df = pd.DataFrame(processed_data)\n        \n        # 데이터셋 정보 출력\n        total_samples = len(self.tomo_df)\n        motor_samples = len(self.tomo_df[self.tomo_df['Motor axis 0'] != -1])\n        print(f\"Dataset statistics:\")\n        print(f\"Total samples: {total_samples}\")\n        print(f\"Samples with motors: {motor_samples}\")\n        print(f\"Samples without motors: {total_samples - motor_samples}\")\n    \n    def cache_file_paths(self):\n        \"\"\"\n        파일 경로 캐싱 메서드 수정\n        original_tomo_id를 사용하여 파일 경로 찾기\n        \"\"\"\n        self.slice_files = {}\n        \n        if self.train:\n            for _, row in self.tomo_df.iterrows():\n                tomo_id = row['original_tomo_id']  # 원본 tomo_id 사용\n                if tomo_id not in self.slice_files:  # 중복 처리 방지\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        else:\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 load_volume(self, tomo_id):\n        \"\"\"\n        볼륨 로딩 메서드 수정\n        original_tomo_id를 사용하여 파일 로드\n        \"\"\"\n        # tomo_id에서 원본 ID 추출\n        original_tomo_id = tomo_id.split('_motor_')[0] if '_motor_' in tomo_id else tomo_id.split('_no_motor')[0]\n        files = self.slice_files[original_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            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')\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 __len__(self):\n        return len(self.tomo_df) if self.train else len(self.tomo_dirs)\n        \n    def __getitem__(self, idx):\n        if self.train:\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 (random jitter to coordinates)\n                if random.random() < 0.5:\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 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 and ensure correct shape\n            # volume shape should be [C, D, H, W] for single sample\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,  # Shape: [1, D, H, W]\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        else:\n            # Test mode\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 and ensure correct shape\n            volume = torch.from_numpy(volume).unsqueeze(0)  # Add channel dimension\n            \n            return {\n                'tomo_id': tomo_id,\n                'volume': volume,  # Shape: [1, D, H, W]\n                'original_shape': torch.tensor(original_shape, dtype=torch.float32)\n            }","metadata":{"execution":{"iopub.status.busy":"2025-03-29T12:03:45.486348Z","iopub.execute_input":"2025-03-29T12:03:45.486664Z","iopub.status.idle":"2025-03-29T12:03:45.510305Z","shell.execute_reply.started":"2025-03-29T12:03:45.486631Z","shell.execute_reply":"2025-03-29T12:03:45.509419Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 2. Model\n","metadata":{}},{"cell_type":"code","source":"from torchvision.ops import StochasticDepth\nfrom typing import List, Dict\nfrom torch import Tensor\n#from torchtune.modules import RotaryPositionalEmbeddings\n\nclass ConvNextStem(nn.Sequential):\n    def __init__(self, in_features: int, out_features: int):\n        super().__init__(\n            nn.Conv3d(in_features, out_features, kernel_size=(1, 2, 2), stride=(1, 2, 2)),\n            nn.GroupNorm(num_groups=1, num_channels=out_features)\n        )\n\nclass LayerScaler(nn.Module):\n    def __init__(self, init_value: float, dimensions: int):\n        super().__init__()\n        self.gamma = nn.Parameter(init_value * torch.ones((dimensions)),\n                                    requires_grad=True)\n\n    def forward(self, x):\n        return self.gamma[None,...,None,None] * x\n\nclass BottleNeckBlock(nn.Module):\n    def __init__(\n        self,\n        in_features: int,\n        out_features: int,\n        expansion: int = 4,\n        drop_p: float = .0,\n        layer_scaler_init_value: float = 1e-6,\n    ):\n        super().__init__()\n        expanded_features = out_features * expansion\n        self.block = nn.Sequential(\n            # narrow -> wide (with depth-wise and bigger kernel)\n            nn.Conv3d(\n                in_features, in_features, kernel_size=(2, 7, 7), padding='same', bias=False, groups=in_features\n            ),\n            # GroupNorm with num_groups=1 is the same as LayerNorm but works for 2D data\n            nn.GroupNorm(num_groups=in_features, num_channels=in_features),\n            # wide -> wide\n            nn.Conv3d(in_features, expanded_features, kernel_size=1),\n            nn.GELU(),\n            # wide -> narrow\n            nn.Conv3d(expanded_features, out_features, kernel_size=1),\n        )\n        #self.layer_scaler = LayerScaler(layer_scaler_init_value, out_features)\n        #self.drop_path = StochasticDepth(drop_p, mode=\"batch\")\n\n\n    def forward(self, x: Tensor) -> Tensor:\n        res = x\n        x = self.block(x)\n        #x = self.layer_scaler(x)\n        #x = self.drop_path(x)\n        x += res\n        return x\n\nclass ConvNexStage(nn.Sequential):\n    def __init__(\n        self, in_features: int, out_features: int, depth: int, **kwargs\n    ):\n        super().__init__(\n            # add the downsampler\n            nn.Sequential(\n                nn.GroupNorm(num_groups=in_features, num_channels=in_features),\n                nn.Conv3d(in_features, out_features, kernel_size=(2, 2, 2), stride=(2, 2, 2))\n            ),\n            *[\n                BottleNeckBlock(out_features, out_features, **kwargs)\n                for _ in range(depth)\n            ],\n        )\n\nclass ConvNextEncoder(nn.Module):\n    def __init__(\n        self,\n        in_channels: int,\n        stem_features: int,\n        depths: List[int],\n        widths: List[int],\n        drop_p: float = .0,\n    ):\n        super().__init__()\n        self.stem = ConvNextStem(in_channels, stem_features)\n\n        in_out_widths = list(zip(widths, widths[1:]))\n        # create drop paths probabilities (one for each stage)\n        drop_probs = [x.item() for x in torch.linspace(0, drop_p, sum(depths))]\n\n        self.stages = nn.ModuleList(\n            [\n                ConvNexStage(stem_features, widths[0], depths[0], drop_p=drop_probs[0]),\n                *[\n                    ConvNexStage(in_features, out_features, depth, drop_p=drop_p)\n                    for (in_features, out_features), depth, drop_p in zip(\n                        in_out_widths, depths[1:], drop_probs[1:]\n                    )\n                ],\n            ]\n        )\n\n\n    def forward(self, x):\n        x = self.stem(x)\n        for stage in self.stages:\n            x = stage(x)\n        return x\n\nclass ConvNextClassifier(nn.Module):\n    def __init__(self, max_slices):\n        super().__init__()\n        self.encoder = ConvNextEncoder(in_channels=1, stem_features=max_slices, depths=[3,3,9,3], widths=[64, 128, 256, 512])\n        self.flatten = nn.Sequential(nn.AdaptiveAvgPool3d((1,1,1)),\n                                    nn.Flatten(1),\n                                    nn.LayerNorm(512),\n                                     )\n        self.l1 = nn.Linear(512, max_slices) # (exists, not exists, exists, ... )\n\n    def forward(self, x, label=None):\n        x = self.encoder(x)\n        x = self.flatten(x)\n        l1 = self.l1(x)\n\n        return l1\n    \nclass ConvNextRegressor(nn.Module):\n    def __init__(self, max_slices):\n        super().__init__()\n        self.encoder = ConvNextEncoder(in_channels=1, stem_features=max_slices, depths=[3,3,9,3], widths=[64, 128, 256, 512])\n        self.flatten = nn.Sequential(nn.AdaptiveAvgPool3d((1,1,1)),\n                                    nn.Flatten(1),\n                                    nn.LayerNorm(512),\n                                     )\n        self.l1 = nn.Linear(512, 3)\n\n    def forward(self, x, label=None):\n        x = self.encoder(x)\n        x = self.flatten(x)\n        l1 = self.l1(x)\n\n        return l1","metadata":{"execution":{"iopub.status.busy":"2025-03-29T12:03:45.511212Z","iopub.execute_input":"2025-03-29T12:03:45.511454Z","iopub.status.idle":"2025-03-29T12:03:46.657170Z","shell.execute_reply.started":"2025-03-29T12:03:45.511434Z","shell.execute_reply":"2025-03-29T12:03:46.656471Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3. Train","metadata":{}},{"cell_type":"code","source":"class FlagellarMotorClassificationLoss(nn.Module):\n    \"\"\"\n    Loss function for motor presence detection (classification task).\n    Uses binary cross-entropy loss.\n    \"\"\"\n    def __init__(self):\n        super(FlagellarMotorClassificationLoss, self).__init__()\n        self.bce_loss = nn.BCELoss()\n    \n    def forward(self, presence_pred, presence_true):\n        \"\"\"\n        Args:\n            presence_pred (Tensor): Predicted probability of motor presence [B, 1]\n            presence_true (Tensor): Ground truth motor presence [B, 1]\n        Returns:\n            Tensor: Classification loss\n        \"\"\"\n        return self.bce_loss(presence_pred, presence_true)\n\nclass FlagellarMotorRegressorLoss(nn.Module):\n    \"\"\"\n    Loss function for motor location regression.\n    Uses MSE loss only for samples with motors present.\n    \"\"\"\n    def __init__(self):\n        super(FlagellarMotorRegressorLoss, self).__init__()\n        self.mse_loss = nn.MSELoss()\n    \n    def forward(self, location_pred, location_true, presence_true):\n        \"\"\"\n        Args:\n            location_pred (Tensor): Predicted motor coordinates [B, 4]\n            location_true (Tensor): Ground truth motor coordinates [B, 4]\n        Returns:\n            Tuple[Tensor, Tensor]: (Location loss, Average Euclidean distance)\n        \"\"\"\n        # Initialize loss and distance\n        location_loss = torch.tensor(0.0, device=location_pred.device)\n        avg_euclidean_dist = torch.tensor(0.0, device=location_pred.device)\n        \n        # Only compute loss for samples with motors\n        has_motor = presence_true.squeeze() > 0.5\n        if torch.sum(has_motor) > 0:\n            location_pred_with_motor = location_pred[has_motor]\n            location_true_with_motor = location_true[has_motor]\n            \n            # Calculate MSE loss\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(\n                (location_pred_with_motor - location_true_with_motor) ** 2, \n                dim=1\n            ))\n            avg_euclidean_dist = euclidean_dist.mean()\n        \n        return location_loss","metadata":{"execution":{"iopub.status.busy":"2025-03-29T12:03:46.657991Z","iopub.execute_input":"2025-03-29T12:03:46.658481Z","iopub.status.idle":"2025-03-29T12:03:46.665239Z","shell.execute_reply.started":"2025-03-29T12:03:46.658447Z","shell.execute_reply":"2025-03-29T12:03:46.664343Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_epoch(model, dataloader, optimizer, criterion, device):\n    \"\"\"Train for one epoch\"\"\"\n    model.train()\n    epoch_loss = 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        location_pred = model(volume)\n        \n        # Calculate loss\n        loss = criterion(\n            location_pred, motor_axes, has_motor\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        \n        # Update progress bar\n        progress_bar.set_postfix({\n            'loss': loss.item(),\n        })\n    \n    # Calculate average metrics\n    num_batches = len(dataloader)\n    avg_loss = epoch_loss / num_batches\n    \n    return {\n        'loss': avg_loss\n    }\n","metadata":{"execution":{"iopub.status.busy":"2025-03-29T12:03:46.667687Z","iopub.execute_input":"2025-03-29T12:03:46.667924Z","iopub.status.idle":"2025-03-29T12:03:46.685960Z","shell.execute_reply.started":"2025-03-29T12:03:46.667904Z","shell.execute_reply":"2025-03-29T12:03:46.685109Z"},"trusted":true},"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    \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            location_pred = model(volume)\n            \n            # Calculate loss\n            loss = criterion(\n                location_pred, motor_axes, has_motor\n            )\n            \n            # Update metrics\n            epoch_loss += loss.item()\n            \n            # Calculate F-beta metrics\n            for i in range(len(location_pred)):\n                # Check if model predicts a motor\n                pred_has_motor = True\n                true_has_motor = True\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            })\n    \n    # Calculate average metrics\n    num_batches = len(dataloader)\n    avg_loss = epoch_loss / 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        '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":{"execution":{"iopub.status.busy":"2025-03-29T12:03:46.686956Z","iopub.execute_input":"2025-03-29T12:03:46.687152Z","iopub.status.idle":"2025-03-29T12:03:46.706144Z","shell.execute_reply.started":"2025-03-29T12:03:46.687134Z","shell.execute_reply":"2025-03-29T12:03:46.705461Z"},"trusted":true},"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            location_pred = model(volume)\n            \n            # Process predictions\n            for i in range(len(location_pred)):\n                tomo_id = tomo_ids[i]\n                pred_has_motor = True\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":{"execution":{"iopub.status.busy":"2025-03-29T12:03:46.706963Z","iopub.execute_input":"2025-03-29T12:03:46.707203Z","iopub.status.idle":"2025-03-29T12:03:46.725719Z","shell.execute_reply.started":"2025-03-29T12:03:46.707184Z","shell.execute_reply":"2025-03-29T12:03:46.725049Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 4. Run","metadata":{}},{"cell_type":"code","source":"\n\ndef run():\n    \"\"\"Train the model and save checkpoints\"\"\"\n    # Configuration\n    config = {\n        'batch_size': 32,\n        'num_workers': 2,\n        'max_slices': 32,\n        'target_size': (128, 128),\n        'learning_rate': 0.0005,\n        'weight_decay': 0.0001,\n        'epochs': 5, \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        exclude_no_motor=True\n    )\n    \n    val_dataset = TomogramDataset(\n        csv_file=val_csv,\n        root_dir=TRAIN_DIR,\n        train=True,\n        max_slices=config['max_slices'],\n        target_size=config['target_size'],\n        exclude_no_motor=True\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 = ConvNextRegressor(\n        max_slices=config['max_slices']\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 = FlagellarMotorRegressorLoss()\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\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":{"execution":{"iopub.status.busy":"2025-03-29T12:09:49.099829Z","iopub.execute_input":"2025-03-29T12:09:49.100217Z","iopub.status.idle":"2025-03-29T12:09:49.118019Z","shell.execute_reply.started":"2025-03-29T12:09:49.100183Z","shell.execute_reply":"2025-03-29T12:09:49.117078Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = run()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-29T12:09:51.216905Z","iopub.execute_input":"2025-03-29T12:09:51.217237Z","iopub.status.idle":"2025-03-29T12:18:23.776960Z","shell.execute_reply.started":"2025-03-29T12:09:51.217209Z","shell.execute_reply":"2025-03-29T12:18:23.775610Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 5. Predict\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': 32,\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 = TomogramDataset(\n        root_dir=TEST_DIR,\n        max_slices=config['max_slices'],\n        target_size=config['target_size'],\n        train=False\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 = ConvNextRegressor(\n            config[\"max_slices\"]\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 = ConvNextRegressor(\n            config[\"max_slices\"]\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":{"execution":{"iopub.status.busy":"2025-03-29T12:05:15.430801Z","iopub.status.idle":"2025-03-29T12:05:15.431067Z","shell.execute_reply":"2025-03-29T12:05:15.430962Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 6. Visualize","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,"execution":{"iopub.status.busy":"2025-03-29T12:05:15.432040Z","iopub.status.idle":"2025-03-29T12:05:15.432433Z","shell.execute_reply":"2025-03-29T12:05:15.432260Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}