{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceType":"competition","sourceId":7878,"databundleVersionId":46689},{"sourceType":"datasetVersion","sourceId":15305173,"datasetId":9789774,"databundleVersionId":16209689}],"dockerImageVersionId":31328,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## Imports","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom sklearn.neighbors import NearestNeighbors\nimport torch\nimport torch.nn as nn","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-23T10:17:08.925072Z","iopub.execute_input":"2026-03-23T10:17:08.925477Z","iopub.status.idle":"2026-03-23T10:17:15.595624Z","shell.execute_reply.started":"2026-03-23T10:17:08.925441Z","shell.execute_reply":"2026-03-23T10:17:15.594910Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pip install torch_geometric > /dev/null","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-23T10:17:18.191764Z","iopub.execute_input":"2026-03-23T10:17:18.192569Z","iopub.status.idle":"2026-03-23T10:17:24.837341Z","shell.execute_reply.started":"2026-03-23T10:17:18.192514Z","shell.execute_reply":"2026-03-23T10:17:24.836066Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch_geometric.data import Data\n\nfrom torch_geometric.nn import GCNConv","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-23T10:17:27.196800Z","iopub.execute_input":"2026-03-23T10:17:27.197717Z","iopub.status.idle":"2026-03-23T10:17:45.462378Z","shell.execute_reply.started":"2026-03-23T10:17:27.197674Z","shell.execute_reply":"2026-03-23T10:17:45.461525Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Load and analyse data","metadata":{}},{"cell_type":"code","source":"hits = pd.read_csv(\"/kaggle/input/datasets/hrishikeshthakur7/trackml/train_100_events/event000001000-hits.csv\")\ntruth = pd.read_csv(\"/kaggle/input/datasets/hrishikeshthakur7/trackml/train_100_events/event000001000-truth.csv\")\n\ndf = hits.merge(truth, on=\"hit_id\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-23T10:17:48.852956Z","iopub.execute_input":"2026-03-23T10:17:48.853512Z","iopub.status.idle":"2026-03-23T10:17:49.255303Z","shell.execute_reply.started":"2026-03-23T10:17:48.853475Z","shell.execute_reply":"2026-03-23T10:17:49.254289Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-23T10:17:51.808219Z","iopub.execute_input":"2026-03-23T10:17:51.808642Z","iopub.status.idle":"2026-03-23T10:17:51.838794Z","shell.execute_reply.started":"2026-03-23T10:17:51.808599Z","shell.execute_reply":"2026-03-23T10:17:51.837728Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"particle = df[df['particle_id'] == df['particle_id'].unique()[0]]\n\nplt.scatter(particle['x'], particle['y'])\nplt.title(\"Single Particle Track (2D)\")\nplt.xlabel(\"x\")\nplt.ylabel(\"y\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-23T10:17:55.261861Z","iopub.execute_input":"2026-03-23T10:17:55.262312Z","iopub.status.idle":"2026-03-23T10:17:55.535533Z","shell.execute_reply.started":"2026-03-23T10:17:55.262274Z","shell.execute_reply":"2026-03-23T10:17:55.534506Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Build Graph","metadata":{}},{"cell_type":"code","source":"coords = df[['x', 'y', 'z']].values\n\nnbrs = NearestNeighbors(n_neighbors=10).fit(coords)\ndistances, indices = nbrs.kneighbors(coords)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-23T10:18:28.708779Z","iopub.execute_input":"2026-03-23T10:18:28.709684Z","iopub.status.idle":"2026-03-23T10:18:29.495851Z","shell.execute_reply.started":"2026-03-23T10:18:28.709642Z","shell.execute_reply":"2026-03-23T10:18:29.494972Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Create Edges","metadata":{}},{"cell_type":"code","source":"edge_index = []\n\nfor i, neighbors in enumerate(indices):\n    for j in neighbors:\n        edge_index.append([i, j])\n\nedge_index = np.array(edge_index).T","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-23T10:18:32.576764Z","iopub.execute_input":"2026-03-23T10:18:32.577830Z","iopub.status.idle":"2026-03-23T10:18:34.973951Z","shell.execute_reply.started":"2026-03-23T10:18:32.577787Z","shell.execute_reply":"2026-03-23T10:18:34.973085Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Create Labels","metadata":{}},{"cell_type":"code","source":"labels = []\n\nfor i, j in edge_index.T:\n    if df.iloc[i]['particle_id'] == df.iloc[j]['particle_id']:\n        labels.append(1)  # same track\n    else:\n        labels.append(0)\n\nlabels = np.array(labels)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-23T10:18:37.816563Z","iopub.execute_input":"2026-03-23T10:18:37.817308Z","iopub.status.idle":"2026-03-23T10:20:35.439935Z","shell.execute_reply.started":"2026-03-23T10:18:37.817269Z","shell.execute_reply":"2026-03-23T10:20:35.439016Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Convert to pytorch geometric data","metadata":{}},{"cell_type":"code","source":"x = torch.tensor(coords, dtype=torch.float)\n\nedge_index = torch.tensor(edge_index, dtype=torch.long) if not isinstance(edge_index, torch.Tensor) else edge_index.clone().detach().long()\n\ny = torch.tensor(labels, dtype=torch.float)\n\ndata = Data(x=x, edge_index=edge_index, y=y)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-23T10:21:33.169110Z","iopub.execute_input":"2026-03-23T10:21:33.169644Z","iopub.status.idle":"2026-03-23T10:21:33.203109Z","shell.execute_reply.started":"2026-03-23T10:21:33.169604Z","shell.execute_reply":"2026-03-23T10:21:33.202346Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## GNN Model","metadata":{}},{"cell_type":"code","source":"class GNN(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.conv1 = GCNConv(3, 64)\n        self.conv2 = GCNConv(64, 32)\n\n        # Edge classifier\n        self.edge_mlp = nn.Sequential(\n            nn.Linear(32 * 2, 64),\n            nn.ReLU(),\n            nn.Linear(64, 1)\n        )\n\n    def forward(self, data):\n        x, edge_index = data.x, data.edge_index\n\n        # Step 1: Node embeddings\n        x = self.conv1(x, edge_index)\n        x = torch.relu(x)\n\n        x = self.conv2(x, edge_index)\n        x = torch.relu(x)\n\n        # Step 2: Get edge node pairs\n        row, col = edge_index\n\n        # Step 3: Concatenate node embeddings\n        edge_features = torch.cat([x[row], x[col]], dim=1)\n\n        # Step 4: Predict per edge\n        out = self.edge_mlp(edge_features)\n\n        return out.squeeze()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-23T10:21:37.781761Z","iopub.execute_input":"2026-03-23T10:21:37.782081Z","iopub.status.idle":"2026-03-23T10:21:37.789434Z","shell.execute_reply.started":"2026-03-23T10:21:37.782050Z","shell.execute_reply":"2026-03-23T10:21:37.788311Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Train model","metadata":{}},{"cell_type":"code","source":"model = GNN()\noptimizer = torch.optim.Adam(model.parameters(), lr=0.001)\nloss_fn = nn.BCEWithLogitsLoss()\n\nfor epoch in range(50):\n    model.train()\n    optimizer.zero_grad()\n\n    out = model(data)  # now matches edge labels\n    loss = loss_fn(out, data.y)\n\n    loss.backward()\n    optimizer.step()\n\n    print(f\"Epoch {epoch}, Loss: {loss.item()}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-23T11:02:25.813338Z","iopub.execute_input":"2026-03-23T11:02:25.813761Z","iopub.status.idle":"2026-03-23T11:04:58.278964Z","shell.execute_reply.started":"2026-03-23T11:02:25.813725Z","shell.execute_reply":"2026-03-23T11:04:58.278106Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Evaluation\n","metadata":{}},{"cell_type":"markdown","source":"### Prediction","metadata":{}},{"cell_type":"code","source":"model.eval()\n\nwith torch.no_grad():\n    logits = model(data)\n    probs = torch.sigmoid(logits)   # convert to probabilities\n    preds = (probs > 0.5).float()   # threshold","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-23T11:06:56.395730Z","iopub.execute_input":"2026-03-23T11:06:56.396076Z","iopub.status.idle":"2026-03-23T11:06:57.773058Z","shell.execute_reply.started":"2026-03-23T11:06:56.396042Z","shell.execute_reply":"2026-03-23T11:06:57.772231Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Accuracy","metadata":{}},{"cell_type":"code","source":"accuracy = (preds == data.y).float().mean()\nprint(\"Edge Accuracy:\", accuracy.item())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-23T11:07:03.466023Z","iopub.execute_input":"2026-03-23T11:07:03.466335Z","iopub.status.idle":"2026-03-23T11:07:03.474478Z","shell.execute_reply.started":"2026-03-23T11:07:03.466303Z","shell.execute_reply":"2026-03-23T11:07:03.473389Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Precision/Recall/F1 Score","metadata":{}},{"cell_type":"code","source":"from sklearn.metrics import precision_score, recall_score, f1_score\n\ny_true = data.y.cpu().numpy()\ny_pred = preds.cpu().numpy()\n\nprecision = precision_score(y_true, y_pred)\nrecall = recall_score(y_true, y_pred)\nf1 = f1_score(y_true, y_pred)\n\nprint(\"Precision:\", precision)\nprint(\"Recall:\", recall)\nprint(\"F1 Score:\", f1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-23T11:07:09.927178Z","iopub.execute_input":"2026-03-23T11:07:09.927500Z","iopub.status.idle":"2026-03-23T11:07:10.217702Z","shell.execute_reply.started":"2026-03-23T11:07:09.927467Z","shell.execute_reply":"2026-03-23T11:07:10.216963Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### ROC AUC Score","metadata":{}},{"cell_type":"code","source":"from sklearn.metrics import roc_auc_score\n\ny_probs = probs.cpu().numpy()\nroc_auc = roc_auc_score(y_true, y_probs)\n\nprint(\"ROC-AUC:\", roc_auc)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-23T11:07:16.901333Z","iopub.execute_input":"2026-03-23T11:07:16.901721Z","iopub.status.idle":"2026-03-23T11:07:17.245894Z","shell.execute_reply.started":"2026-03-23T11:07:16.901686Z","shell.execute_reply":"2026-03-23T11:07:17.244873Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Visualisation","metadata":{}},{"cell_type":"markdown","source":"### True Track Visualisation","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\n# pick a real particle (ignore noise = particle_id 0)\nvalid_particles = df[df['particle_id'] != 0]['particle_id'].unique()\npid = valid_particles[0]\n\ntrack = df[df['particle_id'] == pid]\n\nplt.scatter(track['x'], track['y'], s=10)\nplt.title(\"True Track\")\nplt.xlabel(\"x\")\nplt.ylabel(\"y\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-23T11:07:24.779828Z","iopub.execute_input":"2026-03-23T11:07:24.780635Z","iopub.status.idle":"2026-03-23T11:07:24.951808Z","shell.execute_reply.started":"2026-03-23T11:07:24.780595Z","shell.execute_reply":"2026-03-23T11:07:24.950900Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Predicted Track visualisation","metadata":{}},{"cell_type":"code","source":"coords = data.x.cpu().numpy()\nedge_index_np = data.edge_index.cpu().numpy()\npred_edges = edge_index_np[:, preds.cpu().numpy() == 1]\n\nplt.figure(figsize=(6,6))\n\nfor i, j in pred_edges.T[:500]:  # limit for clarity\n    x_vals = [coords[i][0], coords[j][0]]\n    y_vals = [coords[i][1], coords[j][1]]\n    plt.plot(x_vals, y_vals, 'b-', alpha=0.3)\n\nplt.scatter(coords[:,0], coords[:,1], s=1, color='red')\nplt.title(\"Predicted Tracks (Edges)\")\nplt.xlabel(\"x\")\nplt.ylabel(\"y\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-23T11:07:30.216223Z","iopub.execute_input":"2026-03-23T11:07:30.216979Z","iopub.status.idle":"2026-03-23T11:07:30.703218Z","shell.execute_reply.started":"2026-03-23T11:07:30.216937Z","shell.execute_reply":"2026-03-23T11:07:30.702228Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## True Track vs Predicted Track","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize=(6,6))\n\n# true track\nplt.scatter(track['x'], track['y'], color='green', label='True Track', s=10)\n\n# predicted edges\nfor i, j in pred_edges.T[:300]:\n    plt.plot([coords[i][0], coords[j][0]],\n             [coords[i][1], coords[j][1]],\n             'blue', alpha=0.2)\n\nplt.legend()\nplt.title(\"True vs Predicted Tracks\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-23T11:07:38.319308Z","iopub.execute_input":"2026-03-23T11:07:38.319701Z","iopub.status.idle":"2026-03-23T11:07:38.688738Z","shell.execute_reply.started":"2026-03-23T11:07:38.319666Z","shell.execute_reply":"2026-03-23T11:07:38.687835Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}