{"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":11241677,"sourceType":"datasetVersion","datasetId":7022424,"isSourceIdPinned":true},{"sourceId":11382537,"sourceType":"datasetVersion","datasetId":7061153},{"sourceId":11454081,"sourceType":"datasetVersion","datasetId":7176774},{"sourceId":219146636,"sourceType":"kernelVersion"}],"dockerImageVersionId":30919,"isInternetEnabled":false,"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,"execution":{"iopub.status.busy":"2025-04-18T00:16:44.995993Z","iopub.execute_input":"2025-04-18T00:16:44.996345Z","iopub.status.idle":"2025-04-18T00:17:04.595025Z","shell.execute_reply.started":"2025-04-18T00:16:44.996317Z","shell.execute_reply":"2025-04-18T00:17:04.594166Z"},"_kg_hide-output":true,"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\nsys.path.append('/kaggle/input/gnn-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\nimport networkx as nx\nimport pickle","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-18T00:17:42.119889Z","iopub.execute_input":"2025-04-18T00:17:42.120270Z","iopub.status.idle":"2025-04-18T00:17:42.144345Z","shell.execute_reply.started":"2025-04-18T00:17:42.120238Z","shell.execute_reply":"2025-04-18T00:17:42.143493Z"}},"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 = 4\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 = 8\n    dropout = 0.3\n    conf = 0.5\n    conf_fixed = 0.5\n    graph_augment = True\n    gnn_type = 'GraphSage'\n    loss_type = 'bce'\n    use_stn = False\n    jk = 'lstm'\n    pretrained_weights = '/kaggle/input/gnn-byu/checkpoint/256_8_0.3_8_True_True_1_GraphSage_bce_False_lstm_True_True.pth' \n    use_attn = True\n    use_pretrained=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-18T00:17:47.130586Z","iopub.execute_input":"2025-04-18T00:17:47.131561Z","iopub.status.idle":"2025-04-18T00:17:47.137131Z","shell.execute_reply.started":"2025-04-18T00:17:47.131527Z","shell.execute_reply":"2025-04-18T00:17:47.136355Z"}},"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-18T00:17:47.307711Z","iopub.execute_input":"2025-04-18T00:17:47.308005Z","iopub.status.idle":"2025-04-18T00:17:47.313844Z","shell.execute_reply.started":"2025-04-18T00:17:47.307980Z","shell.execute_reply":"2025-04-18T00:17:47.313138Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model =  GraphModel(MODEL_CONFIG, False)\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-18T00:17:49.777659Z","iopub.execute_input":"2025-04-18T00:17:49.777953Z","iopub.status.idle":"2025-04-18T00:17:51.389310Z","shell.execute_reply.started":"2025-04-18T00:17:49.777930Z","shell.execute_reply":"2025-04-18T00:17:51.388455Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"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-18T00:18:03.562863Z","iopub.execute_input":"2025-04-18T00:18:03.563218Z","iopub.status.idle":"2025-04-18T00:18:03.579835Z","shell.execute_reply.started":"2025-04-18T00:18:03.563178Z","shell.execute_reply":"2025-04-18T00:18:03.578864Z"}},"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    #use highest average probs to select cluster\n    return voted_keypoints[idx]\n\ndef normalize_score(scores, threshold=MODEL_CONFIG.conf):\n    scale = 0.5 / (1.0 - threshold)  # e.g., 0.5 / 0.02 = 25\n    offset = 0.5 - threshold * scale\n    return torch.tensor([score * scale + offset for score in scores])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T00:18:03.743674Z","iopub.execute_input":"2025-04-18T00:18:03.744006Z","iopub.status.idle":"2025-04-18T00:18:03.752512Z","shell.execute_reply.started":"2025-04-18T00:18:03.743977Z","shell.execute_reply":"2025-04-18T00:18:03.751677Z"}},"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                        drop_last = True\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            prob = normalize_score(prob)\n            pred = prob > MODEL_CONFIG.conf_fixed\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                    z_kp = torch.tensor([z]*kp_j.shape[0])\n                    kp_j = torch.cat([z_kp[:, None], kp_j], dim=-1) #(N, 3)\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        probs_all = torch.cat(probs_all, dim=0)\n        pred_kp = postprocessing(kps_all)\n        print(pred_kp)\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        print(-1)\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-18T00:33:42.197094Z","iopub.execute_input":"2025-04-18T00:33:42.197413Z","iopub.status.idle":"2025-04-18T00:35:19.901294Z","shell.execute_reply.started":"2025-04-18T00:33:42.197391Z","shell.execute_reply":"2025-04-18T00:35:19.900499Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_sub.to_csv('submission.csv')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T00:35:24.671654Z","iopub.execute_input":"2025-04-18T00:35:24.671972Z","iopub.status.idle":"2025-04-18T00:35:24.677228Z","shell.execute_reply.started":"2025-04-18T00:35:24.671949Z","shell.execute_reply":"2025-04-18T00:35:24.676241Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_sub","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T00:35:29.721291Z","iopub.execute_input":"2025-04-18T00:35:29.721603Z","iopub.status.idle":"2025-04-18T00:35:29.730592Z","shell.execute_reply.started":"2025-04-18T00:35:29.721580Z","shell.execute_reply":"2025-04-18T00:35:29.729774Z"}},"outputs":[],"execution_count":null}]}