{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","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":"gpu","dataSources":[{"sourceId":91249,"databundleVersionId":11294684,"sourceType":"competition"},{"sourceId":244158115,"sourceType":"kernelVersion"}],"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true},"papermill":{"default_parameters":{},"duration":104.600268,"end_time":"2025-04-09T10:15:05.748280","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2025-04-09T10:13:21.148012","version":"2.6.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"from concurrent.futures import ThreadPoolExecutor\nfrom tqdm.notebook import tqdm\nfrom scipy.spatial import distance\nfrom ultralytics import RTDETR\nfrom PIL import Image\nimport matplotlib.patches as patches\nimport matplotlib.pyplot as plt\nimport networkx as nx\nimport pandas as pd\nimport numpy as np\nimport threading\nimport random\nimport torch\nimport cv2\nimport os","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-output":true,"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","papermill":{"duration":8.645163,"end_time":"2025-04-09T10:13:32.424319","exception":false,"start_time":"2025-04-09T10:13:23.779156","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CFG:\n    dataset_path = \"/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/\"\n    test_image_path = os.path.join(dataset_path, \"test\")\n    model_path = \"/kaggle/input/flagellar-motor-detection-2-3-rt-detr-training/byu-locating-bacterial-flagellar-motors/rtdetr-l_fold_0/weights/best.pt\"\n\n    seed = 42\n    device = 'cuda:0'","metadata":{"papermill":{"duration":0.009792,"end_time":"2025-04-09T10:13:32.438581","exception":false,"start_time":"2025-04-09T10:13:32.428789","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.manual_seed(CFG.seed)\nnp.random.seed(CFG.seed)\nrandom.seed(CFG.seed)","metadata":{"papermill":{"duration":0.012816,"end_time":"2025-04-09T10:13:32.455257","exception":false,"start_time":"2025-04-09T10:13:32.442441","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.backends.cudnn.benchmark = True\ntorch.backends.cudnn.deterministic = False\ntorch.backends.cuda.matmul.allow_tf32 = True\ntorch.backends.cudnn.allow_tf32 = True\n\ngpu_name = torch.cuda.get_device_name(0)\ngpu_mem = torch.cuda.get_device_properties(0).total_memory / 1e9  \n\nfree_mem = gpu_mem - torch.cuda.memory_allocated(0) / 1e9\nBATCH_SIZE = max(8, min(32, int(free_mem * 4)))\n\ntorch.cuda.empty_cache()","metadata":{"papermill":{"duration":0.085198,"end_time":"2025-04-09T10:13:32.542675","exception":false,"start_time":"2025-04-09T10:13:32.457477","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class GPUProfiler:\n    def __init__(self, name):\n        self.name = name\n        \n    def __enter__(self):\n        torch.cuda.synchronize()\n        return self\n        \n    def __exit__(self, *args):\n        torch.cuda.synchronize()","metadata":{"papermill":{"duration":0.008905,"end_time":"2025-04-09T10:13:32.554113","exception":false,"start_time":"2025-04-09T10:13:32.545208","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def 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\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\n\ndef process_tomogram(tomo_id, model, concentration=1):\n    tomo_dir = os.path.join(CFG.test_image_path, tomo_id)\n    slice_files = sorted([f for f in os.listdir(tomo_dir) if f.endswith('.jpg')])\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    all_detections = []\n    streams = [torch.cuda.Stream() for _ in range(min(4, BATCH_SIZE))]\n    next_batch_thread = None\n    \n    for batch_start in tqdm(range(0, len(slice_files), BATCH_SIZE), tomo_id):\n        if next_batch_thread is not None:\n            next_batch_thread.join()\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\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            \n            stream = streams[i % len(streams)]\n            with torch.cuda.stream(stream):\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                \n                with GPUProfiler(f\"Inference batch {i+1}/{len(sub_batches)}\"):\n                    sub_results = model(sub_batch_paths, verbose=False, augment=False)\n                    \n                for j, result in enumerate(sub_results):\n                    if len(result.boxes) > 0:\n                        boxes = result.boxes\n                        for box_idx, confidence in enumerate(boxes.conf):\n                            x1, y1, x2, y2 = boxes.xyxy[box_idx].cpu().numpy()\n                            \n                            x_center = (x1 + x2) / 2\n                            y_center = (y1 + y2) / 2\n\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                                'tomo_id': tomo_id\n                            })\n\n        torch.cuda.synchronize()\n\n    if next_batch_thread is not None:\n        next_batch_thread.join()\n\n    if not all_detections:\n        all_detections.append({\n            'z': -1,\n            'y': -1,\n            'x': -1,\n            'confidence': -1,\n            'tomo_id': tomo_id\n        })\n        \n\n    return all_detections","metadata":{"papermill":{"duration":0.027783,"end_time":"2025-04-09T10:13:32.584265","exception":false,"start_time":"2025-04-09T10:13:32.556482","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = RTDETR(CFG.model_path)\nmodel.to(CFG.device)\nmodel.fuse()\n\nif torch.cuda.get_device_capability(0)[0] >= 7: \n    model.model.half()","metadata":{"papermill":{"duration":1.620339,"end_time":"2025-04-09T10:13:34.207071","exception":false,"start_time":"2025-04-09T10:13:32.586732","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%time\n\ntest_tomos = sorted([d for d in os.listdir(CFG.test_image_path) if os.path.isdir(os.path.join(CFG.test_image_path, d))])\n\nall_detections = []\nwith ThreadPoolExecutor(max_workers=1) as executor:\n    future_to_tomo = {}\n    \n    for i, tomo_id in enumerate(test_tomos, 1):\n        future = executor.submit(process_tomogram, tomo_id, model)\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            torch.cuda.empty_cache()\n            result = future.result()\n            all_detections.extend(result)\n        except Exception as e:\n            print(e)","metadata":{"papermill":{"duration":88.581157,"end_time":"2025-04-09T10:15:02.790942","exception":false,"start_time":"2025-04-09T10:13:34.209785","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"all_detections = pd.DataFrame(all_detections)\nif all_detections.tomo_id.nunique() == 3:\n    all_detections.to_csv(\"all_detections.csv\", index=False)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def postprocess(df, confidence_threshold, group_distance_threshold, min_detection_per_group):\n    df = df[df['confidence'] >= confidence_threshold]\n    \n    df_agg_det_tomo_list = []\n    for tomo_id, df_det_tomo in df.groupby('tomo_id'):\n        dist_mat = distance.cdist(df_det_tomo[['x', 'y', 'z']], df_det_tomo[['x', 'y', 'z']], metric='euclidean')\n        adj_mat = (dist_mat <= group_distance_threshold).astype(int)\n        np.fill_diagonal(adj_mat, 0)\n        G = nx.from_numpy_array(adj_mat)\n        connected_components = list(nx.connected_components(G))\n        agg_det_dict_list = []\n        for group_idx_set in connected_components:\n            df_det_grp = df_det_tomo.iloc[list(group_idx_set)] \n\n            z = (df_det_grp['z'] * df_det_grp['confidence']).sum() / df_det_grp['confidence'].sum()\n            y = (df_det_grp['y'] * df_det_grp['confidence']).sum() / df_det_grp['confidence'].sum()\n            x = (df_det_grp['x'] * df_det_grp['confidence']).sum() / df_det_grp['confidence'].sum()\n            score_mean = df_det_grp['confidence'].mean()\n\n            agg_det = {\n                'tomo_id': tomo_id,\n                'x': x,\n                'y': y,\n                'z': z,\n                'confidence': score_mean,\n                'group_det_count': len(group_idx_set),\n            }\n            agg_det_dict_list.append(agg_det)\n            \n        df_agg_det_tomo = pd.DataFrame(agg_det_dict_list)\n        df_agg_det_tomo = df_agg_det_tomo[df_agg_det_tomo['group_det_count'] >= min_detection_per_group]\n        df_agg_det_tomo = df_agg_det_tomo.sort_values(by=['group_det_count', 'confidence'], ascending=False).iloc[:1]\n        df_agg_det_tomo_list.append(df_agg_det_tomo)\n        \n    return pd.concat(df_agg_det_tomo_list)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def add_no_motor_tomos(predictions, test_tomos):\n    no_motors_ids = list(set(test_tomos) - set(predictions['tomo_id']))\n    \n    no_motors_sf = pd.DataFrame({\n        'tomo_id': no_motors_ids,\n        'x': [-1] * len(no_motors_ids),\n        'y': [-1] * len(no_motors_ids),\n        'z': [-1] * len(no_motors_ids),\n        'confidence': [0] * len(no_motors_ids),\n        'fold': [-1] * len(no_motors_ids)\n    })\n    \n    return pd.concat([predictions, no_motors_sf])\n\ndef fix_dtypes(predictions):\n    predictions[\"x\"] = predictions[\"x\"].astype(int)\n    predictions[\"y\"] = predictions[\"y\"].astype(int)\n    predictions[\"z\"] = predictions[\"z\"].astype(int)\n    return predictions","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"final_predictions = postprocess(all_detections, confidence_threshold=0.6170582213804651, group_distance_threshold=82, min_detection_per_group=3)\nfinal_predictions = add_no_motor_tomos(final_predictions, test_tomos)\nfinal_predictions = fix_dtypes(final_predictions)\nfinal_predictions","metadata":{"papermill":{"duration":0.039209,"end_time":"2025-04-09T10:15:02.832781","exception":false,"start_time":"2025-04-09T10:15:02.793572","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if len(final_predictions) == 3:\n    points = []\n    image_paths = []\n    \n    for i, r in final_predictions.iterrows():\n        if (\n            r['z'] != -1 and \n            r['y'] != -1 and \n            r['x'] != -1\n        ):\n            slice_num = int(r['z'])\n            slice_str = f\"{slice_num:04d}\"\n            image_path = f\"/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/test/{r['tomo_id']}/slice_{slice_str}.jpg\"\n            image_paths.append(image_path)\n            points.append((r['x'], r['y'], r['confidence']))\n\n    fig, axes = plt.subplots(1, 3, figsize=(15, 5))\n    box_size = 64\n    half_box = box_size // 2\n\n    for i, path in enumerate(image_paths):\n        img = cv2.imread(path)\n        img_rgb = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        axes[i].imshow(img_rgb)\n        x, y, confidence = points[i]\n        \n        axes[i].scatter(x, y, color=\"red\")\n        \n        rect = patches.Rectangle((x - half_box, y - half_box), box_size, box_size, linewidth=2, edgecolor='lime', facecolor='none')\n        axes[i].add_patch(rect)\n        \n        axes[i].text(x - half_box, y - half_box - 10, f\"Confidence: {confidence:.4f}\", color='lime', fontsize=10, weight='bold', ha='center')\n        axes[i].set_title(path.split(\"/\")[-2] + \"/\" + path.split(\"/\")[-1])\n        axes[i].axis('off')\n\n    plt.tight_layout()\n    plt.show()","metadata":{"papermill":{"duration":1.165579,"end_time":"2025-04-09T10:15:04.001422","exception":false,"start_time":"2025-04-09T10:15:02.835843","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission = final_predictions.copy()\nsubmission = submission.rename(columns={\"x\": \"Motor axis 2\", \"y\": \"Motor axis 1\", \"z\": \"Motor axis 0\"})\nsubmission = submission[['tomo_id', 'Motor axis 0', 'Motor axis 1', 'Motor axis 2']]\nsubmission.to_csv(\"submission.csv\", index=False)\nsubmission.head()","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}