{"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":"gpu","dataSources":[{"sourceId":91249,"databundleVersionId":11294684,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":11036126,"sourceType":"datasetVersion","datasetId":6873934},{"sourceId":11044741,"sourceType":"datasetVersion","datasetId":6880023},{"sourceId":11195676,"sourceType":"datasetVersion","datasetId":6989638},{"sourceId":11522950,"sourceType":"datasetVersion","datasetId":6899986},{"sourceId":11947116,"sourceType":"datasetVersion","datasetId":7238009},{"sourceId":11978661,"sourceType":"datasetVersion","datasetId":7235188},{"sourceId":11982868,"sourceType":"datasetVersion","datasetId":7536393},{"sourceId":219146636,"sourceType":"kernelVersion"},{"sourceId":227202990,"sourceType":"kernelVersion"}],"dockerImageVersionId":30919,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"<h1 style=\"color: #6cb4e4;  text-align: center;  padding: 0.25em;  border-top: solid 2.5px #6cb4e4;  border-bottom: solid 2.5px #6cb4e4;  background: -webkit-repeating-linear-gradient(-45deg, #f0f8ff, #f0f8ff 3px,#e9f4ff 3px, #e9f4ff 7px);  background: repeating-linear-gradient(-45deg, #f0f8ff, #f0f8ff 3px,#e9f4ff 3px, #e9f4ff 7px);height:45px;\">\n<b>\nOnly Submission(LoadLocalTrainModel)\n</b></h1> ","metadata":{}},{"cell_type":"markdown","source":"<h1 style=\"color: #6cb4e4;  text-align: center;  padding: 0.25em;  border-top: solid 2.5px #6cb4e4;  border-bottom: solid 2.5px #6cb4e4;  background: -webkit-repeating-linear-gradient(-45deg, #f0f8ff, #f0f8ff 3px,#e9f4ff 3px, #e9f4ff 7px);  background: repeating-linear-gradient(-45deg, #f0f8ff, #f0f8ff 3px,#e9f4ff 3px, #e9f4ff 7px);height:45px;\">\n<b>\nInference Pipeline\n</b></h1> ","metadata":{}},{"cell_type":"markdown","source":"## **》》》 [IMPORTANT] Env Params**","metadata":{}},{"cell_type":"code","source":"\"\"\" Train Model \"\"\"\nmodel_path = \"/kaggle/input/yolo11n-byu/yolo/yolo11x_fold1.pt\"\n\n\"\"\" [IMPORTANT]\n* This parameter has a significant impact on the value of LB since it is the threshold for the prediction score inferred by the model.\n* In my experiments, 0.5 to 0.55 is optimal for local CV, but when submitting, 0.35 to 0.45 seems to give better results, so there is a difference.\n\"\"\"\nYOLO_THR = 0.05\nGNN_THR = 0.05\n\nMAX_DETECTIONS_PER_TOMO = 3\nNMS_IOU_THRESHOLD = 0.2\nCONCENTRATION = 1\nBATCH_SIZE = 8","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-28T11:45:58.335055Z","iopub.execute_input":"2025-05-28T11:45:58.335326Z","iopub.status.idle":"2025-05-28T11:45:58.340126Z","shell.execute_reply.started":"2025-05-28T11:45:58.335304Z","shell.execute_reply":"2025-05-28T11:45:58.339176Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **》》》 Ultralytics Offline Install**(v8.3.88[2025/03/11 ReleaseVersion])","metadata":{}},{"cell_type":"code","source":"!pip install /kaggle/input/pip-install-pyg/torch_spline_conv-1.2.2+pt25cu124-cp310-cp310-linux_x86_64.whl\n!pip install /kaggle/input/pip-install-pyg/torch_sparse-0.6.18+pt25cu124-cp310-cp310-linux_x86_64.whl\n!pip install /kaggle/input/pip-install-pyg/pyg_lib-0.4.0+pt25cu124-cp310-cp310-linux_x86_64.whl\n!pip install /kaggle/input/pip-install-pyg/torch_cluster-1.6.3+pt25cu124-cp310-cp310-linux_x86_64.whl\n!pip install /kaggle/input/pip-install-pyg/torch_geometric-2.6.1-py3-none-any.whl","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-28T11:45:58.341416Z","iopub.execute_input":"2025-05-28T11:45:58.341723Z","iopub.status.idle":"2025-05-28T11:46:17.155022Z","shell.execute_reply.started":"2025-05-28T11:45:58.341691Z","shell.execute_reply":"2025-05-28T11:46:17.154192Z"},"_kg_hide-output":true,"_kg_hide-input":true,"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"[INFO]\n* This notebookinstall Ultralytics v8.3.88(2025/03/11 ReleaseVersion)\n  Can use YOLO12 is latest family version. \n* If you need a newer version, you can make it available by running and attaching the notebook.\n  https://www.kaggle.com/code/hideyukizushi/ultralytics-offlineinstall-yolo12-weights\n\"\"\"\n!tar xfvz /kaggle/input/ultralytics-offlineinstall-yolo12-weights/archive.tar.gz\n!pip install --no-index --find-links=./packages ultralytics\n!rm -rf ./packages","metadata":{"trusted":true,"_kg_hide-input":true,"scrolled":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2025-05-28T11:46:17.156429Z","iopub.execute_input":"2025-05-28T11:46:17.156727Z","iopub.status.idle":"2025-05-28T11:47:08.041828Z","shell.execute_reply.started":"2025-05-28T11:46:17.156704Z","shell.execute_reply":"2025-05-28T11:47:08.040734Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **》》》 Import Libs**","metadata":{}},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\nimport torch\nimport cv2\nfrom tqdm.notebook import tqdm\nfrom ultralytics import YOLO\nimport threading\nimport time\nfrom contextlib import nullcontext\nfrom concurrent.futures import ThreadPoolExecutor\nimport sys\nimport gc\nsys.path.append('/kaggle/input/yolo-gnn-refine-byu')","metadata":{"trusted":true,"_kg_hide-input":false,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2025-05-28T11:47:08.043511Z","iopub.execute_input":"2025-05-28T11:47:08.043864Z","iopub.status.idle":"2025-05-28T11:47:14.524294Z","shell.execute_reply.started":"2025-05-28T11:47:08.043829Z","shell.execute_reply":"2025-05-28T11:47:14.523674Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch_geometric.nn.models import GraphSAGE, GAT, CorrectAndSmooth\nimport torch.nn as nn\nimport torch\nfrom torch_geometric.nn import LayerNorm\nimport torch.nn.functional as F\nfrom torch_geometric.utils import dropout_edge\nfrom torch.distributions import Beta\nimport torchvision.transforms as transforms\nfrom torch_geometric.nn.conv import TransformerConv\nfrom torch_geometric.data import Data\nfrom torch_cluster import knn_graph\nfrom torch_geometric.nn import radius_graph\nfrom torch_geometric.transforms import AddRandomWalkPE\nfrom scipy.spatial import Delaunay\nfrom scipy.spatial.distance import pdist, squareform\nfrom scipy.sparse.csgraph import minimum_spanning_tree\nfrom sklearn.neighbors import KDTree","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-28T11:47:14.525080Z","iopub.execute_input":"2025-05-28T11:47:14.525402Z","iopub.status.idle":"2025-05-28T11:47:19.230749Z","shell.execute_reply.started":"2025-05-28T11:47:14.525383Z","shell.execute_reply":"2025-05-28T11:47:19.229992Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **》》》 Seed Fix**","metadata":{}},{"cell_type":"code","source":"np.random.seed(42)\ntorch.manual_seed(42)","metadata":{"trusted":true,"_kg_hide-input":false,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2025-05-28T11:47:19.231534Z","iopub.execute_input":"2025-05-28T11:47:19.232043Z","iopub.status.idle":"2025-05-28T11:47:19.242966Z","shell.execute_reply.started":"2025-05-28T11:47:19.232019Z","shell.execute_reply":"2025-05-28T11:47:19.242152Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data_path = \"/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/\"\ntest_dir = os.path.join(data_path, \"test\")\nsubmission_path = \"/kaggle/working/submission.csv\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-28T11:47:19.243993Z","iopub.execute_input":"2025-05-28T11:47:19.244291Z","iopub.status.idle":"2025-05-28T11:47:19.258020Z","shell.execute_reply.started":"2025-05-28T11:47:19.244263Z","shell.execute_reply":"2025-05-28T11:47:19.257201Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **》》》 GNN Refinement Tools**","metadata":{}},{"cell_type":"code","source":"class DATA_CONFIG:\n    radius = 30\n    num_samples = 20\n    thr = 10\n    thr_sim = 0.25\n    k = 12\n    w = [0.5, 0.5]\n\n\nclass TRAIN_CONFIG:\n    epochs = 200\n    patience = 10\n    batch_size = 8\n    lr = 5e-4\n    weight_decay = 0\n    ckpt_epoch_freq = 1\n\n\nclass MODEL_CONIFG_1:\n    num_layers = 8\n    hidden_channels = 256\n    pos_weight = 8\n    jk = 'lstm'\n    dropout = 0.3\n    conf = 0.04\n    graph_augment = False\n    use_pe = True\n    walk_length = 8\n    weight_path = '/kaggle/input/yolo11x-byu/checkpoint/knn_graph_12_8_4.pth'\n\nclass MODEL_CONIFG_2:\n    num_layers = 8\n    hidden_channels = 256\n    pos_weight = 8\n    jk = 'lstm'\n    dropout = 0.3\n    conf = 0.05\n    graph_augment = False\n    use_pe = True\n    walk_length = 8\n    weight_path = '/kaggle/input/yolo11x-byu/checkpoint/delaunay_graph_15_8_4.pth'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-28T11:47:19.260369Z","iopub.execute_input":"2025-05-28T11:47:19.260614Z","iopub.status.idle":"2025-05-28T11:47:19.273078Z","shell.execute_reply.started":"2025-05-28T11:47:19.260584Z","shell.execute_reply":"2025-05-28T11:47:19.272326Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_feature_map(model):\n    features = {}\n    def make_hook(name):\n        def hook(module, input, output):\n            features[name] = output\n        return hook\n    model.model.model[16].register_forward_hook(make_hook(\"bifpn\"))\n    return features","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-28T11:47:19.274361Z","iopub.execute_input":"2025-05-28T11:47:19.274665Z","iopub.status.idle":"2025-05-28T11:47:19.287037Z","shell.execute_reply.started":"2025-05-28T11:47:19.274637Z","shell.execute_reply":"2025-05-28T11:47:19.286399Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def delaunay_graph(x):\n    points = x.cpu().numpy()  # assuming x is a tensor of shape [num_points, dims]\n    tri = Delaunay(points)\n    edges = set()\n\n    for simplex in tri.simplices:\n        for i in range(len(simplex)):\n            for j in range(i + 1, len(simplex)):\n                edges.add(tuple(sorted((simplex[i], simplex[j]))))\n\n    edge_index = torch.tensor(list(edges), dtype=torch.long).t()\n    return edge_index\n\ndef mst_graph(x):\n    pos_np = x.numpy()\n    dist_matrix = squareform(pdist(pos_np))  # pairwise distances\n    mst = minimum_spanning_tree(dist_matrix)  # sparse MST\n    row, col = mst.nonzero()\n    edge_index = torch.tensor([row, col], dtype=torch.long)\n    return edge_index\n\n\ndef kdtree_graph(x, k):\n    N = x.shape[0]\n    pos_np = x.numpy()\n    tree = KDTree(pos_np)\n    ind = tree.query(pos_np, k = k + 1, return_distance=False)  # k+1 because includes self\n    # Convert neighbors to edge index\n    row_idx = np.repeat(np.arange(N), k)\n    col_idx = ind[:, 1:].reshape(-1)  # exclude self neighbor at index 0\n    edge_index = torch.tensor([row_idx, col_idx], dtype=torch.long)\n    return edge_index","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-28T11:47:19.287811Z","iopub.execute_input":"2025-05-28T11:47:19.288020Z","iopub.status.idle":"2025-05-28T11:47:19.300888Z","shell.execute_reply.started":"2025-05-28T11:47:19.287990Z","shell.execute_reply":"2025-05-28T11:47:19.300065Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def assign_feat(x, feat, tomo_id, vol_size):\n    d, h, w = vol_size\n    fd, fh, fw = feat.shape[0], feat.shape[2], feat.shape[3]\n    z = x[:, 0].astype('int32')\n    y = ((x[:, 1]/h) * fh).astype('int32')\n    x = ((x[:, 2]/w) * fw).astype('int32')\n    extract_feat = feat[z, :, y, x]\n    return extract_feat\n\ndef sample_uniform_3d_ball(info, tomo_id, vol_size, radius=10, num_samples=10):\n    \"\"\"\n    Uniformly sample points in 3D balls around each point without exceeding volume boundary.\n\n    Args:\n        points (np.ndarray): (N, 3) center points\n        radius (float): radius of ball\n        num_samples (int): number of samples per center\n        volume_shape (tuple): (Z, Y, X) shape of 3D volume\n\n    Returns:\n        np.ndarray: (N, num_samples, 3) sampled points clipped to stay in volume\n    \"\"\"\n    Z, Y, X = vol_size\n\n    points = info[:, :-1]\n    N = points.shape[0]\n\n    def uniform_ball(n):\n        vec = np.random.randn(n, 3)\n        vec /= np.linalg.norm(vec, axis=1, keepdims=True)\n        r = np.random.rand(n) ** (1/3)\n        return vec * (r[:, None] * radius)\n\n    result = []\n    np.random.seed(42)\n    for center in points:\n        attempts = 0\n        accepted = []\n\n        # Accept valid samples until we have enough or hit retry limit\n        while len(accepted) < num_samples and attempts < num_samples * 10:\n            samples = uniform_ball(num_samples)\n            candidates = samples + center  # shifted samples\n\n            # Keep only those within volume bounds\n            mask = (\n                (candidates[:, 0] >= 0) & (candidates[:, 0] < Z) &\n                (candidates[:, 1] >= 0) & (candidates[:, 1] < Y) &\n                (candidates[:, 2] >= 0) & (candidates[:, 2] < X)\n            )\n            accepted.extend(candidates[mask])\n            attempts += 1\n\n        # If not enough valid, pad with center point\n        if len(accepted) < num_samples:\n            accepted.extend([center] * (num_samples - len(accepted)))\n\n        result.append(np.array(accepted[:num_samples]))\n\n    result = np.concatenate(result, axis=0)\n    return result\n\ntransform_1 = AddRandomWalkPE(walk_length=MODEL_CONIFG_1.walk_length, attr_name=None)\ntransform_2 = AddRandomWalkPE(walk_length=MODEL_CONIFG_2.walk_length, attr_name=None)\n\ndef extract_tomo(info, feat, tomo_id, radius, num_samples):\n    tomo_dir = os.path.join(test_dir, tomo_id)\n    tomo_names = os.listdir(tomo_dir)\n    D = len(tomo_names)\n    img = cv2.imread(os.path.join(tomo_dir, tomo_names[0]))\n    H, W = img.shape[:2]\n    del img\n    vol_size = [D, H, W]\n    points = sample_uniform_3d_ball(info, tomo_id, vol_size, radius, num_samples)\n    extract_feat = assign_feat(points, feat, tomo_id, vol_size)\n    #convert to torch\n    points = torch.from_numpy(points)\n    extract_feat = torch.from_numpy(extract_feat)\n    batch = torch.zeros(points.shape[0], dtype=torch.int64)\n\n    #multi-graph\n    edge_index_1 = knn_graph(points, k=DATA_CONFIG.k, loop=False)\n    edge_index_2 = delaunay_graph(points)\n    data1= Data(points = points, x = extract_feat, edge_index = edge_index_1, batch = batch)\n    data2= Data(points = points, x = extract_feat, edge_index = edge_index_2, batch = batch)\n\n    if MODEL_CONIFG_1.use_pe:\n        data1 = transform_1(data1)\n\n    if MODEL_CONIFG_2.use_pe:\n        data2 = transform_2(data2)\n    return data1, data2","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-28T11:47:19.301709Z","iopub.execute_input":"2025-05-28T11:47:19.301960Z","iopub.status.idle":"2025-05-28T11:47:19.317030Z","shell.execute_reply.started":"2025-05-28T11:47:19.301935Z","shell.execute_reply":"2025-05-28T11:47:19.316229Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class GraphModel(nn.Module):\n    def __init__(self, cfg):\n        super().__init__()\n        self.gnn = GraphSAGE(384 + cfg.walk_length if cfg.use_pe else 384,\n                             num_layers= cfg.num_layers,\n                             hidden_channels=cfg.hidden_channels,\n                             out_channels=1,\n                             jk=cfg.jk,\n                             dropout=cfg.dropout,\n                             norm=LayerNorm(cfg.hidden_channels),\n                             )\n    def forward(self, data):\n        x, edge_index, batch = data.x, data.edge_index, data.batch\n        edge_index = edge_index.cuda()\n        batch = batch.cuda()\n\n        if edge_index.shape[0] == 0:\n            edge_index = torch.tensor([[0, 0]]).T.cuda()\n\n        x = x.cuda()\n        logits = self.gnn(x, edge_index, batch=batch)\n        return logits","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-28T11:47:19.317814Z","iopub.execute_input":"2025-05-28T11:47:19.318092Z","iopub.status.idle":"2025-05-28T11:47:19.333303Z","shell.execute_reply.started":"2025-05-28T11:47:19.318066Z","shell.execute_reply":"2025-05-28T11:47:19.332554Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **》》》 Inference&Submission**","metadata":{}},{"cell_type":"markdown","source":"* GPU Init","metadata":{}},{"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\n\ndevice = 'cuda:0' if torch.cuda.is_available() else 'cpu'\nif device.startswith('cuda'):\n    # Set CUDA optimization flags\n    torch.backends.cudnn.benchmark = True\n    torch.backends.cudnn.deterministic = False\n    torch.backends.cuda.matmul.allow_tf32 = True  # Allow TF32 on Ampere GPUs\n    torch.backends.cudnn.allow_tf32 = True\n    \n    # Print GPU info\n    gpu_name = torch.cuda.get_device_name(0)\n    gpu_mem = torch.cuda.get_device_properties(0).total_memory / 1e9  # Convert to GB\n    print(f\"Using GPU: {gpu_name} with {gpu_mem:.2f} GB memory\")\n    \n    # Get available GPU memory and set batch size accordingly\n    free_mem = gpu_mem - torch.cuda.memory_allocated(0) / 1e9\n    BATCH_SIZE = max(8, min(32, int(free_mem * 4)))  # 4 images per GB as rough estimate\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  # Reduce batch size for CPU","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-28T11:47:19.334174Z","iopub.execute_input":"2025-05-28T11:47:19.334441Z","iopub.status.idle":"2025-05-28T11:47:19.375811Z","shell.execute_reply.started":"2025-05-28T11:47:19.334415Z","shell.execute_reply":"2025-05-28T11:47:19.375173Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"* Inference","metadata":{}},{"cell_type":"code","source":"def normalize_slice(slice_data):\n    \"\"\"\n    Normalize slice data using 2nd and 98th percentiles for better contrast\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            # Try with PIL as fallback\n            img = np.array(Image.open(path))\n        images.append(img)\n    return images\n\ndef process_tomogram(tomo_id, model, model_gnn_1, model_gnn_2, 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    \n    # Get all slice files for this tomogram\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    # Apply CONCENTRATION to reduce the number of slices processed\n    # This will process approximately CONCENTRATION fraction of all slices\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\n    all_detections = []\n    all_detections_tmp = []\n    all_features_bifpn = []\n    if device.startswith('cuda'):\n        streams = [torch.cuda.Stream() for _ in range(min(4, BATCH_SIZE))]\n    else:\n        streams = [None]\n    \n    # Variables for preloading\n    next_batch_thread = None\n    next_batch_images = None\n    \n    # Process slices in batches\n    for batch_start in tqdm(range(0, len(slice_files), BATCH_SIZE)):\n        # Wait for previous preload thread if it exists\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        # Start preloading next batch\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        # Split batch across streams for parallel processing\n        sub_batches = np.array_split(batch_files, len(streams))\n        sub_batch_results = []\n        \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) if stream and device.startswith('cuda') else nullcontext():\n                # Process sub-batch\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                # Run inference with profiling\n                with GPUProfiler(f\"Inference batch {i+1}/{len(sub_batches)}\"):\n                    with torch.no_grad():\n                        features = get_feature_map(model)\n                        sub_results = model.predict(sub_batch_paths, imgsz=960, verbose=False)\n                bifpn_feat = features['bifpn'].cpu().numpy()\n                all_features_bifpn.append(bifpn_feat)\n                del bifpn_feat\n                \n                # Process each result in this sub-batch\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                            if confidence >= GNN_THR:\n                                # Get bounding box coordinates\n                                x1, y1, x2, y2 = boxes.xyxy[box_idx].cpu().numpy()\n                                \n                                # Calculate center coordinates\n                                x_center = (x1 + x2) / 2\n                                y_center = (y1 + y2) / 2\n\n                                all_detections_tmp.append([round(sub_batch_slice_nums[j]),\n                                                           round(y_center),\n                                                           round(x_center),\n                                                           float(confidence)\n                                                          ])\n                                if confidence >= YOLO_THR:\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        \n        # Synchronize streams\n        if device.startswith('cuda'):\n            torch.cuda.synchronize()\n    \n    # Clean up thread if still running\n    if next_batch_thread is not None:\n        next_batch_thread.join()\n\n    if len(all_detections)==0:\n        return {\n            'tomo_id': tomo_id,\n            'Motor axis 0': -1,\n            'Motor axis 1': -1,\n            'Motor axis 2': -1\n        }\n\n    all_detections_tmp = np.array(all_detections_tmp)\n    all_features_bifpn = np.concatenate(all_features_bifpn, axis=0)\n    data1, data2 = extract_tomo(all_detections_tmp, all_features_bifpn, tomo_id, 30, 20)\n    del all_detections_tmp\n    del all_features_bifpn\n    gc.collect()\n\n    #forward\n    logits_1 = model_gnn_1(data1)\n    logits_2 = model_gnn_2(data2)\n    gnn_pred_1 = logits_1.sigmoid()[:, 0].cpu() #(N, )\n    gnn_pred_2 = logits_2.sigmoid()[:, 0].cpu()\n    \n    gnn_pred = torch.cat([gnn_pred_1[gnn_pred_1>MODEL_CONIFG_1.conf],\n                          gnn_pred_2[gnn_pred_2>MODEL_CONIFG_2.conf]], dim=0).cpu()\n\n    try:\n        points_1 = data1.points[gnn_pred_1>MODEL_CONIFG_1.conf]\n        points_2 = data2.points[gnn_pred_2>MODEL_CONIFG_2.conf]\n        points = torch.cat([points_1, points_2], dim=0)\n        pred_idx = gnn_pred.argmax()\n        point_pred = points[pred_idx].numpy() #(3, )\n        # point_pred = postprocessing(points)\n        return {\n            'tomo_id': tomo_id,\n            'Motor axis 0': point_pred[0],\n            'Motor axis 1': point_pred[1],\n            'Motor axis 2': point_pred[2]\n        }\n    except:\n        return {\n            'tomo_id': tomo_id,\n            'Motor axis 0': -1,\n            'Motor axis 1': -1,\n            'Motor axis 2': -1\n        }\n    \n\ndef debug_image_loading(tomo_id):\n    \"\"\"\n    Debug function to check image loading\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    if not slice_files:\n        print(f\"No image files found in {tomo_dir}\")\n        return\n        \n    sample_file = slice_files[len(slice_files)//2]  # Middle slice\n    img_path = os.path.join(tomo_dir, sample_file)\n    \n    # Try different loading methods\n    try:\n        # Method 1: PIL\n        img_pil = Image.open(img_path)\n        img_array_pil = np.array(img_pil)\n        \n        # Method 2: OpenCV\n        img_cv2 = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n        # print(f\"OpenCV Image shape: {img_cv2.shape}, dtype: {img_cv2.dtype}\")\n        \n        # Method 3: Convert to RGB\n        img_rgb = cv2.cvtColor(cv2.imread(img_path), cv2.COLOR_BGR2RGB)\n    except Exception as e:\n        print(f\"Error loading image {img_path}: {e}\")\n        \n    # Also test with YOLO's built-in loader\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}\")\n\n\ndef load_gnn_model(cfg):\n    model_gnn = GraphModel(cfg)\n    model_gnn.load_state_dict(torch.load(cfg.weight_path))\n    model_gnn.to(device)\n    model_gnn.eval()\n    return model_gnn\n\n\ndef 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    \n    if test_tomos:\n        debug_image_loading(test_tomos[0])\n    \n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n    print('load yolo')\n    model = YOLO(model_path)\n    model.to(device)\n\n    print('load gnn')\n    model_gnn_1 = load_gnn_model(MODEL_CONIFG_1)\n    model_gnn_2 = load_gnn_model(MODEL_CONIFG_2)\n\n    print('start inference')\n    # Additional optimizations for inference\n    if device.startswith('cuda'):\n        # Fuse conv and bn layers for faster inference\n        model.fuse()\n        \n        # Enable model half precision (FP16) if on compatible GPU\n        if torch.cuda.get_device_capability(0)[0] >= 7:  # Volta or newer\n            model.model.half()\n    \n    # Process tomograms with parallelization\n    results = []\n    motors_found = 0\n\n    with ThreadPoolExecutor(max_workers=1) as executor:\n        future_to_tomo = {}\n        \n        # Submit all tomograms for processing\n        for i, tomo_id in enumerate(test_tomos, 1):\n            future = executor.submit(process_tomogram, tomo_id, model, model_gnn_1, model_gnn_2, i, total_tomos)\n            future_to_tomo[future] = tomo_id\n        \n        # Process completed futures as they complete\n        for future in future_to_tomo:\n            tomo_id = future_to_tomo[future]\n            \n            try:\n                # Clear CUDA cache between tomograms\n                if torch.cuda.is_available():\n                    torch.cuda.empty_cache()\n                    \n                result = future.result()\n                results.append(result)\n                \n                # Update motors found count\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: \"\n                          f\"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                    \n                print(f\"Current detection rate: {motors_found}/{len(results)} ({motors_found/len(results)*100:.1f}%)\")\n            \n            except Exception as e:\n                print(f\"Error processing {tomo_id}: {e}\")\n                # Create a default entry for failed tomograms\n                results.append({\n                    'tomo_id': tomo_id,\n                    'Motor axis 0': -1,\n                    'Motor axis 1': -1,\n                    'Motor axis 2': -1\n                })\n    \n    # Create submission dataframe\n    submission_df = pd.DataFrame(results)\n    \n    # Ensure proper column order\n    submission_df = submission_df[['tomo_id', 'Motor axis 0', 'Motor axis 1', 'Motor axis 2']]\n    \n    # Save the submission file\n    submission_df.to_csv(submission_path, index=False)\n    print(\"=\"*50)\n    print(\"= Submission preview:\")\n    print(\"=\"*50)\n    print(submission_df.head())\n    \n    return submission_df","metadata":{"trusted":true,"scrolled":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2025-05-28T11:47:19.376517Z","iopub.execute_input":"2025-05-28T11:47:19.376774Z","iopub.status.idle":"2025-05-28T11:47:19.400097Z","shell.execute_reply.started":"2025-05-28T11:47:19.376755Z","shell.execute_reply":"2025-05-28T11:47:19.399426Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission = generate_submission()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-28T11:47:19.400795Z","iopub.execute_input":"2025-05-28T11:47:19.401021Z","iopub.status.idle":"2025-05-28T11:50:21.801536Z","shell.execute_reply.started":"2025-05-28T11:47:19.400990Z","shell.execute_reply":"2025-05-28T11:50:21.800745Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}