{"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,"sourceType":"competition"},{"sourceId":11279996,"sourceType":"datasetVersion","datasetId":7022424},{"sourceId":11316406,"sourceType":"datasetVersion","datasetId":7061153},{"sourceId":219146636,"sourceType":"kernelVersion"}],"dockerImageVersionId":30919,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## Install Packages","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,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2025-04-07T23:52:07.403011Z","iopub.execute_input":"2025-04-07T23:52:07.403330Z","iopub.status.idle":"2025-04-07T23:52:25.462664Z","shell.execute_reply.started":"2025-04-07T23:52:07.403302Z","shell.execute_reply":"2025-04-07T23:52:25.461833Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Submission","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport torch\nimport torch.nn as nn\nfrom torch_geometric.data import Dataset, Data\nfrom torch_geometric.loader import DataLoader\nimport pandas as pd\nimport sys\n# sys.path.append('/kaggle/input/stn-gnn-byu')\nsys.path.append('/kaggle/input/stn-gnn-v3-byu')\nfrom model import *\nfrom preprocessing_ import *\nimport glob\nimport os\nfrom tqdm import tqdm\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\nimport matplotlib.pyplot as plt\nimport cv2\nimport gc\nimport hdbscan\nfrom collections import defaultdict","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-07T23:52:25.464079Z","iopub.execute_input":"2025-04-07T23:52:25.464315Z","iopub.status.idle":"2025-04-07T23:52:34.424591Z","shell.execute_reply.started":"2025-04-07T23:52:25.464293Z","shell.execute_reply":"2025-04-07T23:52:34.423956Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class TRAIN_CONFIG:\n    epochs = 1000\n    lr = 5e-4\n    ckpt_epoch_freq = 2\n    pos_weight = 64\n    weight_decay = 1e-4\n    patience = 10\n    class_weight = [1, pos_weight]\n    alpha = 0.25\n    gamma = 2\n    beta = 0.1\n    augment = False\n    mix_beta = 1\n    mixup_p = 1\n\n\nclass MODEL_CONFIG:\n    hidden_channels = 256\n    num_layers = 10\n    dropout = 0.3\n    conf = 0.75\n    graph_augment = False\n    gnn_type = 'GraphSage'\n    loss_type = 'bce'\n    use_stn = False\n    jk = 'lstm'\n    pretrained_weights = '/kaggle/input/stn-gnn-v3-byu/checkpoint/256_10_0.3_32_False_True_1_GraphSage_bce_False_lstm_True.pth' \n    use_attn = True\n\n\nclass DATA_CONFIG:\n    train_path = './graph_data_byu/train'\n    val_path = './graph_data_byu/val'\n    num_workers = 4\n    batch_size = 8\n    jump_step = 8\n    min_cluster_size = 4\n    min_samples = 6","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-07T23:52:34.426009Z","iopub.execute_input":"2025-04-07T23:52:34.426488Z","iopub.status.idle":"2025-04-07T23:52:34.431440Z","shell.execute_reply.started":"2025-04-07T23:52:34.426467Z","shell.execute_reply":"2025-04-07T23:52:34.430651Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class TomoGraphDataset(Dataset):\n    def __init__(self, data_path):\n        super().__init__()\n        self.data_path = data_path\n        self.data_names = sorted(os.listdir(data_path))[::DATA_CONFIG.jump_step]\n        \n    def len(self):\n        return len(self.data_names)\n\n    def get(self, idx):\n        data_name = self.data_names[idx]\n        data_path = os.path.join(self.data_path, data_name)\n        patch, kp, img_size = point_feature_extractor(data_path)\n        H, W = img_size\n        edge_index = graph_construction(kp)\n        kp[:, 1] = kp[:, 1] * (H/640)\n        kp[:, 0] = kp[:, 0] * (W/640)\n        patch = torch.tensor(patch, dtype=torch.float32)[:, None, ...]\n        edge_index = torch.tensor(edge_index)\n        kp = torch.tensor(kp, dtype=torch.int32)\n        return Data(x=patch, edge_index=edge_index, y=kp)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-07T23:52:34.432787Z","iopub.execute_input":"2025-04-07T23:52:34.433092Z","iopub.status.idle":"2025-04-07T23:52:34.691533Z","shell.execute_reply.started":"2025-04-07T23:52:34.433064Z","shell.execute_reply":"2025-04-07T23:52:34.690826Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model =  GraphModel(MODEL_CONFIG)\nmodel.load_state_dict(torch.load(MODEL_CONFIG.pretrained_weights))\nmodel = model.cuda()\nmodel.eval()\nprint()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-07T23:52:34.692366Z","iopub.execute_input":"2025-04-07T23:52:34.692602Z","iopub.status.idle":"2025-04-07T23:52:35.893655Z","shell.execute_reply.started":"2025-04-07T23:52:34.692581Z","shell.execute_reply":"2025-04-07T23:52:35.892753Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"tomo_path = '/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/test'\ntomo_names = sorted(os.listdir(tomo_path))\n\ndf_ex = pd.read_csv('/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/sample_submission.csv')\ndf_sub = pd.DataFrame(columns=df_ex.columns)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-07T23:52:35.894622Z","iopub.execute_input":"2025-04-07T23:52:35.894971Z","iopub.status.idle":"2025-04-07T23:52:35.913762Z","shell.execute_reply.started":"2025-04-07T23:52:35.894938Z","shell.execute_reply":"2025-04-07T23:52:35.913191Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def postprocessing(kps):\n    clusterer = hdbscan.HDBSCAN(min_cluster_size=DATA_CONFIG.min_cluster_size,\n                                min_samples = DATA_CONFIG.min_samples)\n    labels = clusterer.fit_predict(kps)\n    clusters = defaultdict(list)\n    for label, point in zip(labels, kps):\n        if label != -1: \n            clusters[label].append(point)\n    voted_keypoints = [\n        tuple(np.mean(points, axis=0))\n        for points in clusters.values()\n    ]\n    cluster_size = [len(points) for points in clusters.values()]\n    idx = np.argmax(cluster_size)\n    print(cluster_size)\n    print(idx)\n    return voted_keypoints[idx]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-07T23:52:35.914537Z","iopub.execute_input":"2025-04-07T23:52:35.914830Z","iopub.status.idle":"2025-04-07T23:52:35.919663Z","shell.execute_reply.started":"2025-04-07T23:52:35.914786Z","shell.execute_reply":"2025-04-07T23:52:35.918939Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for k, tomo_name in enumerate(tomo_names):\n    tomo_path_ = os.path.join(tomo_path, tomo_name)\n    ds = TomoGraphDataset(tomo_path_)\n    loader = DataLoader(ds,\n                        batch_size=DATA_CONFIG.batch_size,\n                        num_workers=DATA_CONFIG.num_workers,\n                        shuffle = False\n                       )\n    probs = []\n    preds = []\n    kps = []\n    kps_all = []\n    probs_all = []\n    with torch.no_grad():\n        for i, batch in enumerate(tqdm(loader)):\n            z = i * DATA_CONFIG.batch_size* DATA_CONFIG.jump_step\n            x = batch.x\n            kp = batch.y\n            logits = model(x, batch)\n            prob = logits.sigmoid()[:, 0].cpu()\n            pred = prob > MODEL_CONFIG.conf\n            del logits \n            torch.cuda.empty_cache()\n            for j in range(DATA_CONFIG.batch_size):\n                try:\n                    z = i * DATA_CONFIG.batch_size* DATA_CONFIG.jump_step + j * DATA_CONFIG.jump_step\n                    pred_j = pred[batch.batch==j]\n                    prob_j = prob[batch.batch==j]\n                    kp_j = kp[batch.batch==j]\n                    best_idx = prob_j.argmax()\n                    pred_best = pred_j[best_idx]\n                    kps_all.append(kp_j[pred_j==1])\n                    probs_all.append(prob_j[pred_j==1])\n                    if pred_best == 0:\n                        kps.append([-1, -1, -1])\n                        probs.append(torch.tensor(0.))\n                        preds.append(torch.tensor(0))\n                    else:\n                        kp_best = kp_j[best_idx]\n                        prob_max = prob_j[best_idx]\n                        kps.append([z, kp_best[0], kp_best[1]])\n                        probs.append(prob_max)\n                        preds.append(pred_best)\n                except:\n                    kps.append([-1, -1, -1])\n                    probs.append(torch.tensor(0.))\n                    preds.append(torch.tensor(0))\n    try:\n        idx = torch.argmax(torch.tensor(probs))\n        pred_z = torch.tensor(kps)[idx].numpy()[0]\n        kps_all = torch.cat(kps_all, dim=0)\n        pred_kp = postprocessing(kps_all)\n        pred_kp = np.array([pred_z, pred_kp[0], pred_kp[1]], dtype='int64')\n        df_sub.loc[k, df_sub.columns[1:]] = pred_kp\n        df_sub.loc[k, df_sub.columns[0]] = tomo_name\n    except:\n        df_sub.loc[k, df_sub.columns[1:]] = [-1, -1, -1]\n        df_sub.loc[k, df_sub.columns[0]] = tomo_name\n    gc.collect()\n    del probs\n    del preds\n    del kps","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-07T23:52:35.920490Z","iopub.execute_input":"2025-04-07T23:52:35.920762Z","iopub.status.idle":"2025-04-07T23:54:34.773987Z","shell.execute_reply.started":"2025-04-07T23:52:35.920734Z","shell.execute_reply":"2025-04-07T23:54:34.773077Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_sub","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-07T23:54:34.776329Z","iopub.execute_input":"2025-04-07T23:54:34.776558Z","iopub.status.idle":"2025-04-07T23:54:34.804570Z","shell.execute_reply.started":"2025-04-07T23:54:34.776539Z","shell.execute_reply":"2025-04-07T23:54:34.803662Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"im = cv2.imread('/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/test/tomo_01a877/slice_0148.jpg')\n#im = cv2.imread('/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/test/tomo_00e047/slice_0168.jpg')\nim = im .transpose(1, 0, 2)\nim = cv2.cvtColor(im, cv2.COLOR_BGR2GRAY)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-07T23:54:34.805497Z","iopub.execute_input":"2025-04-07T23:54:34.805718Z","iopub.status.idle":"2025-04-07T23:54:34.847483Z","shell.execute_reply.started":"2025-04-07T23:54:34.805700Z","shell.execute_reply":"2025-04-07T23:54:34.846626Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.imshow(im, cmap='gray')\nplt.scatter(pred_kp[1], pred_kp[2], c='red')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-07T23:54:34.848444Z","iopub.execute_input":"2025-04-07T23:54:34.848751Z","iopub.status.idle":"2025-04-07T23:54:35.159959Z","shell.execute_reply.started":"2025-04-07T23:54:34.848723Z","shell.execute_reply":"2025-04-07T23:54:35.159045Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.imshow(im, cmap='gray')\nplt.scatter(kps_all[:,0], kps_all[:, 1], c='red')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-07T23:54:35.160722Z","iopub.execute_input":"2025-04-07T23:54:35.161059Z","iopub.status.idle":"2025-04-07T23:54:35.405902Z","shell.execute_reply.started":"2025-04-07T23:54:35.161028Z","shell.execute_reply":"2025-04-07T23:54:35.404959Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}