{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import getpass\nfrom pathlib import Path\nfrom typing import Any, Callable, List, Optional, Sequence, Tuple, Union\n\nimport matplotlib.pyplot as plt\nimport random\nimport os\nimport numpy as np\nimport pandas as pd\nimport pytorch_lightning as pl\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nfrom scipy.interpolate import interp1d\nfrom sklearn.preprocessing import RobustScaler\nfrom torch import LongTensor, Tensor\n# from torch.utils.data import DataLoader, Dataset\nfrom tqdm import tqdm\nfrom transformers import get_cosine_schedule_with_warmup\n\nCOMP_NAME = \"icecube-neutrinos-in-deep-ice\"\n# Return the “login name” of the user\nKERNEL = False if getpass.getuser() == \"anjum\" else True\nif not KERNEL:  # in personal computer\n    INPUT_PATH = Path(f\"/mnt/storage_dimm2/kaggle_data/{COMP_NAME}\")\n    OUTPUT_PATH = Path(f\"/mnt/storage_dimm2/kaggle_output/{COMP_NAME}\")\n    MODEL_CACHE = Path(\"/mnt/storage/model_cache/torch\")\n    TRANSPARENCY_PATH = INPUT_PATH / \"ice_transparency.txt\"\nelse:           # in kaggle\n    INPUT_PATH = Path(f\"/kaggle/input/{COMP_NAME}\")\n    MODEL_CACHE = None\n    TRANSPARENCY_PATH = \"/kaggle/input/icecubetransparency/ice_transparency.txt\"\n\n    # Install packages\n    import subprocess\n\n    if torch.cuda.is_available():\n        whls = [\n            \"/kaggle/input/pytorchgeometric/torch_cluster-1.6.0-cp37-cp37m-linux_x86_64.whl\",\n            \"/kaggle/input/pytorchgeometric/torch_scatter-2.1.0-cp37-cp37m-linux_x86_64.whl\",\n            \"/kaggle/input/pytorchgeometric/torch_sparse-0.6.16-cp37-cp37m-linux_x86_64.whl\",\n            \"/kaggle/input/pytorchgeometric/torch_spline_conv-1.2.1-cp37-cp37m-linux_x86_64.whl\",\n            \"/kaggle/input/pytorchgeometric/torch_geometric-2.2.0-py3-none-any.whl\",\n            \"/kaggle/input/pytorchgeometric/ruamel.yaml-0.17.21-py3-none-any.whl\",\n        ]\n    else:\n        whls = [\n            \"/kaggle/input/pytorch-geometric/PyTorch-Geometric/torch_cluster-1.6.0-cp37-cp37m-linux_x86_64.whl\",\n            \"/kaggle/input/pytorch-geometric/PyTorch-Geometric/torch_scatter-2.0.9-cp37-cp37m-linux_x86_64.whl\",\n            \"/kaggle/input/pytorch-geometric/PyTorch-Geometric/torch_sparse-0.6.15-cp37-cp37m-linux_x86_64.whl\",\n            \"/kaggle/input/pytorch-geometric/PyTorch-Geometric/torch_spline_conv-1.2.1-cp37-cp37m-linux_x86_64.whl\",\n            \"/kaggle/input/pytorch-geometric/PyTorch-Geometric/torch_geometric-2.1.0.post1-py3-none-any.whl\",\n            \"/kaggle/input/pytorchgeometric/ruamel.yaml-0.17.21-py3-none-any.whl\",\n        ]\n\n    for w in whls:\n        print(\"Installing\", w)\n        subprocess.call([\"pip\", \"install\", w, \"--no-deps\", \"--upgrade\"])\n\n    import sys\n#     sys.path.append(\"/kaggle/input/graphnet/graphnet-main/src\")\n\n# from graphnet.models.graph_builders import KNNGraphBuilder\n# from graphnet.models.task.reconstruction import (\n#     AzimuthReconstructionWithKappa,\n#     ZenithReconstruction,\n# )\n# from graphnet.training.loss_functions import VonMisesFisher2DLoss, CosineLoss\n# from graphnet.models.gnn.gnn import GNN\n# from graphnet.models.utils import calculate_xyzt_homophily\n# from graphnet.utilities.config import save_model_config\n\nimport torch_geometric\nimport torch_geometric.nn as pyg_nn\nfrom torch_geometric.data import Data, Dataset\nfrom torch_geometric.loader import DataLoader\n# from torch_geometric.nn import EdgeConv\nfrom torch_geometric.nn import EdgeConv, SAGEConv, ChebConv, GCNConv\n# from torch_geometric.nn.pool import knn_graph\nfrom torch_geometric.nn import knn_graph\nfrom torch_geometric.typing import Adj\n\nimport torch_scatter\nfrom torch_scatter import scatter_max, scatter_mean, scatter_min, scatter_sum\n\nGLOBAL_POOLINGS = {\n    \"min\": scatter_min,\n    \"max\": scatter_max,\n    \"sum\": scatter_sum,\n    \"mean\": scatter_mean,\n}\n\n_dtype = {\n    \"batch_id\": \"int16\",\n    \"event_id\": \"int64\",\n}","metadata":{"_uuid":"632c42d7-0e51-42ee-8205-cc7fec9d0c56","_cell_guid":"7e3ab8b8-8781-4d27-958b-ea47012eb172","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-02-26T06:37:51.444493Z","iopub.execute_input":"2023-02-26T06:37:51.444818Z","iopub.status.idle":"2023-02-26T06:40:09.860922Z","shell.execute_reply.started":"2023-02-26T06:37:51.444736Z","shell.execute_reply":"2023-02-26T06:40:09.859782Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"https://github.com/graphnet-team/graphnet","metadata":{}},{"cell_type":"markdown","source":"## Dataset","metadata":{}},{"cell_type":"code","source":"!cat /kaggle/input/icecubetransparency/ice_transparency.txt","metadata":{"execution":{"iopub.status.busy":"2023-02-26T06:40:09.863273Z","iopub.execute_input":"2023-02-26T06:40:09.864715Z","iopub.status.idle":"2023-02-26T06:40:10.814846Z","shell.execute_reply.started":"2023-02-26T06:40:09.864669Z","shell.execute_reply":"2023-02-26T06:40:10.812261Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# datasets.py\ndef ice_transparency(data_path, datum=1950):\n    # Data from page 31 of https://arxiv.org/pdf/1301.5361.pdf\n    # Datum is from footnote 8 of page 29\n    df = pd.read_csv(data_path, delim_whitespace=True)\n    df[\"z\"] = df[\"depth\"] - datum\n    df[\"z_norm\"] = df[\"z\"] / 500\n    df[[\"scattering_len_norm\", \"absorption_len_norm\"]] = RobustScaler().fit_transform(\n        df[[\"scattering_len\", \"absorption_len\"]]\n    )\n\n    # These are both roughly equivalent after scaling\n    f_scattering = interp1d(df[\"z_norm\"], df[\"scattering_len_norm\"])\n    f_absorption = interp1d(df[\"z_norm\"], df[\"absorption_len_norm\"])\n    return f_scattering, f_absorption","metadata":{"execution":{"iopub.status.busy":"2023-02-26T06:40:10.819948Z","iopub.execute_input":"2023-02-26T06:40:10.823888Z","iopub.status.idle":"2023-02-26T06:40:10.844410Z","shell.execute_reply.started":"2023-02-26T06:40:10.823833Z","shell.execute_reply":"2023-02-26T06:40:10.842874Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class IceCubeSubmissionDataset(Dataset):\n    def __init__(\n        self,\n        batch_id,\n        event_ids,\n        sensor_df,\n        mode=\"test\",\n        y=None,\n        pulse_limit=300,\n        transform=None,\n        pre_transform=None,\n        pre_filter=None,\n    ):\n        super().__init__(transform, pre_transform, pre_filter)\n        self.y = y\n        self.event_ids = event_ids\n        self.sensor_df = sensor_df\n        self.pulse_limit = pulse_limit\n        self.f_scattering, self.f_absorption = ice_transparency(TRANSPARENCY_PATH)\n        self.batch_df = pd.read_parquet(INPUT_PATH / mode / f\"batch_{batch_id}.parquet\")\n\n        self.batch_df[\"time\"] = (self.batch_df[\"time\"] - 1.0e04) / 3.0e4\n        self.batch_df[\"charge\"] = np.log10(self.batch_df[\"charge\"]) / 3.0\n        self.batch_df[\"auxiliary\"] = self.batch_df[\"auxiliary\"].astype(int) - 0.5\n\n    def len(self):\n        return len(self.event_ids)\n\n    def get(self, idx):\n        event_id = self.event_ids[idx]\n        event = self.batch_df.loc[event_id]\n    \n        # represent each event by a single graph\n        event = pd.merge(event, self.sensor_df, on=\"sensor_id\")\n        \n        col = [\"x\", \"y\", \"z\", \"time\", \"charge\", \"qe\", \"auxiliary\"]\n        x = event[col].values\n        x = torch.tensor(x, dtype=torch.float32)\n        \n        data = Data(x=x, n_pulses=torch.tensor(x.shape[0], dtype=torch.int32))\n\n        # Add ice transparency data\n        z = data.x[:, 2].numpy()\n        scattering = torch.tensor(self.f_scattering(z), dtype=torch.float32).view(-1, 1)\n        data.x = torch.cat([data.x, scattering], dim=1)\n\n        # Downsample the large events\n        if data.n_pulses > self.pulse_limit:\n            data.x = data.x[np.random.choice(data.n_pulses, self.pulse_limit)]\n            data.n_pulses = torch.tensor(self.pulse_limit, dtype=torch.int32)\n    \n        # Builds graph from the k-nearest neighbours.\n        data.edge_index = knn_graph(\n            data.x[:, [0, 1, 2]],  # x, y, z\n            k=8,\n            batch=None,\n            loop=False\n        )\n\n        if self.y is not None:\n            y = self.y.loc[idx, :].values\n            y = torch.tensor(y, dtype=torch.float32)\n            data.y = y\n\n        return data","metadata":{"execution":{"iopub.status.busy":"2023-02-26T06:40:10.852196Z","iopub.execute_input":"2023-02-26T06:40:10.852563Z","iopub.status.idle":"2023-02-26T06:40:10.882721Z","shell.execute_reply.started":"2023-02-26T06:40:10.852521Z","shell.execute_reply":"2023-02-26T06:40:10.880311Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# preprocessing.py\ndef prepare_sensors():\n    sensors = pd.read_csv(INPUT_PATH / \"sensor_geometry.csv\").astype(\n        {\n            \"sensor_id\": np.int16,\n            \"x\": np.float32,\n            \"y\": np.float32,\n            \"z\": np.float32,\n        }\n    )\n    sensors[\"string\"] = 0\n    sensors[\"qe\"] = 1\n\n    for i in range(len(sensors) // 60):\n        start, end = i * 60, (i * 60) + 60\n        sensors.loc[start:end, \"string\"] = i\n\n        # High Quantum Efficiency in the lower 50 DOMs - https://arxiv.org/pdf/2209.03042.pdf (Figure 1)\n        if i in range(78, 86):\n            start_veto, end_veto = i * 60, (i * 60) + 10\n            start_core, end_core = end_veto + 1, (i * 60) + 60\n            sensors.loc[start_core:end_core, \"qe\"] = 1.35\n\n    # https://github.com/graphnet-team/graphnet/blob/b2bad25528652587ab0cdb7cf2335ee254cfa2db/src/graphnet/models/detector/icecube.py#L33-L41\n    # Assume that \"rde\" (relative dom efficiency) is equivalent to QE\n    sensors[\"x\"] /= 500\n    sensors[\"y\"] /= 500\n    sensors[\"z\"] /= 500\n    sensors[\"qe\"] -= 1.25\n    sensors[\"qe\"] /= 0.25\n\n    return sensors","metadata":{"execution":{"iopub.status.busy":"2023-02-26T06:40:10.886977Z","iopub.execute_input":"2023-02-26T06:40:10.889226Z","iopub.status.idle":"2023-02-26T06:40:10.905745Z","shell.execute_reply.started":"2023-02-26T06:40:10.889178Z","shell.execute_reply":"2023-02-26T06:40:10.904586Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sensors = prepare_sensors()\nsensors","metadata":{"execution":{"iopub.status.busy":"2023-02-26T06:40:10.907177Z","iopub.execute_input":"2023-02-26T06:40:10.907626Z","iopub.status.idle":"2023-02-26T06:40:11.012875Z","shell.execute_reply.started":"2023-02-26T06:40:10.907592Z","shell.execute_reply":"2023-02-26T06:40:11.012022Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"meta = pd.read_parquet(\n    INPUT_PATH / f\"train_meta.parquet\", columns=[\"batch_id\", \"event_id\", \"azimuth\", \"zenith\"]\n).astype(_dtype)\nmeta","metadata":{"execution":{"iopub.status.busy":"2023-02-26T06:40:11.016911Z","iopub.execute_input":"2023-02-26T06:40:11.019251Z","iopub.status.idle":"2023-02-26T06:40:41.155031Z","shell.execute_reply.started":"2023-02-26T06:40:11.019216Z","shell.execute_reply":"2023-02-26T06:40:41.153432Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch_ids = meta[\"batch_id\"].unique()\nbatch_ids","metadata":{"execution":{"iopub.status.busy":"2023-02-26T06:40:41.160918Z","iopub.execute_input":"2023-02-26T06:40:41.165897Z","iopub.status.idle":"2023-02-26T06:40:41.875325Z","shell.execute_reply.started":"2023-02-26T06:40:41.165799Z","shell.execute_reply":"2023-02-26T06:40:41.874347Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# for i, b in enumerate(batch_ids):\n#     event_ids = meta[meta[\"batch_id\"] == b][\"event_id\"].tolist()\n#     y = meta[meta[\"batch_id\"] == b][['zenith', 'azimuth']].reset_index(drop=True)\n#     dataset = IceCubeSubmissionDataset(\n#         b, event_ids, sensors, mode='train', y=y,\n#     )\n#     print(f'batch {i}')\n#     print(\"num of graph:\", len(dataset), '\\t', dataset[0], '\\t', dataset[1])\n#     if i >= 3:\n#         break","metadata":{"execution":{"iopub.status.busy":"2023-02-26T06:40:41.877007Z","iopub.execute_input":"2023-02-26T06:40:41.877358Z","iopub.status.idle":"2023-02-26T06:40:41.882239Z","shell.execute_reply.started":"2023-02-26T06:40:41.877323Z","shell.execute_reply":"2023-02-26T06:40:41.881243Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# dataset[0].edge_index","metadata":{"execution":{"iopub.status.busy":"2023-02-26T06:40:41.886910Z","iopub.execute_input":"2023-02-26T06:40:41.888067Z","iopub.status.idle":"2023-02-26T06:40:41.893136Z","shell.execute_reply.started":"2023-02-26T06:40:41.888030Z","shell.execute_reply":"2023-02-26T06:40:41.892140Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class IceCubeDataset(Dataset):\n    def __init__(\n        self,\n        batch_ids,\n        sensor_df,\n        mode=\"train\",\n        pulse_limit=300,  # the max nums of node in one graph\n        n_sample=100,     # the max nums of graph in one batch\n        root='',\n        transform=None,\n        pre_transform=None,\n        pre_filter=None,\n    ):\n        self.batch_ids = batch_ids\n        self.sensor_df = sensor_df\n        self.pulse_limit = pulse_limit\n        self.f_scattering, self.f_absorption = ice_transparency(TRANSPARENCY_PATH)\n        self.current_batch = batch_ids[0]\n        self.mode = mode\n        self.n_sample = n_sample\n        self.data = None\n        super().__init__(root, transform, pre_transform, pre_filter)\n\n\n    @property\n    def processed_file_names(self):\n        return ['data_1.pt', 'data_2.pt']\n\n    def len(self):\n        return len(self.batch_ids)*self.n_sample\n\n    def get(self, idx):\n        if idx > self.n_sample*len(self.batch_ids) or self.data is None:\n            batch_id = idx // self.n_sample\n            self.data = torch.load(os.path.join(self.processed_dir, f'data_{batch_id+1}.pt'))\n        return self.data[idx%self.n_sample]\n    \n    def process(self):\n        for batch_id in self.batch_ids:\n            event_ids = meta[meta[\"batch_id\"] == batch_id][\"event_id\"].tolist()\n            \n            batch_df = pd.read_parquet(INPUT_PATH / self.mode / f\"batch_{batch_id}.parquet\")\n            batch_df[\"time\"] = (batch_df[\"time\"] - 1.0e04) / 3.0e4\n            batch_df[\"charge\"] = np.log10(batch_df[\"charge\"]) / 3.0\n            batch_df[\"auxiliary\"] = batch_df[\"auxiliary\"].astype(int) - 0.5\n            \n            if self.mode == 'train':\n                y_df = meta[meta[\"batch_id\"] == batch_id][['zenith', 'azimuth']].reset_index(drop=True)\n            \n            data_list = []\n            event_sample = range(len(event_ids))\n            for i in random.sample(event_sample, self.n_sample):\n                event_id = event_ids[i]\n                event = batch_df.loc[event_id]\n                event = pd.merge(event, self.sensor_df, on=\"sensor_id\")\n                \n                col = [\"x\", \"y\", \"z\", \"time\", \"charge\", \"qe\", \"auxiliary\"]\n                x = event[col].values\n                x = torch.tensor(x, dtype=torch.float32)\n                \n                data = Data(x=x, n_pulses=torch.tensor(x.shape[0], dtype=torch.int32))\n                # Add ice transparency data\n                z = data.x[:, 2].numpy()\n                scattering = torch.tensor(self.f_scattering(z), dtype=torch.float32).view(-1, 1)\n                data.x = torch.cat([data.x, scattering], dim=1)\n\n                # Downsample the large events\n                if data.n_pulses > self.pulse_limit:\n                    data.x = data.x[np.random.choice(data.n_pulses, self.pulse_limit)]\n                    data.n_pulses = torch.tensor(self.pulse_limit, dtype=torch.int32)\n\n                # Builds graph from the k-nearest neighbours.\n                data.edge_index = knn_graph(\n                    data.x[:, [0, 1, 2]],  # x, y, z\n                    k=8,\n                    batch=None,\n                    loop=False\n                )\n                \n                if self.mode == 'train':\n                    y = y_df.loc[batch_id, :].values\n                    y = torch.tensor(y, dtype=torch.float32)\n                    data.y = y\n\n                if self.pre_filter is not None and not self.pre_filter(data):\n                    continue\n                if self.pre_transform is not None:\n                    data = self.pre_transform(data)\n\n                data_list.append(data)\n\n            torch.save(data_list, os.path.join(self.processed_dir, f'data_{batch_id}.pt'))","metadata":{"execution":{"iopub.status.busy":"2023-02-26T06:43:09.832591Z","iopub.execute_input":"2023-02-26T06:43:09.832976Z","iopub.status.idle":"2023-02-26T06:43:09.857598Z","shell.execute_reply.started":"2023-02-26T06:43:09.832942Z","shell.execute_reply":"2023-02-26T06:43:09.856597Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset = IceCubeDataset(batch_ids[0:300], sensors, mode='train', pulse_limit=300, n_sample=10)","metadata":{"execution":{"iopub.status.busy":"2023-02-26T07:19:07.553069Z","iopub.execute_input":"2023-02-26T07:19:07.553463Z","iopub.status.idle":"2023-02-26T07:19:07.572246Z","shell.execute_reply.started":"2023-02-26T07:19:07.553430Z","shell.execute_reply":"2023-02-26T07:19:07.571375Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# !rm -rf ./processed\n!ls ./processed","metadata":{"execution":{"iopub.status.busy":"2023-02-26T07:19:09.305045Z","iopub.execute_input":"2023-02-26T07:19:09.305435Z","iopub.status.idle":"2023-02-26T07:19:10.324073Z","shell.execute_reply.started":"2023-02-26T07:19:09.305397Z","shell.execute_reply":"2023-02-26T07:19:10.322929Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(dataset)","metadata":{"execution":{"iopub.status.busy":"2023-02-26T07:19:13.961423Z","iopub.execute_input":"2023-02-26T07:19:13.961996Z","iopub.status.idle":"2023-02-26T07:19:13.973257Z","shell.execute_reply.started":"2023-02-26T07:19:13.961956Z","shell.execute_reply":"2023-02-26T07:19:13.972157Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset[0]","metadata":{"execution":{"iopub.status.busy":"2023-02-26T07:19:14.788879Z","iopub.execute_input":"2023-02-26T07:19:14.789268Z","iopub.status.idle":"2023-02-26T07:19:14.819563Z","shell.execute_reply.started":"2023-02-26T07:19:14.789228Z","shell.execute_reply":"2023-02-26T07:19:14.818437Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_len = int(0.7*len(dataset))\nval_len = len(dataset) - train_len\ntrain_loader = DataLoader(dataset[0:train_len], batch_size=32, num_workers=1)\nval_loader = DataLoader(dataset[train_len:], batch_size=32, num_workers=1)","metadata":{"execution":{"iopub.status.busy":"2023-02-26T07:19:16.639005Z","iopub.execute_input":"2023-02-26T07:19:16.639389Z","iopub.status.idle":"2023-02-26T07:19:16.646191Z","shell.execute_reply.started":"2023-02-26T07:19:16.639337Z","shell.execute_reply":"2023-02-26T07:19:16.645137Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dynamic graph","metadata":{}},{"cell_type":"code","source":"def calculate_distance_matrix(xyz_coords: Tensor) -> Tensor:\n    \"\"\"Calculate the matrix of pairwise distances between pulses.\n    Args:\n        xyz_coords: (x,y,z)-coordinates of pulses, of shape [nb_doms, 3].\n    Returns:\n        Matrix of pairwise distances, of shape [nb_doms, nb_doms]\n    \"\"\"\n    diff = xyz_coords.unsqueeze(dim=2) - xyz_coords.T.unsqueeze(dim=0)\n    return torch.sqrt(torch.sum(diff**2, dim=1))\n\n\nclass EuclideanGraphBuilder(nn.Module):\n    \"\"\"Builds graph according to Euclidean distance between nodes.\n    See https://arxiv.org/pdf/1809.06166.pdf.\n    \"\"\"\n    def __init__(\n        self,\n        sigma: float,\n        threshold: float = 0.0,\n        columns: List[int] = None,\n    ):\n        \"\"\"Construct `EuclideanGraphBuilder`.\"\"\"\n        # Base class constructor\n        super().__init__()\n\n        # Check(s)\n        if columns is None:\n            columns = [0, 1, 2]\n\n        # Member variable(s)\n        self._sigma = sigma\n        self._threshold = threshold\n        self._columns = columns\n\n    def forward(self, data: Data) -> Data:\n        \"\"\"Forward pass.\"\"\"\n        # Constructs the adjacency matrix from the raw, DOM-level data and\n        # returns this matrix\n        xyz_coords = data.x[:, self._columns]\n\n        # Construct block-diagonal matrix indicating whether pulses belong to\n        # the same event in the batch\n        batch_mask = data.batch.unsqueeze(dim=0) == data.batch.unsqueeze(dim=1)\n\n        distance_matrix = calculate_distance_matrix(xyz_coords)\n        affinity_matrix = torch.exp(\n            -0.5 * distance_matrix**2 / self._sigma**2\n        )\n\n        # Use softmax to normalise all adjacencies to one for each node\n        exp_row_sums = torch.exp(affinity_matrix).sum(axis=1)\n        weighted_adj_matrix = torch.exp(\n            affinity_matrix\n        ) / exp_row_sums.unsqueeze(dim=1)\n\n        # Only include edges with weights that exceed the chosen threshold (and\n        # are part of the same event)\n        sources, targets = torch.where(\n            (weighted_adj_matrix > self._threshold) & (batch_mask)\n        )\n        edge_weights = weighted_adj_matrix[sources, targets]\n\n        data.edge_index = torch.stack((sources, targets))\n        data.edge_weight = edge_weights\n\n        return data","metadata":{"execution":{"iopub.status.busy":"2023-02-26T07:19:18.277171Z","iopub.execute_input":"2023-02-26T07:19:18.277892Z","iopub.status.idle":"2023-02-26T07:19:18.289740Z","shell.execute_reply.started":"2023-02-26T07:19:18.277853Z","shell.execute_reply":"2023-02-26T07:19:18.288385Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model","metadata":{}},{"cell_type":"code","source":"class DenseDynBlock(nn.Module):\n    \"\"\"\n    Dense Dynamic graph convolution block\n    \"\"\"\n    def __init__(self, in_channels, out_channels=64, sigma=0.5):\n        super(DenseDynBlock, self).__init__()\n        self.GraphBuilder = EuclideanGraphBuilder(sigma=sigma)\n        self.gnn = GCNConv(in_channels, out_channels)\n\n    def forward(self, data):\n        data1 = self.GraphBuilder(data)\n        x, edge_index, batch = data1.x, data1.edge_index, data1.batch\n        x = self.gnn(x, edge_index)\n        data1.x = torch.cat((x, data.x), 1)\n        return data1","metadata":{"execution":{"iopub.status.busy":"2023-02-26T07:19:19.361267Z","iopub.execute_input":"2023-02-26T07:19:19.362584Z","iopub.status.idle":"2023-02-26T07:19:19.369623Z","shell.execute_reply.started":"2023-02-26T07:19:19.362534Z","shell.execute_reply":"2023-02-26T07:19:19.368580Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MyGNN(nn.Module):\n    \"\"\"\n    Dynamic graph convolution layer\n    \"\"\"\n    def __init__(self, in_channels, hidden_channels, out_channels, n_blocks):\n        super().__init__()\n        self.n_blocks = n_blocks\n        self.head = SAGEConv(in_channels, hidden_channels)\n        c_growth  = hidden_channels\n        self.gnn = nn.Sequential(*[DenseDynBlock(hidden_channels+i*c_growth, c_growth)\n                                    for i in range(n_blocks-1)])\n        fusion_dims = int(hidden_channels * self.n_blocks + c_growth * ((1 + self.n_blocks - 1) * (self.n_blocks - 1) / 2))\n        self.linear = nn.Linear(fusion_dims, out_channels)\n\n    def forward(self, data):\n        x, edge_index, batch = data.x, data.edge_index, data.batch\n        data.x = self.head(x, edge_index)\n        feats = [data.x]\n        for i in range(self.n_blocks-1):\n            data = self.gnn[i](data)\n            feats.append(data.x)\n        feats = torch.cat(feats, 1)\n        x = pyg_nn.global_mean_pool(feats, data.batch)\n        out = F.relu(self.linear(x))\n        return out","metadata":{"execution":{"iopub.status.busy":"2023-02-26T07:19:19.869319Z","iopub.execute_input":"2023-02-26T07:19:19.870057Z","iopub.status.idle":"2023-02-26T07:19:19.879778Z","shell.execute_reply.started":"2023-02-26T07:19:19.870017Z","shell.execute_reply":"2023-02-26T07:19:19.878627Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = MyGNN(8, 16, 2, 3)\nmodel","metadata":{"execution":{"iopub.status.busy":"2023-02-26T07:19:20.458720Z","iopub.execute_input":"2023-02-26T07:19:20.459630Z","iopub.status.idle":"2023-02-26T07:19:20.475143Z","shell.execute_reply.started":"2023-02-26T07:19:20.459580Z","shell.execute_reply":"2023-02-26T07:19:20.474181Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for sample_batched in train_loader:\n    outputs = model(sample_batched)\n    print(outputs.shape, sample_batched.x.shape, sample_batched.y.shape)\n    break","metadata":{"execution":{"iopub.status.busy":"2023-02-26T07:19:21.033220Z","iopub.execute_input":"2023-02-26T07:19:21.033904Z","iopub.status.idle":"2023-02-26T07:19:22.137949Z","shell.execute_reply.started":"2023-02-26T07:19:21.033865Z","shell.execute_reply":"2023-02-26T07:19:22.136631Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## train","metadata":{}},{"cell_type":"code","source":"epochs = 10\nbatchsize = 32\ncriterion = nn.MSELoss()\noptimizer = torch.optim.AdamW(model.parameters(), lr=0.3)\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint('using ', device)\nmodel = model.to(device)","metadata":{"execution":{"iopub.status.busy":"2023-02-26T07:19:23.799028Z","iopub.execute_input":"2023-02-26T07:19:23.799464Z","iopub.status.idle":"2023-02-26T07:19:29.591760Z","shell.execute_reply.started":"2023-02-26T07:19:23.799424Z","shell.execute_reply":"2023-02-26T07:19:29.590694Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"total_loss = []\nfor epoch_num in range(epochs):\n\n    model.train()\n    total_loss_train = 0\n    total_acc_train = 0\n    for sample_batched in train_loader:\n        sample_batched = sample_batched.to(device)\n        optimizer.zero_grad()\n        outputs = model(sample_batched)\n        # loss\n        label = sample_batched.y.reshape(-1, 2).to(device)\n        loss = criterion(outputs, label)\n        total_loss_train += loss.item()\n        # update\n        loss.backward()\n        optimizer.step()\n\n    model.eval()\n    total_loss_val = 0\n    with torch.no_grad():\n        for sample_batched in val_loader:\n            sample_batched = sample_batched.to(device)\n            outputs = model(sample_batched)\n            # loss\n            label = sample_batched.y.reshape(-1, 2).to(device)\n            loss = criterion(outputs, label)\n            total_loss_val += loss.item()\n\n    print(\n        \"Epoch: [{:0>3d} / {:0>3d}] | Train Loss: {:.4f} | Val Loss: {:.4f}\".format(\n            epoch_num + 1,\n            epochs,\n            total_loss_train / train_len,\n            total_loss_val / val_len,\n        )\n    )\n    total_loss.append([total_loss_train / train_len, total_loss_val / val_len])","metadata":{"execution":{"iopub.status.busy":"2023-02-26T07:19:29.595671Z","iopub.execute_input":"2023-02-26T07:19:29.597610Z","iopub.status.idle":"2023-02-26T07:19:52.496642Z","shell.execute_reply.started":"2023-02-26T07:19:29.597564Z","shell.execute_reply":"2023-02-26T07:19:52.495141Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"total_loss = np.array(total_loss, dtype=np.float32)\nplt.figure()\nplt.xlabel(\"epoch\")\nplt.ylabel(\"loss\")\nplt.plot(total_loss[:,0], 'b-', label=\"train\")\nplt.plot(total_loss[:,1], 'r-', label=\"val\")\nplt.legend()\nplt.grid(True)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-02-26T07:19:52.498324Z","iopub.execute_input":"2023-02-26T07:19:52.499496Z","iopub.status.idle":"2023-02-26T07:19:52.742042Z","shell.execute_reply.started":"2023-02-26T07:19:52.499449Z","shell.execute_reply":"2023-02-26T07:19:52.741056Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Infer","metadata":{}},{"cell_type":"code","source":"def infer(model, loader, device=\"cpu\"):\n    model.to(device)\n    model.eval()\n\n    predictions = []\n    with torch.no_grad():\n        for batch in loader:\n            batch = batch.to(device)\n            pred_angles = model(batch)\n            predictions.append(pred_angles.cpu())\n\n    return torch.cat(predictions, 0)","metadata":{"execution":{"iopub.status.busy":"2023-02-26T07:19:52.745249Z","iopub.execute_input":"2023-02-26T07:19:52.745560Z","iopub.status.idle":"2023-02-26T07:19:52.751382Z","shell.execute_reply.started":"2023-02-26T07:19:52.745531Z","shell.execute_reply":"2023-02-26T07:19:52.750269Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def make_predictions(model, device=\"cpu\", mode=\"test\", batch_size=32):\n    sensors = prepare_sensors()\n\n    meta = pd.read_parquet(\n        INPUT_PATH / f\"{mode}_meta.parquet\", columns=[\"batch_id\", \"event_id\"]\n    ).astype(_dtype)\n    batch_ids = meta[\"batch_id\"].unique()\n\n    if mode == \"train\":\n        batch_ids = batch_ids[:6]\n\n    batch_preds = []\n    for b in batch_ids:\n        event_ids = meta[meta[\"batch_id\"] == b][\"event_id\"].tolist()\n        dataset = IceCubeSubmissionDataset(\n            b, event_ids, sensors, mode=mode,\n        )\n        loader = DataLoader(dataset, batch_size=batch_size, num_workers=1)\n        batch_preds.append(infer(model, loader, device=device))\n        print(\"Finished batch\", b)\n\n        if mode == \"train\" and b == 6:\n            break\n\n    output = torch.cat(batch_preds, 0)\n\n    event_id_labels = []\n    for b in batch_ids:\n        event_id_labels.extend(meta[meta[\"batch_id\"] == b][\"event_id\"].tolist())\n\n    sub = {\n        \"event_id\": event_id_labels,\n        \"azimuth\": output[:, 0],\n        \"zenith\": output[:, 1],\n    }\n\n    sub = pd.DataFrame(sub)\n    sub.to_csv(\"submission.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2023-02-26T07:19:52.752814Z","iopub.execute_input":"2023-02-26T07:19:52.753397Z","iopub.status.idle":"2023-02-26T07:19:52.764923Z","shell.execute_reply.started":"2023-02-26T07:19:52.753343Z","shell.execute_reply":"2023-02-26T07:19:52.763942Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"make_predictions(model, device=\"cuda\", mode=\"test\", batch_size=32)","metadata":{"execution":{"iopub.status.busy":"2023-02-26T07:19:52.766378Z","iopub.execute_input":"2023-02-26T07:19:52.766727Z","iopub.status.idle":"2023-02-26T07:19:53.020194Z","shell.execute_reply.started":"2023-02-26T07:19:52.766691Z","shell.execute_reply":"2023-02-26T07:19:53.019012Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pd.read_csv(\"submission.csv\")","metadata":{"execution":{"iopub.status.busy":"2023-02-26T07:19:53.022873Z","iopub.execute_input":"2023-02-26T07:19:53.023557Z","iopub.status.idle":"2023-02-26T07:19:53.037968Z","shell.execute_reply.started":"2023-02-26T07:19:53.023506Z","shell.execute_reply":"2023-02-26T07:19:53.037005Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}