{"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"},{"sourceId":226864880,"sourceType":"kernelVersion"}],"dockerImageVersionId":31041,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!tar xfvz /kaggle/input/ultralytics-for-offline-install/archive.tar.gz\n!pip install --no-index --find-links=./packages ultralytics\n!rm -rf ./packages","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T06:32:47.756063Z","iopub.execute_input":"2025-05-21T06:32:47.756273Z","iopub.status.idle":"2025-05-21T06:34:29.622310Z","shell.execute_reply.started":"2025-05-21T06:32:47.756255Z","shell.execute_reply":"2025-05-21T06:34:29.621251Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install plotly","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T06:34:54.826936Z","iopub.execute_input":"2025-05-21T06:34:54.827758Z","iopub.status.idle":"2025-05-21T06:34:57.921803Z","shell.execute_reply.started":"2025-05-21T06:34:54.827717Z","shell.execute_reply":"2025-05-21T06:34:57.921139Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import plotly.express as px\nfrom PIL import Image, ImageDraw\nimport random\nimport seaborn as sns\nfrom matplotlib.patches import Rectangle\nfrom ultralytics import YOLO\nimport yaml\nimport json\nimport os\nimport glob\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom tqdm.notebook import tqdm\nfrom sklearn.model_selection import train_test_split\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\nimport cv2\nimport threading\nimport time\nfrom contextlib import nullcontext\nfrom concurrent.futures import ThreadPoolExecutor\nimport math","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T06:35:01.303893Z","iopub.execute_input":"2025-05-21T06:35:01.304238Z","iopub.status.idle":"2025-05-21T06:35:09.410364Z","shell.execute_reply.started":"2025-05-21T06:35:01.304208Z","shell.execute_reply":"2025-05-21T06:35:09.409789Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define global constants for dataset directories\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 if they don't exist\nos.makedirs(OUTPUT_DIR, exist_ok=True)\nos.makedirs(MODEL_DIR, exist_ok=True)\n\n# Set device: Use GPU if available; otherwise, fall back to CPU\nDEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {DEVICE}\")\n\n# Set random 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-05-21T06:35:13.791664Z","iopub.execute_input":"2025-05-21T06:35:13.792192Z","iopub.status.idle":"2025-05-21T06:35:13.883992Z","shell.execute_reply.started":"2025-05-21T06:35:13.792167Z","shell.execute_reply":"2025-05-21T06:35:13.883347Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load the training labels CSV into a pandas DataFrame\ntrain_labels = pd.read_csv(TRAIN_CSV)\n\n# Display basic dataset information\nprint(\"Training dataset shape:\", train_labels.shape)\nprint(\"\\nColumns in the dataset:\")\nprint(train_labels.columns.tolist())\n\n# Display basic statistics for numerical columns\nprint(\"\\nBasic statistics:\")\ndisplay(train_labels.describe())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T06:35:21.532787Z","iopub.execute_input":"2025-05-21T06:35:21.533113Z","iopub.status.idle":"2025-05-21T06:35:21.586688Z","shell.execute_reply.started":"2025-05-21T06:35:21.533067Z","shell.execute_reply":"2025-05-21T06:35:21.586007Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Count unique tomograms in the dataset\nunique_tomo_count = train_labels['tomo_id'].nunique()\nprint(f\"\\nNumber of unique tomograms: {unique_tomo_count}\")\n\n# Compute distribution of motors per tomogram\nmotors_per_tomo = train_labels.groupby('tomo_id')['Number of motors'].first().value_counts().sort_index()\nprint(\"\\nDistribution of motors per tomogram:\")\nprint(motors_per_tomo)\n\n# Visualize the distribution with a bar plot\nplt.figure(figsize=(8, 5))\nmotors_per_tomo.plot(kind='bar', color='skyblue', edgecolor='black')\nplt.title('Distribution of Motors per Tomogram')\nplt.xlabel('Number of Motors')\nplt.ylabel('Frequency')\nplt.xticks(rotation=0)\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T06:35:26.138077Z","iopub.execute_input":"2025-05-21T06:35:26.138666Z","iopub.status.idle":"2025-05-21T06:35:26.560433Z","shell.execute_reply.started":"2025-05-21T06:35:26.138644Z","shell.execute_reply":"2025-05-21T06:35:26.559723Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Display a few sample rows from the training labels\nprint(\"\\nSample rows from training labels:\")\ndisplay(train_labels.head())\n\n# Check for missing values in each column\nprint(\"\\nMissing values per column:\")\ndisplay(train_labels.isnull().sum())\n\n# Explore the range of tomogram sizes along each axis\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# Display voxel spacing distribution\nprint(\"\\nVoxel spacing distribution:\")\nvoxel_spacing_counts = train_labels['Voxel spacing'].value_counts().sort_index()\ndisplay(voxel_spacing_counts)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T06:35:31.667825Z","iopub.execute_input":"2025-05-21T06:35:31.668129Z","iopub.status.idle":"2025-05-21T06:35:31.688931Z","shell.execute_reply.started":"2025-05-21T06:35:31.668105Z","shell.execute_reply":"2025-05-21T06:35:31.688258Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"fig_motor = px.scatter_3d(\n    train_labels, \n    x='Motor axis 0', \n    y='Motor axis 1', \n    z='Motor axis 2',\n    color='Number of motors', \n    color_continuous_scale=\"viridis\",  # Using a vibrant color scheme\n    size_max=8, \n    width=900, \n    height=600, \n    opacity=0.85, \n    template=\"plotly_white\",  # Lighter theme for better contrast\n    title=\"🚀 3D Scatter Plot: Motor Axes\"\n)\n\nfig_motor.update_layout(\n    font_size=10,\n    legend_font_size=14,\n    margin=dict(l=10, r=10, b=10, t=40)\n)\n\nfig_motor.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T06:35:36.927622Z","iopub.execute_input":"2025-05-21T06:35:36.927902Z","iopub.status.idle":"2025-05-21T06:35:38.357938Z","shell.execute_reply.started":"2025-05-21T06:35:36.927881Z","shell.execute_reply":"2025-05-21T06:35:38.357278Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"fig_shape = px.scatter_3d(\n    train_labels, \n    x='Array shape (axis 0)', \n    y='Array shape (axis 1)', \n    z='Array shape (axis 2)',\n    color='Number of motors', \n    color_continuous_scale=\"magma\",  # More contrast for clarity\n    size_max=8, \n    width=900, \n    height=600, \n    opacity=0.85, \n    template=\"seaborn\",  # New theme for a scientific feel\n    title=\"🧬 3D Scatter Plot: Tomogram Shapes\"\n)\n\nfig_shape.update_layout(\n    font_size=10,\n    legend_font_size=14,\n    margin=dict(l=10, r=10, b=10, t=40)\n)\n\nfig_shape.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T06:35:45.258757Z","iopub.execute_input":"2025-05-21T06:35:45.259360Z","iopub.status.idle":"2025-05-21T06:35:45.380882Z","shell.execute_reply.started":"2025-05-21T06:35:45.259336Z","shell.execute_reply":"2025-05-21T06:35:45.379944Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Show descriptive statistics\ndisplay(train_labels.describe().loc[['mean', 'min', 'max']].T)\n\n# Improved histogram design\ntrain_labels.hist(\n    bins=30, \n    figsize=(14, 8), \n    layout=(3, 4), \n    edgecolor=\"black\", \n    color=\"#4CAF50\"  # Greenish color theme\n)\nplt.suptitle(\"Feature Distributions\", fontsize=16, fontweight='bold', color=\"darkblue\")\nplt.tight_layout()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T06:35:59.443355Z","iopub.execute_input":"2025-05-21T06:35:59.443654Z","iopub.status.idle":"2025-05-21T06:36:01.055394Z","shell.execute_reply.started":"2025-05-21T06:35:59.443633Z","shell.execute_reply":"2025-05-21T06:36:01.054712Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.figure(figsize=(9, 5), facecolor=\"white\")\nsns.heatmap(\n    data=train_labels.corr(numeric_only=True),\n    cmap=\"coolwarm\",  # Strong contrast for positive/negative correlations\n    vmin=-1, vmax=1,\n    linecolor=\"white\", linewidth=0.6,\n    annot=True,\n    fmt=\".2f\"\n)\nplt.title('Correlation Heatmap', fontsize=14, fontweight='bold', color=\"black\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T06:36:08.537960Z","iopub.execute_input":"2025-05-21T06:36:08.538670Z","iopub.status.idle":"2025-05-21T06:36:08.950646Z","shell.execute_reply.started":"2025-05-21T06:36:08.538644Z","shell.execute_reply":"2025-05-21T06:36:08.949906Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plotImages(title, directory, n_images=16, img_size=(128, 128)):\n    \"\"\"\n    Display a grid of images from the specified directory.\n    \n    Args:\n        title (str): Title to print before displaying images.\n        directory (str): Glob pattern for image files.\n        n_images (int): Number of images to display.\n        img_size (tuple): Size to resize images for display.\n    \"\"\"\n    print(f\"🖼 {title}\")\n    image_files = glob.glob(directory)\n    \n    if not image_files:\n        print(\"No images found.\")\n        return\n    \n    plt.figure(figsize=(12, 12))\n    plt.subplots_adjust(wspace=0.1, hspace=0.1)\n    \n    for i, file_path in enumerate(image_files[:n_images]):\n        img = cv2.imread(file_path)\n        if img is None:\n            continue\n        img = cv2.resize(img, img_size)\n        plt.subplot(4, 4, i+1)\n        plt.imshow(cv2.cvtColor(img, cv2.COLOR_BGR2RGB))\n        plt.axis('off')\n    \n    plt.suptitle(title, fontsize=14, fontweight='bold', color=\"darkred\")\n    plt.show()\n\nplotImages(\"Bacterial Flagellar Motors - Train Images\", \"../input/byu-locating-bacterial-flagellar-motors-2025/train/***/**\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T06:36:14.598670Z","iopub.execute_input":"2025-05-21T06:36:14.599161Z","iopub.status.idle":"2025-05-21T06:36:23.359670Z","shell.execute_reply.started":"2025-05-21T06:36:14.599137Z","shell.execute_reply":"2025-05-21T06:36:23.358729Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_images(path, n_images=12, is_random=True, figsize=(14, 14)):\n      \n    plt.figure(figsize=figsize)\n    \n    image_names = os.listdir(path)\n    if is_random:\n        image_names = random.sample(image_names, min(len(image_names), n_images))\n    else:\n        image_names = image_names[:n_images]\n    \n    w = int(math.sqrt(n_images))\n    h = math.ceil(n_images / w)\n    \n    for ind, image_name in enumerate(image_names):\n        img_path = os.path.join(path, image_name)\n        img = cv2.imread(img_path)\n        if img is None:\n            continue\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        plt.subplot(h, w, ind + 1)\n        plt.imshow(img)\n        plt.xticks([])\n        plt.yticks([])\n    \n    plt.suptitle(\"Sample Tomogram Images\", fontsize=14, fontweight='bold', color=\"darkblue\")\n    plt.show()\n\nvisualize_images(\"/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/train/tomo_098751\", n_images=9)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T06:36:36.555239Z","iopub.execute_input":"2025-05-21T06:36:36.556137Z","iopub.status.idle":"2025-05-21T06:36:38.087688Z","shell.execute_reply.started":"2025-05-21T06:36:36.556101Z","shell.execute_reply":"2025-05-21T06:36:38.086743Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Select a sample tomogram ID to visualize\nsample_tomo_id = train_labels['tomo_id'].iloc[0]\nprint(f\"\\nVisualizing sample tomogram: {sample_tomo_id}\")\n\n# Construct the folder path for the selected tomogram\nsample_folder = os.path.join(TRAIN_DIR, sample_tomo_id)\n\nif os.path.exists(sample_folder):\n    # Get all JPEG slice files from the tomogram folder\n    slice_files = sorted(glob.glob(os.path.join(sample_folder, '*.jpg')))\n    print(f\"Number of slice files in tomogram '{sample_tomo_id}': {len(slice_files)}\")\n    \n    if slice_files:\n        # Load the first slice to check its dimensions\n        sample_slice = Image.open(slice_files[0])\n        print(f\"Dimensions of a sample slice: {sample_slice.size}\")\n        \n        # Plot slices from the beginning, middle, and end of the tomogram\n        fig, axes = plt.subplots(1, 3, figsize=(15, 5))\n        slice_indices = [0, len(slice_files)//2, len(slice_files)-1]\n        for i, idx in enumerate(slice_indices):\n            img = Image.open(slice_files[idx])\n            axes[i].imshow(img, cmap='gray')\n            axes[i].set_title(f\"Slice {idx}\")\n            axes[i].axis('off')\n        plt.tight_layout()\n        plt.show()\n    else:\n        print(\"No slice files found in the folder.\")\nelse:\n    print(f\"Folder '{sample_folder}' does not exist. Please check the dataset directory.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T06:36:47.737919Z","iopub.execute_input":"2025-05-21T06:36:47.738234Z","iopub.status.idle":"2025-05-21T06:36:48.834345Z","shell.execute_reply.started":"2025-05-21T06:36:47.738212Z","shell.execute_reply":"2025-05-21T06:36:48.833523Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install ipywidgets\n!pip install progress\nfrom tqdm import tqdm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T06:36:55.664203Z","iopub.execute_input":"2025-05-21T06:36:55.664472Z","iopub.status.idle":"2025-05-21T06:40:19.651874Z","shell.execute_reply.started":"2025-05-21T06:36:55.664452Z","shell.execute_reply":"2025-05-21T06:40:19.651128Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define YOLO dataset structure and parameters\ndata_path = \"/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/\"\ntrain_dir = os.path.join(data_path, \"train\")\n\n# Output directories for YOLO dataset (adjust as needed)\nyolo_dataset_dir = \"/kaggle/working/yolo_dataset\"\nyolo_images_train = os.path.join(yolo_dataset_dir, \"images\", \"train\")\nyolo_images_val = os.path.join(yolo_dataset_dir, \"images\", \"val\")\nyolo_labels_train = os.path.join(yolo_dataset_dir, \"labels\", \"train\")\nyolo_labels_val = os.path.join(yolo_dataset_dir, \"labels\", \"val\")\n\n# Create necessary directories\nfor dir_path in [yolo_images_train, yolo_images_val, yolo_labels_train, yolo_labels_val]:\n    os.makedirs(dir_path, exist_ok=True)\n\n# Define constants for processing\nTRUST = 4       # Number of slices above and below center slice (total slices = 2*TRUST + 1)\nBOX_SIZE = 24   # Bounding box size (in pixels)\nTRAIN_SPLIT = 0.8  # 80% training, 20% validation\n\n# Define a helper function for image normalization using percentile-based contrast enhancement.\ndef normalize_slice(slice_data):\n    p2 = np.percentile(slice_data, 2)\n    p98 = np.percentile(slice_data, 98)\n    clipped_data = np.clip(slice_data, p2, p98)\n    normalized = 255 * (clipped_data - p2) / (p98 - p2)\n    return np.uint8(normalized)\n\n# Define the preprocessing function to extract slices, normalize, and generate YOLO annotations.\ndef prepare_yolo_dataset(trust=TRUST, train_split=TRAIN_SPLIT):\n\n    # Load the labels CSV\n    labels_df = pd.read_csv(os.path.join(data_path, \"train_labels.csv\"))\n    \n    total_motors = labels_df['Number of motors'].sum()\n    print(f\"Total number of motors in the dataset: {total_motors}\")\n    \n    # Consider only tomograms with at least one motor\n    tomo_df = labels_df[labels_df['Number of motors'] > 0].copy()\n    unique_tomos = tomo_df['tomo_id'].unique()\n    print(f\"Found {len(unique_tomos)} unique tomograms with motors\")\n    \n    # Shuffle and split tomograms into train and validation sets\n    np.random.shuffle(unique_tomos)\n    split_idx = int(len(unique_tomos) * train_split)\n    train_tomos = unique_tomos[:split_idx]\n    val_tomos = unique_tomos[split_idx:]\n    print(f\"Split: {len(train_tomos)} tomograms for training, {len(val_tomos)} tomograms for validation\")\n    \n    # Helper function to process a list of tomograms\n    def process_tomogram_set(tomogram_ids, images_dir, labels_dir, set_name):\n        motor_counts = []\n        for tomo_id in tomogram_ids:\n            # Get motor annotations for the current tomogram\n            tomo_motors = labels_df[labels_df['tomo_id'] == tomo_id]\n            for _, motor in tomo_motors.iterrows():\n                if pd.isna(motor['Motor axis 0']):\n                    continue\n                motor_counts.append(\n                    (tomo_id, \n                     int(motor['Motor axis 0']), \n                     int(motor['Motor axis 1']), \n                     int(motor['Motor axis 2']),\n                     int(motor['Array shape (axis 0)']))\n                )\n        \n        print(f\"Will process approximately {len(motor_counts) * (2 * trust + 1)} slices for {set_name}\")\n        processed_slices = 0\n        \n        # Loop over each motor annotation\n        for tomo_id, z_center, y_center, x_center, z_max in tqdm(motor_counts, desc=f\"Processing {set_name} motors\"):\n            z_min = max(0, z_center - trust)\n            z_max_bound = min(z_max - 1, z_center + trust)\n            for z in range(z_min, z_max_bound + 1):\n                # Create the slice filename and source path\n                slice_filename = f\"slice_{z:04d}.jpg\"\n                src_path = os.path.join(train_dir, tomo_id, slice_filename)\n                if not os.path.exists(src_path):\n                    print(f\"Warning: {src_path} does not exist, skipping.\")\n                    continue\n                \n                # Load, normalize, and save the image slice\n                img = Image.open(src_path)\n                img_array = np.array(img)\n                normalized_img = normalize_slice(img_array)\n                dest_filename = f\"{tomo_id}_z{z:04d}_y{y_center:04d}_x{x_center:04d}.jpg\"\n                dest_path = os.path.join(images_dir, dest_filename)\n                Image.fromarray(normalized_img).save(dest_path)\n                \n                # Prepare YOLO bounding box annotation (normalized values)\n                img_width, img_height = img.size\n                x_center_norm = x_center / img_width\n                y_center_norm = y_center / img_height\n                box_width_norm = BOX_SIZE / img_width\n                box_height_norm = BOX_SIZE / img_height\n                label_path = os.path.join(labels_dir, dest_filename.replace('.jpg', '.txt'))\n                with open(label_path, 'w') as f:\n                    f.write(f\"0 {x_center_norm} {y_center_norm} {box_width_norm} {box_height_norm}\\n\")\n                \n                processed_slices += 1\n        \n        return processed_slices, len(motor_counts)\n    \n    # Process training tomograms\n    train_slices, train_motors = process_tomogram_set(train_tomos, yolo_images_train, yolo_labels_train, \"training\")\n    # Process validation tomograms\n    val_slices, val_motors = process_tomogram_set(val_tomos, yolo_images_val, yolo_labels_val, \"validation\")\n    \n    # Generate YAML configuration for YOLO training\n    yaml_content = {\n        'path': yolo_dataset_dir,\n        'train': 'images/train',\n        'val': 'images/val',\n        'names': {0: 'motor'}\n    }\n    with open(os.path.join(yolo_dataset_dir, 'dataset.yaml'), 'w') as f:\n        yaml.dump(yaml_content, f, default_flow_style=False)\n    \n    print(f\"\\nProcessing Summary:\")\n    print(f\"- Train set: {len(train_tomos)} tomograms, {train_motors} motors, {train_slices} slices\")\n    print(f\"- Validation set: {len(val_tomos)} tomograms, {val_motors} motors, {val_slices} slices\")\n    print(f\"- Total: {len(train_tomos) + len(val_tomos)} tomograms, {train_motors + val_motors} motors, {train_slices + val_slices} slices\")\n    \n    return {\n        \"dataset_dir\": yolo_dataset_dir,\n        \"yaml_path\": os.path.join(yolo_dataset_dir, 'dataset.yaml'),\n        \"train_tomograms\": len(train_tomos),\n        \"val_tomograms\": len(val_tomos),\n        \"train_motors\": train_motors,\n        \"val_motors\": val_motors,\n        \"train_slices\": train_slices,\n        \"val_slices\": val_slices\n    }\n\n# Run the preprocessing\nsummary = prepare_yolo_dataset(TRUST)\nprint(f\"\\nPreprocessing Complete:\")\nprint(f\"- Training data: {summary['train_tomograms']} tomograms, {summary['train_motors']} motors, {summary['train_slices']} slices\")\nprint(f\"- Validation data: {summary['val_tomograms']} tomograms, {summary['val_motors']} motors, {summary['val_slices']} slices\")\nprint(f\"- Dataset directory: {summary['dataset_dir']}\")\nprint(f\"- YAML configuration: {summary['yaml_path']}\")\nprint(\"\\nReady for YOLO training!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T06:41:08.644594Z","iopub.execute_input":"2025-05-21T06:41:08.644923Z","iopub.status.idle":"2025-05-21T06:43:58.594864Z","shell.execute_reply.started":"2025-05-21T06:41:08.644901Z","shell.execute_reply":"2025-05-21T06:43:58.594124Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Set paths for the preprocessed YOLO training images and labels.\nimages_train_dir = os.path.join(yolo_dataset_dir, \"images\", \"train\")\nlabels_train_dir = os.path.join(yolo_dataset_dir, \"labels\", \"train\")\n\n# %%\ndef visualize_random_training_samples(num_samples=4):\n    # Get all image files from the train directory (support multiple image extensions)\n    image_files = []\n    for ext in ['*.jpg', '*.jpeg', '*.png']:\n        image_files.extend(glob.glob(os.path.join(images_train_dir, \"**\", ext), recursive=True))\n    \n    if len(image_files) == 0:\n        print(\"No image files found in the train directory!\")\n        return\n        \n    num_samples = min(num_samples, len(image_files))\n    random_images = random.sample(image_files, num_samples)\n    \n    # Create subplots for visualization\n    rows = int(np.ceil(num_samples / 2))\n    cols = min(num_samples, 2)\n    fig, axes = plt.subplots(rows, cols, figsize=(14, 5 * rows))\n    \n    if num_samples == 1:\n        axes = np.array([axes])\n    axes = axes.flatten()\n    \n    for i, img_path in enumerate(random_images):\n        try:\n            # Determine corresponding label file (YOLO format)\n            relative_path = os.path.relpath(img_path, images_train_dir)\n            label_path = os.path.join(labels_train_dir, os.path.splitext(relative_path)[0] + '.txt')\n            \n            # Load and normalize image for display\n            img = Image.open(img_path)\n            img_width, img_height = img.size\n            img_array = np.array(img)\n            p2 = np.percentile(img_array, 2)\n            p98 = np.percentile(img_array, 98)\n            normalized = np.clip(img_array, p2, p98)\n            normalized = 255 * (normalized - p2) / (p98 - p2)\n            img_normalized = Image.fromarray(np.uint8(normalized))\n            \n            # Convert to RGB for annotation drawing\n            img_rgb = img_normalized.convert('RGB')\n            overlay = Image.new('RGBA', img_rgb.size, (0, 0, 0, 0))\n            draw = ImageDraw.Draw(overlay)\n            \n            # Load YOLO annotations if available\n            annotations = []\n            if os.path.exists(label_path):\n                with open(label_path, 'r') as f:\n                    for line in f:\n                        # YOLO format: class x_center y_center width height (normalized values)\n                        values = line.strip().split()\n                        class_id = int(values[0])\n                        x_center = float(values[1]) * img_width\n                        y_center = float(values[2]) * img_height\n                        width = float(values[3]) * img_width\n                        height = float(values[4]) * img_height\n                        annotations.append({\n                            'class_id': class_id,\n                            'x_center': x_center,\n                            'y_center': y_center,\n                            'width': width,\n                            'height': height\n                        })\n            \n            # Draw annotations on the overlay\n            for ann in annotations:\n                x_center = ann['x_center']\n                y_center = ann['y_center']\n                width = ann['width']\n                height = ann['height']\n                x1 = max(0, int(x_center - width/2))\n                y1 = max(0, int(y_center - height/2))\n                x2 = min(img_width, int(x_center + width/2))\n                y2 = min(img_height, int(y_center + height/2))\n                draw.rectangle([x1, y1, x2, y2], fill=(255, 0, 0, 64), outline=(255, 0, 0, 200))\n                draw.text((x1, y1-10), f\"Class {ann['class_id']}\", fill=(255, 0, 0, 255))\n            \n            # Indicate if no annotations were found\n            if not annotations:\n                draw.text((10, 10), \"No annotations found\", fill=(255, 0, 0, 255))\n            \n            # Composite overlay and display image\n            img_rgb = Image.alpha_composite(img_rgb.convert('RGBA'), overlay).convert('RGB')\n            axes[i].imshow(np.array(img_rgb))\n            img_name = os.path.basename(img_path)\n            axes[i].set_title(f\"Image: {img_name}\\nAnnotations: {len(annotations)}\")\n            axes[i].axis('on')\n            \n        except Exception as e:\n            print(f\"Error processing image {img_path}: {e}\")\n            axes[i].text(0.5, 0.5, f\"Error loading image: {os.path.basename(img_path)}\",\n                         horizontalalignment='center', verticalalignment='center')\n            axes[i].axis('off')\n    \n    # Turn off any extra subplots\n    for j in range(i + 1, len(axes)):\n        axes[j].axis('off')\n    \n    plt.tight_layout()\n    plt.show()\n    print(f\"Displayed {num_samples} random images with YOLO annotations\")\n\n# Run visualization of random training samples\nvisualize_random_training_samples(4)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T07:36:17.482506Z","iopub.execute_input":"2025-05-21T07:36:17.483249Z","iopub.status.idle":"2025-05-21T07:36:18.850918Z","shell.execute_reply.started":"2025-05-21T07:36:17.483228Z","shell.execute_reply":"2025-05-21T07:36:18.849991Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Set random seeds for reproducibility\nnp.random.seed(42)\nrandom.seed(42)\ntorch.manual_seed(42)\n\n# Define paths for the Kaggle environment\nyolo_dataset_dir = \"/kaggle/working/yolo_dataset\"\nyolo_weights_dir = \"/kaggle/working/yolo_weights\"\nyolo_pretrained_weights = \"/kaggle/input/ultralytics-for-offline-install/yolov8n.pt\"  # Pre-downloaded weights\n\n# Create the weights directory if it does not exist\nos.makedirs(yolo_weights_dir, exist_ok=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T07:36:32.303840Z","iopub.execute_input":"2025-05-21T07:36:32.304162Z","iopub.status.idle":"2025-05-21T07:36:32.310380Z","shell.execute_reply.started":"2025-05-21T07:36:32.304131Z","shell.execute_reply":"2025-05-21T07:36:32.309649Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def fix_yaml_paths(yaml_path):\n    print(f\"Fixing YAML paths in {yaml_path}\")\n    with open(yaml_path, 'r') as f:\n        yaml_data = yaml.safe_load(f)\n    \n    if 'path' in yaml_data:\n        yaml_data['path'] = yolo_dataset_dir\n    \n    fixed_yaml_path = \"/kaggle/working/fixed_dataset.yaml\"\n    with open(fixed_yaml_path, 'w') as f:\n        yaml.dump(yaml_data, f)\n    \n    print(f\"Created fixed YAML at {fixed_yaml_path} with path: {yaml_data.get('path')}\")\n    return fixed_yaml_path","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T07:36:36.278552Z","iopub.execute_input":"2025-05-21T07:36:36.278823Z","iopub.status.idle":"2025-05-21T07:36:36.283589Z","shell.execute_reply.started":"2025-05-21T07:36:36.278803Z","shell.execute_reply":"2025-05-21T07:36:36.282839Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_dfl_loss_curve(run_dir):\n    results_csv = os.path.join(run_dir, 'results.csv')\n    if not os.path.exists(results_csv):\n        print(f\"Results file not found at {results_csv}\")\n        return\n    \n    results_df = pd.read_csv(results_csv)\n    train_dfl_col = [col for col in results_df.columns if 'train/dfl_loss' in col]\n    val_dfl_col = [col for col in results_df.columns if 'val/dfl_loss' in col]\n    \n    if not train_dfl_col or not val_dfl_col:\n        print(\"DFL loss columns not found in results CSV\")\n        print(f\"Available columns: {results_df.columns.tolist()}\")\n        return\n    \n    train_dfl_col = train_dfl_col[0]\n    val_dfl_col = val_dfl_col[0]\n    \n    best_epoch = results_df[val_dfl_col].idxmin()\n    best_val_loss = results_df.loc[best_epoch, val_dfl_col]\n    \n    plt.figure(figsize=(10, 6))\n    plt.plot(results_df['epoch'], results_df[train_dfl_col], label='Train DFL Loss')\n    plt.plot(results_df['epoch'], results_df[val_dfl_col], label='Validation DFL Loss')\n    plt.axvline(x=results_df.loc[best_epoch, 'epoch'], color='r', linestyle='--', \n                label=f'Best Model (Epoch {int(results_df.loc[best_epoch, \"epoch\"])}, Val Loss: {best_val_loss:.4f})')\n    plt.xlabel('Epoch')\n    plt.ylabel('DFL Loss')\n    plt.title('Training and Validation DFL Loss')\n    plt.legend()\n    plt.grid(True, linestyle='--', alpha=0.7)\n    \n    plot_path = os.path.join(run_dir, 'dfl_loss_curve.png')\n    plt.savefig(plot_path)\n    plt.savefig(os.path.join('/kaggle/working', 'dfl_loss_curve.png'))\n    \n    print(f\"Loss curve saved to {plot_path}\")\n    plt.close()\n    \n    return best_epoch, best_val_loss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T07:36:43.036537Z","iopub.execute_input":"2025-05-21T07:36:43.037104Z","iopub.status.idle":"2025-05-21T07:36:43.044015Z","shell.execute_reply.started":"2025-05-21T07:36:43.037062Z","shell.execute_reply":"2025-05-21T07:36:43.043318Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_yolo_model(yaml_path, pretrained_weights_path, epochs=30, batch_size=16, img_size=640):\n    print(f\"Loading pre-trained weights from: {pretrained_weights_path}\")\n    model = YOLO(pretrained_weights_path)\n    \n    results = model.train(\n        data=yaml_path,\n        epochs=epochs,\n        batch=batch_size,\n        imgsz=img_size,\n        project=yolo_weights_dir,\n        name='motor_detector',\n        exist_ok=True,\n        patience=5,\n        save_period=5,\n        val=True,\n        verbose=True\n    )\n    \n    run_dir = os.path.join(yolo_weights_dir, 'motor_detector')\n    best_epoch_info = plot_dfl_loss_curve(run_dir)\n    if best_epoch_info:\n        best_epoch, best_val_loss = best_epoch_info\n        print(f\"\\nBest model found at epoch {best_epoch} with validation DFL loss: {best_val_loss:.4f}\")\n    \n    return model, results","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T07:36:52.095736Z","iopub.execute_input":"2025-05-21T07:36:52.096005Z","iopub.status.idle":"2025-05-21T07:36:52.101059Z","shell.execute_reply.started":"2025-05-21T07:36:52.095986Z","shell.execute_reply":"2025-05-21T07:36:52.100517Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def predict_on_samples(model, num_samples=4):\n    val_dir = os.path.join(yolo_dataset_dir, 'images', 'val')\n    if not os.path.exists(val_dir):\n        print(f\"Validation directory not found at {val_dir}\")\n        val_dir = os.path.join(yolo_dataset_dir, 'images', 'train')\n        print(f\"Using train directory for predictions instead: {val_dir}\")\n        \n    if not os.path.exists(val_dir):\n        print(\"No images directory found for predictions\")\n        return\n    \n    val_images = os.listdir(val_dir)\n    if len(val_images) == 0:\n        print(\"No images found for prediction\")\n        return\n    \n    num_samples = min(num_samples, len(val_images))\n    samples = random.sample(val_images, num_samples)\n    \n    fig, axes = plt.subplots(2, 2, figsize=(12, 12))\n    axes = axes.flatten()\n    \n    for i, img_file in enumerate(samples):\n        if i >= len(axes):\n            break\n            \n        img_path = os.path.join(val_dir, img_file)\n        results = model.predict(img_path, conf=0.25)[0]\n        img = Image.open(img_path)\n        axes[i].imshow(np.array(img), cmap='gray')\n        \n        # Draw ground truth box if available (extracted from filename)\n        try:\n            parts = img_file.split('_')\n            y_part = [p for p in parts if p.startswith('y')]\n            x_part = [p for p in parts if p.startswith('x')]\n            if y_part and x_part:\n                y_gt = int(y_part[0][1:])\n                x_gt = int(x_part[0][1:].split('.')[0])\n                box_size = 24\n                rect_gt = Rectangle((x_gt - box_size//2, y_gt - box_size//2), box_size, box_size,\n                                      linewidth=1, edgecolor='g', facecolor='none')\n                axes[i].add_patch(rect_gt)\n        except:\n            pass\n        \n        if len(results.boxes) > 0:\n            boxes = results.boxes.xyxy.cpu().numpy()\n            confs = results.boxes.conf.cpu().numpy()\n            for box, conf in zip(boxes, confs):\n                x1, y1, x2, y2 = box\n                rect_pred = Rectangle((x1, y1), x2-x1, y2-y1, linewidth=1, edgecolor='r', facecolor='none')\n                axes[i].add_patch(rect_pred)\n                axes[i].text(x1, y1-5, f'{conf:.2f}', color='red')\n        \n        axes[i].set_title(f\"Image: {img_file}\\nGT (green) vs Pred (red)\")\n    \n    plt.tight_layout()\n    plt.savefig(os.path.join('/kaggle/working', 'predictions.png'))\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T07:36:57.607307Z","iopub.execute_input":"2025-05-21T07:36:57.607556Z","iopub.status.idle":"2025-05-21T07:36:57.617820Z","shell.execute_reply.started":"2025-05-21T07:36:57.607541Z","shell.execute_reply":"2025-05-21T07:36:57.616958Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T07:37:06.524664Z","iopub.execute_input":"2025-05-21T07:37:06.525243Z","iopub.status.idle":"2025-05-21T07:37:06.528863Z","shell.execute_reply.started":"2025-05-21T07:37:06.525221Z","shell.execute_reply":"2025-05-21T07:37:06.528122Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def prepare_dataset():\n    train_images_dir = os.path.join(yolo_dataset_dir, 'images', 'train')\n    val_images_dir = os.path.join(yolo_dataset_dir, 'images', 'val')\n    train_labels_dir = os.path.join(yolo_dataset_dir, 'labels', 'train')\n    val_labels_dir = os.path.join(yolo_dataset_dir, 'labels', 'val')\n    \n    print(f\"Directory status:\")\n    print(f\"- Train images exists: {os.path.exists(train_images_dir)}\")\n    print(f\"- Val images exists: {os.path.exists(val_images_dir)}\")\n    print(f\"- Train labels exists: {os.path.exists(train_labels_dir)}\")\n    print(f\"- Val labels exists: {os.path.exists(val_labels_dir)}\")\n    \n    original_yaml_path = os.path.join(yolo_dataset_dir, 'dataset.yaml')\n    if os.path.exists(original_yaml_path):\n        print(f\"Found original dataset.yaml at {original_yaml_path}\")\n        return fix_yaml_paths(original_yaml_path)\n    else:\n        print(\"Original dataset.yaml not found, creating a new one\")\n        yaml_data = {\n            'path': yolo_dataset_dir,\n            'train': 'images/train',\n            'val': 'images/train' if not os.path.exists(val_images_dir) else 'images/val',\n            'names': {0: 'motor'}\n        }\n        new_yaml_path = \"/kaggle/working/dataset.yaml\"\n        with open(new_yaml_path, 'w') as f:\n            yaml.dump(yaml_data, f)\n        print(f\"Created new YAML at {new_yaml_path}\")\n        return new_yaml_path\n\ndef main():\n    print(\"Starting YOLO training process...\")\n    yaml_path = prepare_dataset()\n    print(f\"Using YAML file: {yaml_path}\")\n    with open(yaml_path, 'r') as f:\n        print(f\"YAML contents:\\n{f.read()}\")\n    \n    print(\"\\nStarting YOLO training...\")\n    model, results = train_yolo_model(\n        yaml_path,\n        pretrained_weights_path=yolo_pretrained_weights,\n        epochs=30  # For demonstration, using 30 epochs\n    )\n    \n    print(\"\\nTraining complete!\")\n    print(\"\\nRunning predictions on sample images...\")\n    predict_on_samples(model, num_samples=4)\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T07:37:10.963479Z","iopub.execute_input":"2025-05-21T07:37:10.963761Z","iopub.status.idle":"2025-05-21T07:50:40.014665Z","shell.execute_reply.started":"2025-05-21T07:37:10.963741Z","shell.execute_reply":"2025-05-21T07:50:40.013547Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Set random seed for reproducibility\nnp.random.seed(42)\ntorch.manual_seed(42)\n\n# Define paths for the test data and submission\ndata_path = \"/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/\"\ntest_dir = os.path.join(data_path, \"test\")\nsubmission_path = \"/kaggle/working/submission.csv\"\n\n# Path to the best trained model (adjust if necessary)\nmodel_path = \"/kaggle/working/yolo_weights/motor_detector/weights/best.pt\"\n\n# Define detection and processing parameters\nCONFIDENCE_THRESHOLD = 0.45\nMAX_DETECTIONS_PER_TOMO = 3\nNMS_IOU_THRESHOLD = 0.2\nCONCENTRATION = 1  # Process a fraction of slices for fast submission\n\n# GPU profiling context manager for timing\nclass GPUProfiler:\n    def __init__(self, name):\n        self.name = name\n        self.start_time = None\n        \n    def __enter__(self):\n        if torch.cuda.is_available():\n            torch.cuda.synchronize()\n        self.start_time = time.time()\n        return self\n        \n    def __exit__(self, *args):\n        if torch.cuda.is_available():\n            torch.cuda.synchronize()\n        elapsed = time.time() - self.start_time\n        print(f\"[PROFILE] {self.name}: {elapsed:.3f}s\")\n\n# Set device and dynamic batch size\ndevice = 'cuda:0' if torch.cuda.is_available() else 'cpu'\nBATCH_SIZE = 8\nif device.startswith('cuda'):\n    torch.backends.cudnn.benchmark = True\n    torch.backends.cudnn.deterministic = False\n    torch.backends.cuda.matmul.allow_tf32 = True\n    torch.backends.cudnn.allow_tf32 = True\n    gpu_name = torch.cuda.get_device_name(0)\n    gpu_mem = torch.cuda.get_device_properties(0).total_memory / 1e9\n    print(f\"Using GPU: {gpu_name} with {gpu_mem:.2f} GB memory\")\n    free_mem = gpu_mem - torch.cuda.memory_allocated(0) / 1e9\n    BATCH_SIZE = max(8, min(32, int(free_mem * 4)))\n    print(f\"Dynamic batch size set to {BATCH_SIZE} based on {free_mem:.2f}GB free memory\")\nelse:\n    print(\"GPU not available, using CPU\")\n    BATCH_SIZE = 4","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T07:51:03.867481Z","iopub.execute_input":"2025-05-21T07:51:03.867788Z","iopub.status.idle":"2025-05-21T07:51:03.879111Z","shell.execute_reply.started":"2025-05-21T07:51:03.867762Z","shell.execute_reply":"2025-05-21T07:51:03.878223Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def normalize_slice(slice_data):\n\n    p2 = np.percentile(slice_data, 2)\n    p98 = np.percentile(slice_data, 98)\n    clipped_data = np.clip(slice_data, p2, p98)\n    normalized = 255 * (clipped_data - p2) / (p98 - p2)\n    return np.uint8(normalized)\n\ndef preload_image_batch(file_paths):\n    images = []\n    for path in file_paths:\n        img = cv2.imread(path)\n        if img is None:\n            img = np.array(Image.open(path))\n        images.append(img)\n    return images\n\ndef perform_3d_nms(detections, iou_threshold):\n    if not detections:\n        return []\n    \n    detections = sorted(detections, key=lambda x: x['confidence'], reverse=True)\n    final_detections = []\n    def distance_3d(d1, d2):\n        return np.sqrt((d1['z'] - d2['z'])**2 + (d1['y'] - d2['y'])**2 + (d1['x'] - d2['x'])**2)\n    \n    box_size = 24\n    distance_threshold = box_size * iou_threshold\n    \n    while detections:\n        best_detection = detections.pop(0)\n        final_detections.append(best_detection)\n        detections = [d for d in detections if distance_3d(d, best_detection) > distance_threshold]\n    \n    return final_detections\n\ndef process_tomogram(tomo_id, model, index=0, total=1):\n    print(f\"Processing tomogram {tomo_id} ({index}/{total})\")\n    tomo_dir = os.path.join(test_dir, tomo_id)\n    slice_files = sorted([f for f in os.listdir(tomo_dir) if f.endswith('.jpg')])\n    \n    selected_indices = np.linspace(0, len(slice_files)-1, int(len(slice_files) * CONCENTRATION))\n    selected_indices = np.round(selected_indices).astype(int)\n    slice_files = [slice_files[i] for i in selected_indices]\n    \n    print(f\"Processing {len(slice_files)} out of {len(os.listdir(tomo_dir))} slices (CONCENTRATION={CONCENTRATION})\")\n    all_detections = []\n    \n    if device.startswith('cuda'):\n        streams = [torch.cuda.Stream() for _ in range(min(4, BATCH_SIZE))]\n    else:\n        streams = [None]\n    \n    next_batch_thread = None\n    next_batch_images = None\n    \n    for batch_start in range(0, len(slice_files), BATCH_SIZE):\n        if next_batch_thread is not None:\n            next_batch_thread.join()\n            next_batch_images = None\n            \n        batch_end = min(batch_start + BATCH_SIZE, len(slice_files))\n        batch_files = slice_files[batch_start:batch_end]\n        \n        next_batch_start = batch_end\n        next_batch_end = min(next_batch_start + BATCH_SIZE, len(slice_files))\n        next_batch_files = slice_files[next_batch_start:next_batch_end] if next_batch_start < len(slice_files) else []\n        if next_batch_files:\n            next_batch_paths = [os.path.join(tomo_dir, f) for f in next_batch_files]\n            next_batch_thread = threading.Thread(target=preload_image_batch, args=(next_batch_paths,))\n            next_batch_thread.start()\n        else:\n            next_batch_thread = None\n        \n        sub_batches = np.array_split(batch_files, len(streams))\n        for i, sub_batch in enumerate(sub_batches):\n            if len(sub_batch) == 0:\n                continue\n            stream = streams[i % len(streams)]\n            with torch.cuda.stream(stream) if stream and device.startswith('cuda') else nullcontext():\n                sub_batch_paths = [os.path.join(tomo_dir, slice_file) for slice_file in sub_batch]\n                sub_batch_slice_nums = [int(slice_file.split('_')[1].split('.')[0]) for slice_file in sub_batch]\n                with GPUProfiler(f\"Inference batch {i+1}/{len(sub_batches)}\"):\n                    sub_results = model(sub_batch_paths, verbose=False)\n                for j, result in enumerate(sub_results):\n                    if len(result.boxes) > 0:\n                        for box_idx, confidence in enumerate(result.boxes.conf):\n                            if confidence >= CONFIDENCE_THRESHOLD:\n                                x1, y1, x2, y2 = result.boxes.xyxy[box_idx].cpu().numpy()\n                                x_center = (x1 + x2) / 2\n                                y_center = (y1 + y2) / 2\n                                all_detections.append({\n                                    'z': round(sub_batch_slice_nums[j]),\n                                    'y': round(y_center),\n                                    'x': round(x_center),\n                                    'confidence': float(confidence)\n                                })\n        if device.startswith('cuda'):\n            torch.cuda.synchronize()\n    \n    if next_batch_thread is not None:\n        next_batch_thread.join()\n    \n    final_detections = perform_3d_nms(all_detections, NMS_IOU_THRESHOLD)\n    final_detections.sort(key=lambda x: x['confidence'], reverse=True)\n    \n    if not final_detections:\n        return {'tomo_id': tomo_id, 'Motor axis 0': -1, 'Motor axis 1': -1, 'Motor axis 2': -1}\n    \n    best_detection = final_detections[0]\n    return {\n        'tomo_id': tomo_id,\n        'Motor axis 0': round(best_detection['z']),\n        'Motor axis 1': round(best_detection['y']),\n        'Motor axis 2': round(best_detection['x'])\n    }\n\ndef debug_image_loading(tomo_id):\n    tomo_dir = os.path.join(test_dir, tomo_id)\n    slice_files = sorted([f for f in os.listdir(tomo_dir) if f.endswith('.jpg')])\n    if not slice_files:\n        print(f\"No image files found in {tomo_dir}\")\n        return\n        \n    print(f\"Found {len(slice_files)} image files in {tomo_dir}\")\n    sample_file = slice_files[len(slice_files)//2]\n    img_path = os.path.join(tomo_dir, sample_file)\n    \n    try:\n        img_pil = Image.open(img_path)\n        print(f\"PIL Image shape: {np.array(img_pil).shape}, dtype: {np.array(img_pil).dtype}\")\n        img_cv2 = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n        print(f\"OpenCV Image shape: {img_cv2.shape}, dtype: {img_cv2.dtype}\")\n        img_rgb = cv2.cvtColor(cv2.imread(img_path), cv2.COLOR_BGR2RGB)\n        print(f\"OpenCV RGB Image shape: {img_rgb.shape}, dtype: {img_rgb.dtype}\")\n        print(\"Image loading successful!\")\n    except Exception as e:\n        print(f\"Error loading image {img_path}: {e}\")\n        \n    try:\n        test_model = YOLO(model_path)\n        test_results = test_model([img_path], verbose=False)\n        print(\"YOLO model successfully processed the test image\")\n    except Exception as e:\n        print(f\"Error with YOLO processing: {e}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T07:51:10.825509Z","iopub.execute_input":"2025-05-21T07:51:10.825772Z","iopub.status.idle":"2025-05-21T07:51:10.844852Z","shell.execute_reply.started":"2025-05-21T07:51:10.825757Z","shell.execute_reply":"2025-05-21T07:51:10.844209Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def generate_submission():\n    test_tomos = sorted([d for d in os.listdir(test_dir) if os.path.isdir(os.path.join(test_dir, d))])\n    total_tomos = len(test_tomos)\n    print(f\"Found {total_tomos} tomograms in test directory\")\n    \n    if test_tomos:\n        debug_image_loading(test_tomos[0])\n    \n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n    \n    print(f\"Loading YOLO model from {model_path}\")\n    model = YOLO(model_path)\n    model.to(device)\n    if device.startswith('cuda'):\n        model.fuse()\n        if torch.cuda.get_device_capability(0)[0] >= 7:\n            model.model.half()\n            print(\"Using half precision (FP16) for inference\")\n    \n    results = []\n    motors_found = 0\n    \n    with ThreadPoolExecutor(max_workers=1) as executor:\n        future_to_tomo = {}\n        for i, tomo_id in enumerate(test_tomos, 1):\n            future = executor.submit(process_tomogram, tomo_id, model, i, total_tomos)\n            future_to_tomo[future] = tomo_id\n        \n        for future in future_to_tomo:\n            tomo_id = future_to_tomo[future]\n            try:\n                if torch.cuda.is_available():\n                    torch.cuda.empty_cache()\n                result = future.result()\n                results.append(result)\n                has_motor = not pd.isna(result['Motor axis 0'])\n                if has_motor:\n                    motors_found += 1\n                    print(f\"Motor found in {tomo_id} at position: z={result['Motor axis 0']}, y={result['Motor axis 1']}, x={result['Motor axis 2']}\")\n                else:\n                    print(f\"No motor detected in {tomo_id}\")\n                print(f\"Current detection rate: {motors_found}/{len(results)} ({motors_found/len(results)*100:.1f}%)\")\n            except Exception as e:\n                print(f\"Error processing {tomo_id}: {e}\")\n                results.append({'tomo_id': tomo_id, 'Motor axis 0': -1, 'Motor axis 1': -1, 'Motor axis 2': -1})\n    \n    submission_df = pd.DataFrame(results)\n    submission_df = submission_df[['tomo_id', 'Motor axis 0', 'Motor axis 1', 'Motor axis 2']]\n    submission_df.to_csv(submission_path, index=False)\n    \n    print(f\"\\nSubmission complete!\")\n    print(f\"Motors detected: {motors_found}/{total_tomos} ({motors_found/total_tomos*100:.1f}%)\")\n    print(f\"Submission saved to: {submission_path}\")\n    print(\"\\nSubmission preview:\")\n    print(submission_df.head())\n    return submission_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T07:51:17.147945Z","iopub.execute_input":"2025-05-21T07:51:17.148685Z","iopub.status.idle":"2025-05-21T07:51:17.157614Z","shell.execute_reply.started":"2025-05-21T07:51:17.148659Z","shell.execute_reply":"2025-05-21T07:51:17.156796Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if __name__ == \"__main__\":\n    start_time = time.time()\n    submission = generate_submission()\n    elapsed = time.time() - start_time\n    print(f\"\\nTotal execution time: {elapsed:.2f} seconds ({elapsed/60:.2f} minutes)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T07:51:21.472755Z","iopub.execute_input":"2025-05-21T07:51:21.473048Z","iopub.status.idle":"2025-05-21T07:52:07.759584Z","shell.execute_reply.started":"2025-05-21T07:51:21.473027Z","shell.execute_reply":"2025-05-21T07:52:07.758934Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}