{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":91249,"databundleVersionId":11294684,"sourceType":"competition"}],"dockerImageVersionId":31041,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# PART 0 – Visualization & Exploratory Data Analysis (EDA)\n","metadata":{}},{"cell_type":"markdown","source":"## Step 0.1 – Explore the File Structure\n","metadata":{}},{"cell_type":"code","source":"import os\nfrom glob import glob\nfrom collections import defaultdict\nfrom PIL import Image\nimport numpy as np\nimport matplotlib.pyplot as plt\n\n# Set base paths\nTRAIN_DIR = \"/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/train\"\nTEST_DIR = \"/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/test\"\n\n# List tomograms\ntrain_tomos = sorted(os.listdir(TRAIN_DIR))\ntest_tomos = sorted(os.listdir(TEST_DIR))\n\nprint(f\"🧪 Number of train tomograms: {len(train_tomos)}\")\nprint(f\"🧪 Number of test tomograms: {len(test_tomos)}\")\n\n# Count number of slices per tomogram\nslices_per_tomo = {}\nfor tomo in train_tomos:\n    slices = glob(os.path.join(TRAIN_DIR, tomo, '*.jpg'))\n    slices_per_tomo[tomo] = len(slices)\n\n# Summary statistics for slices\nslice_counts = list(slices_per_tomo.values())\nprint(f\"\\n🧩 Slices per tomogram (train):\")\nprint(f\"  ➤ Min: {np.min(slice_counts)}\")\nprint(f\"  ➤ Max: {np.max(slice_counts)}\")\nprint(f\"  ➤ Mean: {np.mean(slice_counts):.2f}\")\nprint(f\"  ➤ Median: {np.median(slice_counts)}\")\n\n# Plot histogram\nplt.figure(figsize=(10,6))\nplt.hist(slice_counts, bins=20, color='skyblue', edgecolor='black')\nplt.title(\"Distribution of Number of Slices per Tomogram (Train Set)\")\nplt.xlabel(\"Number of slices\")\nplt.ylabel(\"Number of tomograms\")\nplt.grid(True)\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-03T14:39:47.711135Z","iopub.execute_input":"2025-07-03T14:39:47.711535Z","iopub.status.idle":"2025-07-03T14:40:00.889255Z","shell.execute_reply.started":"2025-07-03T14:39:47.711512Z","shell.execute_reply":"2025-07-03T14:40:00.888538Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Step 0.2 – Load Annotations (e.g. CSV or JSON)\n","metadata":{}},{"cell_type":"code","source":"\nimport pandas as pd\n\nannotations_path = '/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/train_labels.csv'\ndf = pd.read_csv(annotations_path)\n\n# Preview the first few rows\ndf.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-03T14:40:00.890533Z","iopub.execute_input":"2025-07-03T14:40:00.890791Z","iopub.status.idle":"2025-07-03T14:40:01.171623Z","shell.execute_reply.started":"2025-07-03T14:40:00.890773Z","shell.execute_reply":"2025-07-03T14:40:01.171062Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Step 0.3 – Basic Summary of the Dataset\n","metadata":{}},{"cell_type":"code","source":"# - How many tomograms do we have?\n# - How many contain motors?\n# - How many motors per tomogram?\n\nmotor_counts = df.groupby('tomo_id')['Number of motors'].max()\n\nprint(\"🔢 Total tomograms with annotations:\", df['tomo_id'].nunique())\nprint(\"🛑 Tomograms with no motors:\", (motor_counts == 0).sum())\nprint(\"✅ Tomograms with at least 1 motor:\", (motor_counts > 0).sum())\n\nprint(\"\\n📊 Motor Count Statistics:\")\nprint(\"  ➤ Min:\", motor_counts.min())\nprint(\"  ➤ Max:\", motor_counts.max())\nprint(\"  ➤ Mean:\", motor_counts.mean())\nprint(\"  ➤ Median:\", motor_counts.median())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-03T14:40:01.172358Z","iopub.execute_input":"2025-07-03T14:40:01.172611Z","iopub.status.idle":"2025-07-03T14:40:01.187612Z","shell.execute_reply.started":"2025-07-03T14:40:01.172587Z","shell.execute_reply":"2025-07-03T14:40:01.186907Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Step 0.4 – Plot: Number of Motors per Tomogram\n","metadata":{}},{"cell_type":"code","source":"import seaborn as sns\nimport matplotlib.pyplot as plt\n\nplt.figure(figsize=(10,6))\nsns.histplot(motor_counts, bins=range(0, motor_counts.max()+2), discrete=True, color='purple')\nplt.title(\"Number of Motors per Tomogram\")\nplt.xlabel(\"Number of Motors\")\nplt.ylabel(\"Number of Tomograms\")\nplt.grid(True)\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-03T14:40:01.188592Z","iopub.execute_input":"2025-07-03T14:40:01.188832Z","iopub.status.idle":"2025-07-03T14:40:02.080056Z","shell.execute_reply.started":"2025-07-03T14:40:01.188811Z","shell.execute_reply":"2025-07-03T14:40:02.07937Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Step 0.5 – Analyze Motor Coordinates (x, y, z)\n","metadata":{}},{"cell_type":"code","source":"valid_motors = df[df['Number of motors'] > 0].copy()\nvalid_motors = valid_motors[valid_motors[['Motor axis 0', 'Motor axis 1', 'Motor axis 2']].min(axis=1) >= 0]\n\nfig, axs = plt.subplots(1, 3, figsize=(18,5))\nfor ax, coord in zip(axs, ['Motor axis 0', 'Motor axis 1', 'Motor axis 2']):\n    sns.histplot(valid_motors[coord], bins=50, ax=ax)\n    ax.set_title(f\"Distribution of {coord}\")\n    ax.set_xlabel(coord)\n    ax.set_ylabel(\"Frequency\")\n    ax.grid(True)\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-03T14:40:02.081787Z","iopub.execute_input":"2025-07-03T14:40:02.082286Z","iopub.status.idle":"2025-07-03T14:40:02.784346Z","shell.execute_reply.started":"2025-07-03T14:40:02.082266Z","shell.execute_reply":"2025-07-03T14:40:02.783637Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Step 0.6 – Optional: 3D Scatter Plot of Motor Positions\n","metadata":{}},{"cell_type":"code","source":"from mpl_toolkits.mplot3d import Axes3D\n\nfig = plt.figure(figsize=(10, 8))\nax = fig.add_subplot(111, projection='3d')\n\nax.scatter(valid_motors['Motor axis 0'],\n           valid_motors['Motor axis 1'],\n           valid_motors['Motor axis 2'],\n           c='red', alpha=0.5, s=10)\n\nax.set_xlabel(\"X (Axis 0)\")\nax.set_ylabel(\"Y (Axis 1)\")\nax.set_zlabel(\"Z (Axis 2)\")\nax.set_title(\"3D Distribution of Motor Coordinates\")\n\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-03T14:40:02.785075Z","iopub.execute_input":"2025-07-03T14:40:02.785379Z","iopub.status.idle":"2025-07-03T14:40:02.942315Z","shell.execute_reply.started":"2025-07-03T14:40:02.785354Z","shell.execute_reply":"2025-07-03T14:40:02.94146Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Step 0.7 – Analyze Array Shape (Volume Dimensions)\n","metadata":{}},{"cell_type":"code","source":"z_dim = df['Array shape (axis 0)']\ny_dim = df['Array shape (axis 1)']\nx_dim = df['Array shape (axis 2)']\n\nprint(\"📦 Tomogram Dimensions (Z,Y,X):\")\nprint(f\"  ➤ Z-axis (depth):   min={z_dim.min()}, max={z_dim.max()}, mean={z_dim.mean():.2f}\")\nprint(f\"  ➤ Y-axis (height):  min={y_dim.min()}, max={y_dim.max()}, mean={y_dim.mean():.2f}\")\nprint(f\"  ➤ X-axis (width):   min={x_dim.min()}, max={x_dim.max()}, mean={x_dim.mean():.2f}\")\n\n# Plot distributions\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\nfig, axs = plt.subplots(1, 3, figsize=(18, 5))\nfor ax, data, label in zip(axs, [z_dim, y_dim, x_dim], ['Z (depth)', 'Y (height)', 'X (width)']):\n    sns.histplot(data, bins=20, ax=ax)\n    ax.set_title(f'Distribution of {label}')\n    ax.set_xlabel(label)\n    ax.set_ylabel('Number of Tomograms')\n    ax.grid(True)\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-03T14:40:02.943175Z","iopub.execute_input":"2025-07-03T14:40:02.943418Z","iopub.status.idle":"2025-07-03T14:40:03.464976Z","shell.execute_reply.started":"2025-07-03T14:40:02.943399Z","shell.execute_reply":"2025-07-03T14:40:03.464356Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Step 0.8 – Analyze Voxel Spacing Distribution\n","metadata":{}},{"cell_type":"code","source":"\nvoxel_spacing = df['Voxel spacing']\n\nprint(\"🧮 Voxel Spacing Statistics:\")\nprint(f\"  ➤ Min:   {voxel_spacing.min()}\")\nprint(f\"  ➤ Max:   {voxel_spacing.max()}\")\nprint(f\"  ➤ Mean:  {voxel_spacing.mean():.2f}\")\nprint(f\"  ➤ Median:{voxel_spacing.median()}\")\n\n# Plot voxel spacing distribution\nplt.figure(figsize=(8, 5))\nsns.histplot(voxel_spacing, bins=20, color='teal', edgecolor='black')\nplt.title(\"Distribution of Voxel Spacing\")\nplt.xlabel(\"Voxel Spacing\")\nplt.ylabel(\"Number of Tomograms\")\nplt.grid(True)\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# PART 1 – Initial Candidate Detection using UNet2D\n","metadata":{}},{"cell_type":"markdown","source":"## 🧾 Step 1.1 – Filter Positive Tomograms Only\n","metadata":{}},{"cell_type":"code","source":"import pandas as pd\n\n# Load annotation file\ndf = pd.read_csv('/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/train_labels.csv')\n\n# Filter only tomograms that contain at least one motor\ndf_pos = df[df['Number of motors'] > 0]\npositive_tomo_ids = df_pos['tomo_id'].unique()\n\nprint(f\"Total positive tomograms: {len(positive_tomo_ids)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-03T14:41:47.068673Z","iopub.execute_input":"2025-07-03T14:41:47.068923Z","iopub.status.idle":"2025-07-03T14:41:47.084498Z","shell.execute_reply.started":"2025-07-03T14:41:47.068904Z","shell.execute_reply":"2025-07-03T14:41:47.08391Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 🧾 Step 1.2 – Build Slice-Level Dataset for Segmentation\n\n","metadata":{}},{"cell_type":"markdown","source":"### Original way - problem with Huge RAM","metadata":{}},{"cell_type":"code","source":"import os\nimport numpy as np\nfrom PIL import Image\nfrom glob import glob\nimport time\n\n# Path to training images\nTRAIN_IMG_DIR = \"/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/train\"\n\n# Create list of (image_path, mask) pairs\nslice_dataset = []\n\nstart_time = time.time()\nprint(f\"🔁 Starting slice dataset creation for {len(positive_tomo_ids)} tomograms...\")\n\nfor idx, tomo_id in enumerate(positive_tomo_ids):\n    if idx % 5 == 0:\n        print(f\"🧪 Processing tomogram {idx + 1}/{len(positive_tomo_ids)}: {tomo_id}\")\n\n    tomo_path = os.path.join(TRAIN_IMG_DIR, tomo_id)\n    slice_paths = sorted(glob(os.path.join(tomo_path, \"*.jpg\")))\n\n    # Get all motor coordinates for this tomogram\n    motor_coords = df_pos[df_pos['tomo_id'] == tomo_id][['Motor axis 0', 'Motor axis 1', 'Motor axis 2']].values\n\n    for slice_path in slice_paths:\n        # Extract slice index from filename\n        slice_index = int(os.path.basename(slice_path).split(\"_\")[-1].split(\".\")[0])\n        \n        # Check if any motor appears in this slice\n        motors_in_slice = []\n        for z, y, x in motor_coords:\n            if int(z) == slice_index:\n                motors_in_slice.append((int(y), int(x)))  # (row, col)\n\n        # Load image\n        img = Image.open(slice_path).convert(\"L\")\n        img = np.array(img)\n\n        # Create binary mask\n        mask = np.zeros_like(img, dtype=np.uint8)\n        for y, x in motors_in_slice:\n            if 0 <= y < img.shape[0] and 0 <= x < img.shape[1]:\n                mask[y, x] = 1  # optionally: draw small disk or gaussian here\n\n        # Add to dataset\n        slice_dataset.append((slice_path, mask))\n\nelapsed = time.time() - start_time\nprint(f\"✅ Slice dataset built with {len(slice_dataset)} samples in {elapsed:.2f} seconds.\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Better way to save RAM","metadata":{}},{"cell_type":"code","source":"import os\nimport numpy as np\nfrom PIL import Image\nfrom glob import glob\nimport time\n\n# 📁 Path to training images\nTRAIN_IMG_DIR = \"/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/train\"\n\n# 📦 Final dataset: list of (image_path, slice_index, list of (y,x) coords)\nslice_dataset = []\n\nstart_time = time.time()\nprint(f\"🔁 Starting lightweight slice dataset creation for {len(positive_tomo_ids)} tomograms...\")\n\nfor idx, tomo_id in enumerate(positive_tomo_ids):\n    if idx % 5 == 0:\n        print(f\"🧪 Processing tomogram {idx + 1}/{len(positive_tomo_ids)}: {tomo_id}\")\n\n    tomo_path = os.path.join(TRAIN_IMG_DIR, tomo_id)\n    slice_paths = sorted(glob(os.path.join(tomo_path, \"*.jpg\")))\n\n    # ➕ Get motor coordinates for this tomogram\n    motor_coords = df_pos[df_pos['tomo_id'] == tomo_id][['Motor axis 0', 'Motor axis 1', 'Motor axis 2']].values\n\n    for slice_path in slice_paths:\n        # 🎯 Get slice index from filename\n        slice_index = int(os.path.basename(slice_path).split(\"_\")[-1].split(\".\")[0])\n\n        # 🧠 Filter motors in this slice\n        motors_in_slice = [\n            (int(y), int(x)) for z, y, x in motor_coords if int(z) == slice_index\n        ]\n\n        # 📝 Save tuple with image path, index, and list of coords\n        slice_dataset.append((slice_path, slice_index, motors_in_slice))\n\nelapsed = time.time() - start_time\nprint(f\"✅ Slice metadata dataset created with {len(slice_dataset)} samples in {elapsed:.2f} seconds.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T22:41:38.038004Z","iopub.execute_input":"2025-06-30T22:41:38.038275Z","iopub.status.idle":"2025-06-30T22:41:39.303531Z","shell.execute_reply.started":"2025-06-30T22:41:38.038254Z","shell.execute_reply":"2025-06-30T22:41:39.302879Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 🧾 Step 1.3 – Define UNet2D Model","metadata":{}},{"cell_type":"code","source":"!pip install -q segmentation-models-pytorch\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T22:41:56.927477Z","iopub.execute_input":"2025-06-30T22:41:56.927919Z","iopub.status.idle":"2025-06-30T22:42:00.081069Z","shell.execute_reply.started":"2025-06-30T22:41:56.927896Z","shell.execute_reply":"2025-06-30T22:42:00.080068Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport segmentation_models_pytorch as smp\n\n# Define model using segmentation_models_pytorch\nmodel = smp.Unet(\n    encoder_name=\"resnet18\",        # lightweight encoder\n    encoder_weights=\"imagenet\",     # optional (can be None)\n    in_channels=1,                  # grayscale input\n    classes=1,                      # binary segmentation\n)\n\n# Example: move model to GPU if available\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel = model.to(device)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T22:41:44.847965Z","iopub.execute_input":"2025-06-30T22:41:44.848226Z","iopub.status.idle":"2025-06-30T22:41:49.273315Z","shell.execute_reply.started":"2025-06-30T22:41:44.8482Z","shell.execute_reply":"2025-06-30T22:41:49.272502Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 🧾  Step 1.4 – Dataset and Dataloaders with train/val/test split by tomograms\n\n\n","metadata":{}},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nfrom sklearn.model_selection import train_test_split\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import DataLoader, Dataset\nimport torchvision.transforms as T\nimport time\nfrom PIL import Image\nimport matplotlib.pyplot as plt\n\nclass SliceSegmentationDataset(Dataset):\n    def __init__(self, dataset_list, image_transform=None, mask_transform=None):\n        \"\"\"\n        dataset_list: list of tuples (image_path, slice_index, list of (y,x) coords)\n        \"\"\"\n        self.data = dataset_list\n        self.image_transform = image_transform\n        self.mask_transform = mask_transform\n\n    def __len__(self):\n        return len(self.data)\n\n    def __getitem__(self, idx):\n        img_path, slice_index, coords = self.data[idx]\n\n        # Load grayscale image\n        image = Image.open(img_path).convert(\"L\")\n        image_np = np.array(image)\n\n        # Create empty binary mask\n        mask_np = np.zeros_like(image_np, dtype=np.uint8)\n        for y, x in coords:\n            if 0 <= y < mask_np.shape[0] and 0 <= x < mask_np.shape[1]:\n                mask_np[y, x] = 1\n\n        # Convert back to PIL for transform\n        image = Image.fromarray(image_np)\n        mask = Image.fromarray(mask_np)\n\n        # Apply transforms\n        if self.image_transform:\n            image = self.image_transform(image)\n        if self.mask_transform:\n            mask = self.mask_transform(mask)\n\n        return image, mask\n\n\n# Transforms\nimage_transform = T.Compose([\n    T.Resize((512, 512)),\n    T.ToTensor(),\n    T.Normalize(mean=[0.485], std=[0.229])  # ImageNet normalization for grayscale\n])\n\nmask_transform = T.Compose([\n    T.Resize((512, 512)),\n    T.ToTensor(),\n])\n\n# Split by tomogram ID\nall_tomos = list(set([os.path.basename(os.path.dirname(p)) for p, _, _ in slice_dataset]))\ntrain_val_tomos, test_tomos = train_test_split(all_tomos, test_size=0.15, random_state=42)\ntrain_tomos, val_tomos = train_test_split(train_val_tomos, test_size=0.2, random_state=42)\n\n# Filter dataset\ndef filter_by_tomos(dataset, tomo_list):\n    return [item for item in dataset if os.path.basename(os.path.dirname(item[0])) in tomo_list]\n\ntrain_list = filter_by_tomos(slice_dataset, train_tomos)\nval_list = filter_by_tomos(slice_dataset, val_tomos)\ntest_list = filter_by_tomos(slice_dataset, test_tomos)\n\n# Build PyTorch datasets\ntrain_dataset = SliceSegmentationDataset(train_list, image_transform, mask_transform)\nval_dataset = SliceSegmentationDataset(val_list, image_transform, mask_transform)\ntest_dataset = SliceSegmentationDataset(test_list, image_transform, mask_transform)\n\n# DataLoaders\ntrain_loader = DataLoader(train_dataset, batch_size=16, shuffle=True)\nval_loader = DataLoader(val_dataset, batch_size=16, shuffle=False)\ntest_loader = DataLoader(test_dataset, batch_size=16, shuffle=False)\n\nprint(f\"Train samples: {len(train_dataset)} | Val: {len(val_dataset)} | Test: {len(test_dataset)}\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 🧩  Step 1.5 – Training Loop with Early Stopping, Loss Logging and Time Reporting\n\n\n","metadata":{}},{"cell_type":"code","source":"# Loss and optimizer\ncriterion = nn.BCEWithLogitsLoss()\noptimizer = torch.optim.Adam(model.parameters(), lr=1e-3)\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, patience=5, factor=0.5)\n\n# Training loop\nEPOCHS = 50\npatience = 8\ntrain_losses, val_losses = [], []\nbest_val_loss = float('inf')\nepochs_no_improve = 0\n\nprint(f\"🔁 Starting training for {EPOCHS} epochs...\")\nprint(f\"📊 Train samples: {len(train_dataset)} | Val: {len(val_dataset)} | Batches per epoch: {len(train_loader)}\")\n\nfor epoch in range(EPOCHS):\n    start_time = time.time()\n    model.train()\n    train_loss = 0\n\n    for batch_idx, (images, masks) in enumerate(train_loader):\n        images, masks = images.to(device), masks.to(device)\n        optimizer.zero_grad()\n        outputs = model(images)\n        loss = criterion(outputs, masks)\n        loss.backward()\n        optimizer.step()\n        train_loss += loss.item()\n\n        # Print progress only in epoch 1\n        if epoch == 0 and batch_idx % 5 == 0:\n            print(f\"🌀 Epoch 1 | Batch {batch_idx+1}/{len(train_loader)} | Batch Loss: {loss.item():.4f}\")\n\n    train_loss /= len(train_loader)\n    train_losses.append(train_loss)\n\n    model.eval()\n    val_loss = 0\n    with torch.no_grad():\n        for images, masks in val_loader:\n            images, masks = images.to(device), masks.to(device)\n            outputs = model(images)\n            loss = criterion(outputs, masks)\n            val_loss += loss.item()\n    val_loss /= len(val_loader)\n    val_losses.append(val_loss)\n\n    scheduler.step(val_loss)\n    epoch_time = time.time() - start_time\n    lr = optimizer.param_groups[0]['lr']\n\n    print(f\"📘 Epoch {epoch+1}/{EPOCHS} | 🏋️‍♂️ Train Loss: {train_loss:.4f} | 🧪 Val Loss: {val_loss:.4f} | ⏱️ Time: {epoch_time:.2f}s | LR: {lr:.6f}\")\n\n    if val_loss < best_val_loss:\n        best_val_loss = val_loss\n        epochs_no_improve = 0\n        torch.save(model.state_dict(), \"best_model.pth\")  # Save best model\n    else:\n        epochs_no_improve += 1\n        if epochs_no_improve >= patience:\n            print(\"⛔ Early stopping triggered.\")\n            break\n\n# Plot training and validation losses\nplt.figure(figsize=(10,5))\nplt.plot(train_losses, label=\"Train Loss\")\nplt.plot(val_losses, label=\"Val Loss\")\nplt.title(\"Training and Validation Losses\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Loss\")\nplt.legend()\nplt.grid(True)\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 🟢 Step 1.6 – Evaluation on Internal Test Set\n\n","metadata":{}},{"cell_type":"code","source":"from sklearn.metrics import jaccard_score, f1_score\n\n# Load best model\nmodel.load_state_dict(torch.load(\"best_model.pth\"))\nmodel.eval()\n\ntest_loss = 0\nious, dices = [], []\n\nwith torch.no_grad():\n    for images, masks in test_loader:\n        images, masks = images.to(device), masks.to(device)\n        outputs = model(images)\n\n        loss = criterion(outputs, masks)\n        test_loss += loss.item()\n\n        preds = torch.sigmoid(outputs) > 0.5  # Convert logits to binary\n        preds = preds.cpu().numpy().astype(np.uint8).flatten()\n        targets = masks.cpu().numpy().astype(np.uint8).flatten()\n\n        iou = jaccard_score(targets, preds, zero_division=0)\n        dice = f1_score(targets, preds, zero_division=0)\n\n        ious.append(iou)\n        dices.append(dice)\n\ntest_loss /= len(test_loader)\nmean_iou = np.mean(ious)\nmean_dice = np.mean(dices)\n\nprint(f\"📊 Test Loss: {test_loss:.4f}\")\nprint(f\"📐 Mean IoU:  {mean_iou:.4f}\")\nprint(f\"🎯 Mean Dice: {mean_dice:.4f}\")\n\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\nmodel.eval()\nnum_samples = 5  # כמה דוגמאות להציג\n\nwith torch.no_grad():\n    for i, (images, masks) in enumerate(test_loader):\n        images = images.to(device)\n        outputs = model(images)\n        preds = torch.sigmoid(outputs) > 0.5  # לוגיסט → בינארי\n        break  # מציג רק את הבאטץ' הראשון\n\n# הצגה\nimages = images.cpu()\nmasks = masks.cpu()\npreds = preds.cpu()\n\nfor idx in range(min(num_samples, images.size(0))):\n    fig, axs = plt.subplots(1, 3, figsize=(12, 4))\n    axs[0].imshow(images[idx][0], cmap=\"gray\")\n    axs[0].set_title(\"Input Image\")\n    \n    axs[1].imshow(masks[idx][0], cmap=\"gray\")\n    axs[1].set_title(\"Ground Truth Mask\")\n\n    axs[2].imshow(preds[idx][0], cmap=\"gray\")\n    axs[2].set_title(\"Predicted Mask\")\n\n    for ax in axs:\n        ax.axis(\"off\")\n    plt.tight_layout()\n    plt.show()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# PART 2 – Connected Components\n","metadata":{}},{"cell_type":"markdown","source":"## Step 2.1 – Threshold the Prediction Volume\n","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport scipy.ndimage as ndi\n\n# 🔶 Step 2.1 – Threshold the 3D prediction volume (after UNet2D inference on slices)\nthreshold = 0.5\nbinary_mask = (pred_volume >= threshold).astype(np.uint8)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Step 2.2 – Label Connected Components\n","metadata":{}},{"cell_type":"code","source":"# This finds all connected blobs (using 26-connectivity by default)\nlabeled_mask, num_features = ndi.label(binary_mask)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 🔶 Step 2.3 – Compute Centroids of the blobs\n","metadata":{}},{"cell_type":"code","source":"# This returns list of (z, y, x) float coordinates\ncentroids = ndi.center_of_mass(binary_mask, labeled_mask, range(1, num_features + 1))\n\n# 🔶 Optional: round and display first few\ncentroid_list = [tuple(np.round(c, 1)) for c in centroids]\nprint(f\"📌 Found {len(centroid_list)} candidate blobs\")\nprint(\"Sample centroids:\", centroid_list[:5])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 🟪 Step 2.4 – Visualize Candidate Blobs on Slices\n","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\n# 🔷 בחר slice להצגה\nz_idx = 32  # שווה לשנות ולבחון אחרים לפי הצורך\n\n# 🔷 תצוגת ה־heatmap עם מועמדים\nplt.figure(figsize=(8, 6))\nplt.imshow(pred_volume[z_idx], cmap='hot', interpolation='nearest')\nplt.title(f\"Heatmap with candidate centroids – Slice z={z_idx}\")\nplt.xlabel(\"X axis\")\nplt.ylabel(\"Y axis\")\n\n# 🔷 סימון Centroids שנמצאים בסלייס הזה (z קרוב ל־z_idx)\nfor (z, y, x) in centroids:\n    if int(round(z)) == z_idx:\n        plt.scatter(x, y, c='cyan', s=30, edgecolors='black', label='Centroid')\n\nplt.colorbar(label='Predicted probability')\nplt.grid(False)\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# PART 3 – Refinement of Candidate Predictions\n","metadata":{}},{"cell_type":"markdown","source":"## Step 3.1 – Extract 3D Crops Around Centroids\n","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn.functional as F\nimport numpy as np\nimport scipy.ndimage as ndi\n\n# Crop parameters\ncrop_depth = 16\ncrop_height = 128\ncrop_width = 128\n\nrefined_centroids = []\n\nfor (z, y, x) in centroids:\n    z, y, x = int(round(z)), int(round(y)), int(round(x))\n\n    # ───────────────────────────────\n    # 🔹 Step 3.1 – Extract 3D Crop\n    # ───────────────────────────────\n    z_start = max(z - crop_depth // 2, 0)\n    z_end = min(z_start + crop_depth, pred_volume.shape[0])\n    y_start = max(y - crop_height // 2, 0)\n    y_end = min(y_start + crop_height, pred_volume.shape[1])\n    x_start = max(x - crop_width // 2, 0)\n    x_end = min(x_start + crop_width, pred_volume.shape[2])\n\n    # Adjust in case crop is too close to border\n    z_start = z_end - crop_depth if z_end - z_start < crop_depth else z_start\n    y_start = y_end - crop_height if y_end - y_start < crop_height else y_start\n    x_start = x_end - crop_width if x_end - x_start < crop_width else x_start\n\n    crop = pred_volume[z_start:z_end, y_start:y_end, x_start:x_end]","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 🔹 Step 3.2 – Rerun UNet2D on Crop Slices\n","metadata":{}},{"cell_type":"code","source":" refined_crop = []\n    for i in range(crop.shape[0]):\n        slice_tensor = torch.tensor(crop[i]).float().unsqueeze(0).unsqueeze(0).to(device)  # [1, 1, H, W]\n        with torch.no_grad():\n            pred = model(slice_tensor)  # [1, 1, H, W]\n            pred_slice = torch.sigmoid(pred).squeeze().cpu().numpy()\n        refined_crop.append(pred_slice)\n\n    refined_crop = np.stack(refined_crop)  # shape: (D, H, W)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 🔹 Step 3.3 – Compute Refined Centroid","metadata":{}},{"cell_type":"code","source":" binary_refined = (refined_crop >= 0.4).astype(np.uint8)\n    labeled_refined, num = ndi.label(binary_refined)\n    if num == 0:\n        continue  # Skip if no connected components found\n\n    # Select the largest component\n    sizes = ndi.sum(binary_refined, labeled_refined, range(1, num + 1))\n    largest_label = np.argmax(sizes) + 1\n    mask_largest = (labeled_refined == largest_label)\n\n    # Local centroid in crop space\n    local_centroid = ndi.center_of_mass(mask_largest)\n\n    # Convert to global coordinates\n    global_centroid = (\n        z_start + local_centroid[0],\n        y_start + local_centroid[1],\n        x_start + local_centroid[2],\n    )\n    refined_centroids.append(tuple(np.round(global_centroid, 1)))\n\n# ✅ Final output: refined_centroids\nprint(f\"✅ Total refined centroids: {len(refined_centroids)}\")\nprint(\"Sample refined centroids:\", refined_centroids[:5])","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}