{"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,"sourceType":"competition"},{"sourceId":11274338,"sourceType":"datasetVersion","datasetId":6821877},{"sourceId":296173,"sourceType":"modelInstanceVersion","modelInstanceId":253402,"modelId":274853}],"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"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\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-04-05T08:36:37.375958Z","iopub.execute_input":"2025-04-05T08:36:37.376328Z","iopub.status.idle":"2025-04-05T08:36:37.382180Z","shell.execute_reply.started":"2025-04-05T08:36:37.376297Z","shell.execute_reply":"2025-04-05T08:36:37.381322Z"},"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Export your models to torchscript and try","metadata":{}},{"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\nmodel_path ='/kaggle/input/yolo-v11l/yolov9-s-960_v8_map-0.56.torchscript'\n\n# Define detection and processing parameters\nCONFIDENCE_THRESHOLD = 0.5\nMAX_DETECTIONS_PER_TOMO = 1\nDISTANCE_THRESHOLD = 60\nCONCENTRATION = 1 # Process a fraction of slices for fast submission\nSIZE = (960, 960)\n\n# Set device and dynamic batch size\ndevices = ['cuda:0','cuda:1']\nBATCH_SIZE = 32","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-05T08:36:37.383567Z","iopub.execute_input":"2025-04-05T08:36:37.383821Z","iopub.status.idle":"2025-04-05T08:36:37.412512Z","shell.execute_reply.started":"2025-04-05T08:36:37.383801Z","shell.execute_reply":"2025-04-05T08:36:37.411587Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class 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\ntorch.backends.cudnn.benchmark = True\ntorch.backends.cudnn.deterministic = False\ntorch.backends.cuda.matmul.allow_tf32 = True \ntorch.backends.cudnn.allow_tf32 = True","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-05T08:36:37.414096Z","iopub.execute_input":"2025-04-05T08:36:37.414405Z","iopub.status.idle":"2025-04-05T08:36:37.432940Z","shell.execute_reply.started":"2025-04-05T08:36:37.414384Z","shell.execute_reply":"2025-04-05T08:36:37.432348Z"}},"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, p98 = np.percentile(slice_data, [2, 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 letterbox(image, size=SIZE, color=(114, 114, 114), min_pad=2):\n    \"\"\"Resizes image with unchanged aspect ratio and pads it.\"\"\"\n    new_w, new_h = size\n\n    h,w = image.shape[:2]  # Current shape [height, width]\n      # Scale ratio (new / old)\n    ratio = min(new_w / w, new_h / h )\n    new_unpad_w,new_unpad_h  = (int(round(w * ratio) - 2 * min_pad), int(round(h * ratio)) - 2 * min_pad)  # Rescaled width, height\n    dw, dh = new_w - new_unpad_w, new_h - new_unpad_h  # Padding\n    ratio = min(new_unpad_w / w, new_unpad_h / h )\n    dw /= 2  # Divide padding into 2 sides\n    dh /= 2\n\n    # Resize image\n    image = cv2.resize(image, (new_unpad_w,new_unpad_h), interpolation=cv2.INTER_LINEAR)\n\n    # Padding\n    top, bottom = int(round(dh - 0.1)), int(round(dh + 0.1))\n    left, right = int(round(dw - 0.1)), int(round(dw + 0.1))\n    image_padded = cv2.copyMakeBorder(image, top, bottom, left, right, cv2.BORDER_CONSTANT, value=color)\n\n    return image_padded, ratio, left, top\n\ndef preprocess_image(image, size):\n    image, ratio, pad_w, pad_h = letterbox(image, size)\n    image = normalize_slice(image)\n    image = cv2.cvtColor(image, cv2.COLOR_GRAY2BGR)\n    #img = image_resized[:, :, ::-1].transpose(2, 0, 1)  # BGR to RGB, HWC to CHW\n    img = image.transpose(2, 0, 1)  # BGR to RGB, HWC to CHW\n\n    # Normalize and convert to float32\n    img = np.ascontiguousarray(img, dtype=np.float32) / 255.0\n    img_tensor = torch.from_numpy(img)  # Add batch dimension\n    return img_tensor, ratio, pad_w, pad_h\n\n\n\ndef yolo_batch_infer_impl(model,img_tensor,ratios,paddings, conf_thres):\n    device = next(model.parameters()).device\n    img_tensor = img_tensor.to(device).half()\n    with torch.no_grad():\n        outputs = model(img_tensor)  # Shape [B, 5, 8400]\n        if isinstance(outputs, list): #yolov9\n            outputs = outputs[0]\n\n    outputs = outputs.cpu().numpy()\n    ratios, paddings_w, paddings_h = ratios.numpy(),paddings[0].numpy(),paddings[1].numpy()\n    batch_results = []\n    for i, output in enumerate(outputs):  # Loop directly over model output\n        x_center, y_center, width, height, confidence = output\n        mask = confidence > conf_thres\n        if not mask.any():\n            batch_results.append([])\n            continue\n\n        x, y, w, h, conf = x_center[mask], y_center[mask], width[mask], height[mask], confidence[mask]\n\n        # Convert to original image coordinates\n        ratio, pad_w, pad_h = ratios[i], paddings_w[i], paddings_h[i]\n        x1, y1, x2, y2 = (x - w / 2 - pad_w) / ratio, (y - h / 2 - pad_h) / ratio, (x + w / 2 - pad_w) / ratio, (y + h / 2 - pad_h) / ratio\n        batch_results.append([(_x1,_y1,_x2,_y2,_c) for _x1,_y1,_x2,_y2,_c in zip(x1, y1, x2, y2, conf)])\n\n    return batch_results\n\n\ndef apply_clahe(img_tensor):\n    \"\"\"Applies CLAHE to a PyTorch tensor of shape (B, C, H, W).\"\"\"\n    B, C, H, W = img_tensor.shape\n    img_np = img_tensor.clone().cpu().numpy()  # Convert to NumPy if needed\n    img_np = img_np * 255\n    img_np = img_np.astype(np.uint8)\n    # CLAHE setup\n    clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))\n\n    # Process each image in the batch\n    for b in range(B):\n        for c in range(C):\n            img_np[b, c] = clahe.apply(img_np[b, c])\n\n    img_np = np.ascontiguousarray(img_np, dtype=np.float32) / 255.0\n    return torch.tensor(img_np, dtype=img_tensor.dtype, device=img_tensor.device)  # Convert back to tensor\n\ndef apply_hist_equalization(img_tensor):\n    \"\"\"Applies histogram equalization to a PyTorch tensor of shape (B, C, H, W).\"\"\"\n    B, C, H, W = img_tensor.shape\n    img_np = img_tensor.clone().cpu().numpy()\n\n    img_np = img_np * 255\n    img_np = img_np.astype(np.uint8)\n    # Process each image in the batch\n    for b in range(B):\n        for c in range(C):\n            img_np[b, c] = cv2.equalizeHist(img_np[b, c]) \n\n    img_np = np.ascontiguousarray(img_np, dtype=np.float32) / 255.0\n    return torch.tensor(img_np, dtype=img_tensor.dtype, device=img_tensor.device)\n\ndef iou(box1, box2):\n    \"\"\"Compute Intersection over Union (IoU) between two boxes.\"\"\"\n    x1, y1, x2, y2, _ = box1\n    x1g, y1g, x2g, y2g, _ = box2\n\n    inter_x1 = max(x1, x1g)\n    inter_y1 = max(y1, y1g)\n    inter_x2 = min(x2, x2g)\n    inter_y2 = min(y2, y2g)\n\n    inter_area = max(0, inter_x2 - inter_x1) * max(0, inter_y2 - inter_y1)\n    box1_area = (x2 - x1) * (y2 - y1)\n    box2_area = (x2g - x1g) * (y2g - y1g)\n\n    union_area = box1_area + box2_area - inter_area\n    return inter_area / union_area if union_area > 0 else 0\n\ndef weighted_box_fusion(boxes, iou_threshold=0.5):\n    \"\"\"Applies Weighted Box Fusion to combine overlapping bounding boxes.\"\"\"\n    fused_boxes = []\n    used = [False] * len(boxes)\n\n    for i in range(len(boxes)):\n        if used[i]:\n            continue\n\n        similar_boxes = [boxes[i]]\n        used[i] = True\n\n        for j in range(i + 1, len(boxes)):\n            if used[j]:\n                continue\n\n            if iou(boxes[i], boxes[j]) > iou_threshold:\n                similar_boxes.append(boxes[j])\n                used[j] = True\n\n        # Compute weighted average for the final box\n        similar_boxes = np.array(similar_boxes)\n        confidences = similar_boxes[:, 4]\n        weights = confidences / confidences.sum()\n\n        fused_x1 = np.sum(similar_boxes[:, 0] * weights)\n        fused_y1 = np.sum(similar_boxes[:, 1] * weights)\n        fused_x2 = np.sum(similar_boxes[:, 2] * weights)\n        fused_y2 = np.sum(similar_boxes[:, 3] * weights)\n        fused_confidence = np.mean(confidences) #np.max(confidences)  # Take max confidence\n\n        fused_boxes.append((fused_x1, fused_y1, fused_x2, fused_y2, fused_confidence))\n\n    return fused_boxes\n\ndef tta_yolo_batch_infer(model, img_tensor, ratios, paddings, conf_thres):\n\n    augmentations = {\n        \"original\": lambda x: x,\n        #\"h_flip\": lambda x: torch.flip(x, dims=[3]),  # Flip width (horizontal)\n        #\"v_flip\": lambda x: torch.flip(x, dims=[2]),  # Flip height (vertical)\n        \"hv_flip\": lambda x: torch.flip(x, dims=[2, 3]),  # Flip both\n        #\"hist\": lambda x: apply_hist_equalization(x),\n        #\"clahe\":  lambda x: apply_clahe(x),\n        \"shift\": lambda x: x + 0.2,\n        #\"scale\": lambda x: x * 1.2,\n        #\"inverse\": lambda x: 1 - x,\n    }\n\n    B,C,H,W = img_tensor.shape\n    results = [[] for _ in range(B)]\n    paddings_w, paddings_h, ratios_ = paddings[0].numpy(), paddings[1].numpy(), ratios.numpy()\n    \n    for aug_name, aug_fn in augmentations.items():\n        img_tensor_ = aug_fn(img_tensor.clone())  # Apply augmentation\n        # Run inference\n        with torch.no_grad():\n            outputs = yolo_batch_infer_impl(model,img_tensor_,ratios,paddings, conf_thres)\n\n        for i, output in enumerate(outputs):\n            pad_w, pad_h, ratio = paddings_w[i], paddings_h[i], ratios_[i]\n            for x1, y1, x2, y2, conf in output:\n\n                #Reverse augmentation\n                if aug_name == \"h_flip\":\n                    x1 = (W - pad_w*2 - x1*ratio)/ratio\n                    x2 = (W - pad_w*2 - x2*ratio)/ratio\n                elif aug_name == \"v_flip\":\n                    y1 = (H - pad_h*2 - y1*ratio)/ratio\n                    y2 = (H - pad_h*2 - y2*ratio)/ratio\n                elif aug_name == \"hv_flip\":\n                    y1 = (H - pad_h*2 - y1*ratio)/ratio\n                    y2 = (H - pad_h*2 - y2*ratio)/ratio\n                    x1 = (W - pad_w*2 - x1*ratio)/ratio\n                    x2 = (W - pad_w*2 - x2*ratio)/ratio\n                    \n                results[i].append((x1, y1, x2, y2, conf))\n\n    return [weighted_box_fusion(r, 0.1) for r in results]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-05T08:36:37.478797Z","iopub.execute_input":"2025-04-05T08:36:37.479043Z","iopub.status.idle":"2025-04-05T08:36:37.868860Z","shell.execute_reply.started":"2025-04-05T08:36:37.479022Z","shell.execute_reply":"2025-04-05T08:36:37.868097Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class MyDataset(torch.utils.data.Dataset):\n    def __init__(self,tomo_dir, files, size):\n        self.files = files\n        self.tomo_dir = tomo_dir\n        self.size = size\n\n    def __getitem__(self, index):\n        image_path = os.path.join(self.tomo_dir, self.files[index])\n        image = cv2.imread(image_path, cv2.IMREAD_GRAYSCALE)\n        if image is None:\n            image = np.array(Image.open(image_path))\n        return (*preprocess_image(image, self.size), int(self.files[index].split('_')[1].split('.')[0]))\n        \n    def __len__(self):\n        return len(self.files)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-05T08:36:37.869849Z","iopub.execute_input":"2025-04-05T08:36:37.870075Z","iopub.status.idle":"2025-04-05T08:36:37.899968Z","shell.execute_reply.started":"2025-04-05T08:36:37.870056Z","shell.execute_reply":"2025-04-05T08:36:37.899388Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def perform_3d_nms(detections, distance_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    while detections:\n        best_detection = detections.pop(0)\n\n        final_detections.append(best_detection)\n\n        detections = [d for d in detections if distance_3d(d, best_detection) > distance_threshold]\n    return final_detections\n\n        \ndef process_tomogram(tomo_id, model, index=0, total=1, device=''):\n    \"\"\"\n    Process a single tomogram and return the most confident motor detection.\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    \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    \n    if device.startswith('cuda'):\n        streams = [torch.cuda.Stream(device=device) for _ in range(min(4, BATCH_SIZE))]\n    else:\n        streams = [None]\n\n    dataset = MyDataset(tomo_dir, slice_files, SIZE)\n    dataloader = torch.utils.data.DataLoader(dataset, batch_size=BATCH_SIZE, num_workers = 2)\n    \n    for batch in tqdm(dataloader, desc=f\"Processing tomogram {tomo_id} ({index}/{total}) | {len(slice_files)} out of {len(os.listdir(tomo_dir))} slices (CONCENTRATION={CONCENTRATION}): \"):\n\n        sub_batch_results = []\n        \n        images,ratios,paddings_w,paddings_h,indexes = batch\n        sub_batches = (torch.tensor_split(images, len(streams)),\n                        torch.tensor_split(ratios, len(streams)),\n                        torch.tensor_split(paddings_w, len(streams)),\n                        torch.tensor_split(paddings_h, len(streams)),\n                        torch.tensor_split(indexes, len(streams))\n                      )\n\n        for i, (sub_images,sub_ratios,sub_paddings_w,sub_paddings_h,sub_indexes) in enumerate(zip(*sub_batches)):\n            if len(sub_images) == 0:\n                continue\n            stream = streams[i % len(streams)]\n            with torch.cuda.stream(stream) if stream and device.startswith('cuda') else nullcontext():\n                with GPUProfiler(f\"Inference batch {i+1}/{len(dataloader)}\"):\n                    sub_results = tta_yolo_batch_infer(model, sub_images,sub_ratios,(sub_paddings_w,sub_paddings_h),CONFIDENCE_THRESHOLD)\n                # Process each result in this sub-batch\n                for j, output in enumerate(sub_results):\n                    for x1, y1, x2, y2, conf in output:\n                        x_center = (x1 + x2) / 2\n                        y_center = (y1 + y2) / 2\n                        # Store detection with 3D coordinates\n                        #print(sub_indexes[j])\n                        all_detections.append({\n                            'z': indexes[j].item(),\n                            'y': round(y_center),\n                            'x': round(x_center),\n                            'confidence': float(conf)\n                        })\n        \n        # Synchronize streams\n        if device.startswith('cuda'):\n            torch.cuda.synchronize()\n\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n        \n    final_detections = perform_3d_nms(all_detections, DISTANCE_THRESHOLD)\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    final_detections.sort(key=lambda x: x['confidence'], reverse=True)\n    best_detection = final_detections[0]\n\n\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","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-05T08:36:37.901683Z","iopub.execute_input":"2025-04-05T08:36:37.901960Z","iopub.status.idle":"2025-04-05T08:36:37.926454Z","shell.execute_reply.started":"2025-04-05T08:36:37.901923Z","shell.execute_reply":"2025-04-05T08:36:37.925816Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from joblib import Parallel, delayed\n\ndef generate_submission(test_tomos):\n    \"\"\"\n    Main function to generate the submission file.\n    \"\"\"\n\n    total_tomos = len(test_tomos)\n    print(f\"Found {total_tomos} tomograms in test directory\")\n    \n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n    \n    print(f\"Loading YOLO model from {model_path}\")\n    models = [torch.jit.load(model_path, map_location=device).eval().half() for device in devices]\n    \n    results = []\n    motors_found = 0\n    future_to_tomo = {}\n    \n    with ThreadPoolExecutor(max_workers=4) as executor:\n        \n        for i, tomo_id in enumerate(test_tomos):\n            future = executor.submit(process_tomogram, tomo_id, models[i%2], i, total_tomos, devices[i%2])\n            future_to_tomo[future] = tomo_id\n            \n    for future,tomo_id in future_to_tomo.items():\n        try:\n            result = future.result()\n            results.append(result)\n        except Exception as e:\n            print(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\n    return submission_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-05T08:36:37.927279Z","iopub.execute_input":"2025-04-05T08:36:37.927595Z","iopub.status.idle":"2025-04-05T08:36:37.952080Z","shell.execute_reply.started":"2025-04-05T08:36:37.927568Z","shell.execute_reply":"2025-04-05T08:36:37.951225Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"start_time = time.time()\n\ntest_tomos = sorted([d for d in os.listdir(test_dir) if os.path.isdir(os.path.join(test_dir, d))])\n\nsubmission = generate_submission(test_tomos)\nelapsed = time.time() - start_time\nprint(f\"\\nTotal execution time: {elapsed:.2f} seconds ({elapsed/60:.2f} minutes)\")\n\nprint(\"\\nSubmission preview:\")\nsubmission.to_csv(submission_path, index=False)\nprint(submission.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-05T08:36:37.953180Z","iopub.execute_input":"2025-04-05T08:36:37.953574Z","iopub.status.idle":"2025-04-05T08:37:54.480606Z","shell.execute_reply.started":"2025-04-05T08:36:37.953543Z","shell.execute_reply":"2025-04-05T08:37:54.479676Z"}},"outputs":[],"execution_count":null}]}