{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":91249,"databundleVersionId":11294684,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":5351953,"sourceType":"datasetVersion","datasetId":3106911},{"sourceId":10981225,"sourceType":"datasetVersion","datasetId":6830489},{"sourceId":10986608,"sourceType":"datasetVersion","datasetId":6838043},{"sourceId":226864880,"sourceType":"kernelVersion"}],"dockerImageVersionId":30918,"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-03-12T09:35:31.715366Z","iopub.execute_input":"2025-03-12T09:35:31.715684Z","iopub.status.idle":"2025-03-12T09:36:32.620468Z","shell.execute_reply.started":"2025-03-12T09:35:31.715658Z","shell.execute_reply":"2025-03-12T09:36:32.619389Z"}},"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-03-12T09:36:32.621710Z","iopub.execute_input":"2025-03-12T09:36:32.622048Z","iopub.status.idle":"2025-03-12T09:36:40.821919Z","shell.execute_reply.started":"2025-03-12T09:36:32.622017Z","shell.execute_reply":"2025-03-12T09:36:40.821187Z"}},"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-03-12T09:36:40.823564Z","iopub.execute_input":"2025-03-12T09:36:40.824132Z","iopub.status.idle":"2025-03-12T09:36:40.914719Z","shell.execute_reply.started":"2025-03-12T09:36:40.824108Z","shell.execute_reply":"2025-03-12T09:36:40.913921Z"}},"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-03-12T09:36:40.915790Z","iopub.execute_input":"2025-03-12T09:36:40.916126Z","iopub.status.idle":"2025-03-12T09:36:40.985577Z","shell.execute_reply.started":"2025-03-12T09:36:40.916091Z","shell.execute_reply":"2025-03-12T09:36:40.984957Z"}},"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='green', 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-03-12T09:36:40.986311Z","iopub.execute_input":"2025-03-12T09:36:40.986570Z","iopub.status.idle":"2025-03-12T09:36:41.345461Z","shell.execute_reply.started":"2025-03-12T09:36:40.986538Z","shell.execute_reply":"2025-03-12T09:36:41.344503Z"}},"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-03-12T09:36:41.346359Z","iopub.execute_input":"2025-03-12T09:36:41.346631Z","iopub.status.idle":"2025-03-12T09:36:41.371432Z","shell.execute_reply.started":"2025-03-12T09:36:41.346597Z","shell.execute_reply":"2025-03-12T09:36:41.370678Z"}},"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-03-12T09:36:41.372403Z","iopub.execute_input":"2025-03-12T09:36:41.372650Z","iopub.status.idle":"2025-03-12T09:36:43.562817Z","shell.execute_reply.started":"2025-03-12T09:36:41.372629Z","shell.execute_reply":"2025-03-12T09:36:43.561900Z"}},"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-03-12T09:36:43.565095Z","iopub.execute_input":"2025-03-12T09:36:43.565351Z","iopub.status.idle":"2025-03-12T09:36:43.753653Z","shell.execute_reply.started":"2025-03-12T09:36:43.565329Z","shell.execute_reply":"2025-03-12T09:36:43.752889Z"}},"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=\"red\"  # 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-03-12T09:36:43.755258Z","iopub.execute_input":"2025-03-12T09:36:43.755532Z","iopub.status.idle":"2025-03-12T09:36:45.736563Z","shell.execute_reply.started":"2025-03-12T09:36:43.755510Z","shell.execute_reply":"2025-03-12T09:36:45.735645Z"}},"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=\"spring\",  # 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-03-12T09:36:45.737619Z","iopub.execute_input":"2025-03-12T09:36:45.737951Z","iopub.status.idle":"2025-03-12T09:36:46.140867Z","shell.execute_reply.started":"2025-03-12T09:36:45.737923Z","shell.execute_reply":"2025-03-12T09:36:46.140083Z"}},"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-03-12T09:36:46.141774Z","iopub.execute_input":"2025-03-12T09:36:46.142063Z","iopub.status.idle":"2025-03-12T09:36:57.523853Z","shell.execute_reply.started":"2025-03-12T09:36:46.142039Z","shell.execute_reply":"2025-03-12T09:36:57.522952Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_images(path, n_images=12, is_random=True, figsize=(14, 14)):\n    \"\"\"\n    Visualize a set of images from a directory.\n    \n    Args:\n        path (str): Directory path containing images.\n        n_images (int): Number of images to display.\n        is_random (bool): If True, display random images; else, the first n_images.\n        figsize (tuple): Size of the figure.\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-03-12T09:36:57.524755Z","iopub.execute_input":"2025-03-12T09:36:57.525084Z","iopub.status.idle":"2025-03-12T09:36:59.135365Z","shell.execute_reply.started":"2025-03-12T09:36:57.525059Z","shell.execute_reply":"2025-03-12T09:36:59.134099Z"}},"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-03-12T09:36:59.136362Z","iopub.execute_input":"2025-03-12T09:36:59.136624Z","iopub.status.idle":"2025-03-12T09:37:00.294184Z","shell.execute_reply.started":"2025-03-12T09:36:59.136600Z","shell.execute_reply":"2025-03-12T09:37:00.293365Z"}},"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    \"\"\"\n    Normalize slice data using the 2nd and 98th percentiles.\n    \n    Args:\n        slice_data (numpy.array): Input image slice.\n    \n    Returns:\n        np.uint8: Normalized image in the range [0, 255].\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\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    Extract slices containing motors and save images with corresponding YOLO annotations.\n    \n    Steps:\n    - Load the motor labels.\n    - Perform a train/validation split by tomogram.\n    - For each motor, extract slices in a range (± trust parameter).\n    - Normalize each slice and save it.\n    - Generate YOLO format bounding box annotations with a fixed box size.\n    - Create a YAML configuration file for YOLO training.\n    \n    Returns:\n        dict: A summary containing dataset statistics and file paths.\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-03-12T09:37:00.295153Z","iopub.execute_input":"2025-03-12T09:37:00.295450Z","iopub.status.idle":"2025-03-12T09:40:04.112340Z","shell.execute_reply.started":"2025-03-12T09:37:00.295423Z","shell.execute_reply":"2025-03-12T09:40:04.111421Z"}},"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    \"\"\"\n    Visualize random training samples with YOLO annotations.\n    \n    Args:\n        num_samples (int): Number of random images to display.\n    \"\"\"\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)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-12T09:40:04.113395Z","iopub.execute_input":"2025-03-12T09:40:04.113701Z","iopub.status.idle":"2025-03-12T09:40:05.847730Z","shell.execute_reply.started":"2025-03-12T09:40:04.113677Z","shell.execute_reply":"2025-03-12T09:40:05.846882Z"}},"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/yolov8-original-pretrained-for-detection/detection/yolov8m.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-03-12T09:40:05.848763Z","iopub.execute_input":"2025-03-12T09:40:05.849019Z","iopub.status.idle":"2025-03-12T09:40:05.854626Z","shell.execute_reply.started":"2025-03-12T09:40:05.848999Z","shell.execute_reply":"2025-03-12T09:40:05.853845Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def fix_yaml_paths(yaml_path):\n    \"\"\"\n    Fix the paths in the YAML file to match the actual Kaggle directories.\n    \n    Args:\n        yaml_path (str): Path to the original dataset YAML file.\n        \n    Returns:\n        str: Path to the fixed YAML file.\n    \"\"\"\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-03-12T09:40:05.855453Z","iopub.execute_input":"2025-03-12T09:40:05.855706Z","iopub.status.idle":"2025-03-12T09:40:05.867489Z","shell.execute_reply.started":"2025-03-12T09:40:05.855684Z","shell.execute_reply":"2025-03-12T09:40:05.866593Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_dfl_loss_curve(run_dir):\n    \"\"\"\n    Plot the DFL loss curves for training and validation, marking the best model.\n    \n    Args:\n        run_dir (str): Directory where the training results are stored.\n    \"\"\"\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-03-12T09:40:05.868466Z","iopub.execute_input":"2025-03-12T09:40:05.868665Z","iopub.status.idle":"2025-03-12T09:40:05.891604Z","shell.execute_reply.started":"2025-03-12T09:40:05.868647Z","shell.execute_reply":"2025-03-12T09:40:05.891131Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nfrom ultralytics import YOLO\n\ndef train_yolo_model(yaml_path, pretrained_weights_path, epochs=50, batch_size=16, img_size=640):\n    \"\"\"\n    Train a YOLO model on the prepared dataset with optimized accuracy settings.\n\n    Args:\n        yaml_path (str): Path to the dataset YAML file.\n        pretrained_weights_path (str): Path to pre-downloaded weights file.\n        epochs (int): Number of training epochs.\n        batch_size (int): Batch size for training.\n        img_size (int): Image size for training.\n\n    Returns:\n        model (YOLO): Trained YOLO model.\n        results: Training results.\n    \"\"\"\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=30,  # Stop training if no improvement after 10 epochs\n        save_period=5,  # Save model every 5 epochs\n        val=True,\n        verbose=True,\n        optimizer=\"AdamW\",  # AdamW optimizer for stability\n        lr0=0.001,  # Initial learning rate\n        lrf=0.01,  # Final learning rate factor\n        cos_lr=True,  # Use cosine learning rate decay\n        weight_decay=0.0005,  # Prevent overfitting\n        momentum=0.937,  # Momentum for better gradient updates\n        close_mosaic=10,  # Disable mosaic augmentation after 10 epochs\n        mixup=0.2,  # Apply mixup augmentation\n        workers=4,  # Speed up data loading\n        augment=True,  # Enable additional augmentations\n        amp=True,  # Mixed precision training for faster performance\n        dropout=0.1,\n    )\n\n    run_dir = os.path.join(yolo_weights_dir, 'motor_detector')\n    \n    # If function is defined, plot loss curves for better insights\n    if 'plot_dfl_loss_curve' in globals():\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-03-12T09:40:05.892517Z","iopub.execute_input":"2025-03-12T09:40:05.892829Z","iopub.status.idle":"2025-03-12T09:40:05.915591Z","shell.execute_reply.started":"2025-03-12T09:40:05.892783Z","shell.execute_reply":"2025-03-12T09:40:05.915009Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def predict_on_samples(model, num_samples=4):\n    \"\"\"\n    Run predictions on random validation samples and display results.\n    \n    Args:\n        model: Trained YOLO model.\n        num_samples (int): Number of random samples to test.\n    \"\"\"\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-03-12T09:40:05.916374Z","iopub.execute_input":"2025-03-12T09:40:05.916652Z","iopub.status.idle":"2025-03-12T09:40:05.937410Z","shell.execute_reply.started":"2025-03-12T09:40:05.916631Z","shell.execute_reply":"2025-03-12T09:40:05.936788Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def prepare_dataset():\n    \"\"\"\n    Check if the dataset exists and create/fix a proper YAML file for training.\n    \n    Returns:\n        str: Path to the YAML file to use for training.\n    \"\"\"\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=100  # 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\n# if __name__ == \"__main__\":\n#     main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-12T09:40:05.938169Z","iopub.execute_input":"2025-03-12T09:40:05.938357Z","iopub.status.idle":"2025-03-12T09:40:05.960850Z","shell.execute_reply.started":"2025-03-12T09:40:05.938339Z","shell.execute_reply":"2025-03-12T09:40:05.960263Z"}},"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/input/byu-a-106-yolov8l-daexp-mixupexp/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-03-12T09:40:05.961552Z","iopub.execute_input":"2025-03-12T09:40:05.961778Z","iopub.status.idle":"2025-03-12T09:40:06.016885Z","shell.execute_reply.started":"2025-03-12T09:40:05.961748Z","shell.execute_reply":"2025-03-12T09:40:06.016238Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def normalize_slice(slice_data):\n    \"\"\"\n    Normalize slice data using the 2nd and 98th percentiles.\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    \"\"\"Preload a batch of images to CPU memory.\"\"\"\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    \"\"\"\n    Perform 3D Non-Maximum Suppression on detections to merge nearby motors.\n    \"\"\"\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    \"\"\"\n    Process a single tomogram and return the most confident motor detection.\n    \"\"\"\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    \"\"\"\n    Debug function to test image loading methods.\n    \"\"\"\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-03-12T09:40:06.019247Z","iopub.execute_input":"2025-03-12T09:40:06.019466Z","iopub.status.idle":"2025-03-12T09:40:06.038551Z","shell.execute_reply.started":"2025-03-12T09:40:06.019446Z","shell.execute_reply":"2025-03-12T09:40:06.037943Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def generate_submission():\n    \"\"\"\n    Main function to generate the submission file.\n    \"\"\"\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-03-12T09:40:06.039448Z","iopub.execute_input":"2025-03-12T09:40:06.039759Z","iopub.status.idle":"2025-03-12T09:40:06.062602Z","shell.execute_reply.started":"2025-03-12T09:40:06.039738Z","shell.execute_reply":"2025-03-12T09:40:06.061873Z"}},"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-03-12T09:40:06.063531Z","iopub.execute_input":"2025-03-12T09:40:06.063848Z","iopub.status.idle":"2025-03-12T09:41:49.397456Z","shell.execute_reply.started":"2025-03-12T09:40:06.063816Z","shell.execute_reply":"2025-03-12T09:41:49.396645Z"}},"outputs":[],"execution_count":null}]}