{"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":"markdown","source":"# Imports","metadata":{}},{"cell_type":"code","source":"import torch \ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(device)","metadata":{"execution":{"iopub.status.busy":"2023-03-08T09:19:57.858388Z","iopub.execute_input":"2023-03-08T09:19:57.859407Z","iopub.status.idle":"2023-03-08T09:19:59.941224Z","shell.execute_reply.started":"2023-03-08T09:19:57.859301Z","shell.execute_reply":"2023-03-08T09:19:59.940075Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if device.type == 'cpu':\n    # CPU version\n    ! pip install -q torch-scatter torch-sparse torch-cluster torch-spline-conv torch-geometric --no-index --find-links=file:///kaggle/input/pytorch-geometric/PyTorch-Geometric\nelif device.type == 'cuda':\n    # GPU version\n    ! pip install -q torch-scatter torch-sparse torch-cluster torch-spline-conv torch-geometric --no-index --find-links=file:///kaggle/input/pytorchgeometric\nelse:\n    raise Exception('Bruh.')","metadata":{"execution":{"iopub.status.busy":"2023-03-08T09:19:59.943747Z","iopub.execute_input":"2023-03-08T09:19:59.944460Z","iopub.status.idle":"2023-03-08T09:20:12.800889Z","shell.execute_reply.started":"2023-03-08T09:19:59.944415Z","shell.execute_reply":"2023-03-08T09:20:12.799754Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\nimport gc\n\nfrom torch_geometric.loader import DataLoader\n\n# Add python files\nimport sys\nsys.path.append('/kaggle/input/icecube-py')\nfrom dataset import MyOwnDataset\nfrom metrics import angular_dist_score\nimport pred_to_angles","metadata":{"execution":{"iopub.status.busy":"2023-03-08T09:20:12.803535Z","iopub.execute_input":"2023-03-08T09:20:12.803962Z","iopub.status.idle":"2023-03-08T09:20:16.643450Z","shell.execute_reply.started":"2023-03-08T09:20:12.803918Z","shell.execute_reply":"2023-03-08T09:20:16.642320Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load a pre-trained model","metadata":{}},{"cell_type":"code","source":"import torch\nfrom torch.nn import Linear, LeakyReLU\nfrom torch_geometric.nn import DynamicEdgeConv\nfrom torch_geometric.nn import global_mean_pool\nfrom torch import Tensor, LongTensor\nfrom torch_scatter import scatter_mean\nfrom torch_geometric.utils.homophily import homophily\nfrom torch_geometric.nn.aggr import MultiAggregation, AttentionalAggregation\n\n\nclass EdgeConvMLP(torch.nn.Module):\n    \"\"\"Basic convolutional block.\"\"\"\n    def __init__(self, dim_in, dim_hidden, dim_out):\n        super().__init__()\n\n        self.sequential = torch.nn.Sequential(\n            Linear(dim_in, dim_hidden),\n            LeakyReLU(),\n            Linear(dim_hidden, dim_out),\n            LeakyReLU(),\n        )\n\n    def forward(self, x):\n        return self.sequential(x)\n    \nclass GateMLP(torch.nn.Module):\n    \"\"\"Basic convolutional block.\"\"\"\n    def __init__(self, dim_in, dim_hidden, dim_out):\n        super().__init__()\n\n        self.sequential = torch.nn.Sequential(\n            Linear(dim_in, dim_hidden),\n            LeakyReLU(),\n            Linear(dim_hidden, dim_out),\n            LeakyReLU(),\n        )\n\n    def forward(self, x):\n        return self.sequential(x)\n\n\nclass DynEdgeAttention(torch.nn.Module):\n    \"\"\"Dynedge model from https://iopscience.iop.org/article/10.1088/1748-0221/17/11/P11003)\"\"\"\n    def __init__(self, num_node_features, dim_output, dropout_rate=0.):\n        super(DynEdgeAttention, self).__init__()\n        \n        torch.manual_seed(12345)\n        self.num_node_features = num_node_features\n        self.dim_output = dim_output\n        self.dropout_rate = dropout_rate\n        self.K = 8\n#         self.aggrs_list = ['mean', 'min' , 'max', 'sum', AttentionalAggregation(gate_nn=GateMLP(256, 512, 256))]\n        self.aggrs_list = [AttentionalAggregation(gate_nn=GateMLP(256, 64, 1))]\n\n        self.conv1 = DynamicEdgeConv(nn=EdgeConvMLP(2 * self.num_node_features, 336, 256), k=self.K)\n        self.conv2 = DynamicEdgeConv(nn=EdgeConvMLP(512, 336, 256), k=self.K)\n        self.conv3 = DynamicEdgeConv(nn=EdgeConvMLP(512, 336, 256), k=self.K)\n        self.conv4 = DynamicEdgeConv(nn=EdgeConvMLP(512, 336, 256), k=self.K)\n        \n        # final regressor\n        self.mlp1 = torch.nn.Sequential(\n            Linear(256 * 4 + self.num_node_features, 336),\n            LeakyReLU(),\n            Linear(336, 256),\n            LeakyReLU(),\n        )\n        \n        self.global_pool =  MultiAggregation(aggrs=self.aggrs_list)\n        \n#         mode_kwargs = {'in_channels': 256, 'out_channels': 256, 'num_heads': 16}\n#         self.global_pool =  MultiAggregation(aggrs=self.aggrs_list, mode='attn', mode_kwargs=mode_kwargs)\n\n#         AttentionalAggregation\n#         self.attentional_aggr = AttentionalAggregation(gate_nn=GateMLP(256, 512, 256))\n\n        self.mlp2 =  torch.nn.Sequential(\n            Linear(len(self.aggrs_list) * 256 + (4 + self.num_node_features), 128), # input depends of number of aggregating fns + 4 homophily + mean_node\n#             Linear(256 + (4 + self.num_node_features), 128),\n            LeakyReLU(),\n            Linear(128, self.dim_output)\n        )\n\n\n    def _calculate_global_variables(\n        self,\n        x: Tensor,\n        edge_index: LongTensor,\n        batch: LongTensor,\n    ) -> Tensor:\n        \"\"\"Calculate global variables.\"\"\"\n        # Calculate homophily (scalar variables)\n        h_x = homophily(edge_index, x[:, 0], batch).reshape(-1, 1)\n        h_y = homophily(edge_index, x[:, 1], batch).reshape(-1, 1)\n        h_z = homophily(edge_index, x[:, 2], batch).reshape(-1, 1)\n        h_t = homophily(edge_index, x[:, 3], batch).reshape(-1, 1)\n        \n        # Calculate mean features\n        global_means = scatter_mean(x, batch, dim=0)\n\n        # Add global variables\n        global_variables = torch.cat([global_means, h_x, h_y, h_z, h_t], dim=-1)\n\n        return global_variables\n\n    def forward(self, x, edge_index, batch):\n        # 0. Obtain global variables\n        global_x = self._calculate_global_variables(x, edge_index, batch)\n        \n        # 1. Obtain node embeddings at various embedding depths\n        x1 = self.conv1(x, batch)\n        x2 = self.conv2(x1, batch)\n        x3 = self.conv3(x2, batch)\n        x4 = self.conv4(x3, batch)\n\n        x = torch.cat([x, x1, x2, x3, x4], dim=-1)\n        \n        x = self.mlp1(x)\n\n        # 2. Pooling        \n        x = self.global_pool(x, batch)\n            \n        x = torch.cat([global_x, x], dim=-1)\n\n        # 3. Apply a final MLP regressor\n        x = self.mlp2(x)\n        \n        return x","metadata":{"execution":{"iopub.status.busy":"2023-03-08T09:20:16.648114Z","iopub.execute_input":"2023-03-08T09:20:16.648705Z","iopub.status.idle":"2023-03-08T09:20:16.674770Z","shell.execute_reply.started":"2023-03-08T09:20:16.648674Z","shell.execute_reply":"2023-03-08T09:20:16.673578Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Initalize your final model\nmodel = DynEdgeAttention(\n    num_node_features=5, \n    dim_output=3, \n    dropout_rate=0.\n).to(device)\n\n# Load model from path\nPATH_LOAD = '/kaggle/input/icecube-models/03-06-dynedgeattentionxyzvmfloss1.07.pt'\n\ntarget_mode = 'xyz' # angles / cossin / xyz\n\n\nif device.type == 'cpu':\n    model.load_state_dict(torch.load(PATH_LOAD, map_location=torch.device('cpu')))\nelse: # GPU - cuda\n    model.load_state_dict(torch.load(PATH_LOAD))\n    \nmodel.eval()","metadata":{"execution":{"iopub.status.busy":"2023-03-08T09:21:04.266183Z","iopub.execute_input":"2023-03-08T09:21:04.266546Z","iopub.status.idle":"2023-03-08T09:21:04.386162Z","shell.execute_reply.started":"2023-03-08T09:21:04.266516Z","shell.execute_reply":"2023-03-08T09:21:04.385090Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Loading","metadata":{}},{"cell_type":"code","source":"%%time\ntest_meta_df = pd.read_parquet('/kaggle/input/icecube-neutrinos-in-deep-ice/test_meta.parquet')\n# test_meta_df = pd.read_parquet('/kaggle/input/smallermeta/val_meta_660_small.parquet')\n\ntest_meta_df","metadata":{"execution":{"iopub.status.busy":"2023-03-08T09:24:57.970898Z","iopub.execute_input":"2023-03-08T09:24:57.971521Z","iopub.status.idle":"2023-03-08T09:24:57.989877Z","shell.execute_reply.started":"2023-03-08T09:24:57.971485Z","shell.execute_reply":"2023-03-08T09:24:57.988776Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_batch_ids = test_meta_df.batch_id.unique()\ntest_batch_ids","metadata":{"execution":{"iopub.status.busy":"2023-03-08T09:24:58.164385Z","iopub.execute_input":"2023-03-08T09:24:58.166486Z","iopub.status.idle":"2023-03-08T09:24:58.175044Z","shell.execute_reply.started":"2023-03-08T09:24:58.166457Z","shell.execute_reply":"2023-03-08T09:24:58.173951Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Prediction loop","metadata":{}},{"cell_type":"code","source":"bs_pred = 100 # batchsize for predictions\nmetrics = False","metadata":{"execution":{"iopub.status.busy":"2023-03-08T09:24:58.495737Z","iopub.execute_input":"2023-03-08T09:24:58.496069Z","iopub.status.idle":"2023-03-08T09:24:58.500315Z","shell.execute_reply.started":"2023-03-08T09:24:58.496032Z","shell.execute_reply":"2023-03-08T09:24:58.499333Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\nlist_azs = []\nlist_zens = []\nlist_event_ids = []\n\nfor batch_id in test_batch_ids:\n    \n    print(f'=========== START PREDICTIONS BATCH {batch_id} ===========')\n        \n    dataset = MyOwnDataset(\n        batch_id, \n#         path_batch=f'/kaggle/input/smallermeta/batch_660_small.parquet',\n        path_batch=f'/kaggle/input/icecube-neutrinos-in-deep-ice/test/batch_{batch_id}.parquet',\n#         path_meta='/kaggle/input/smallermeta/val_meta_660_small.parquet',         \n        path_meta='/kaggle/input/icecube-neutrinos-in-deep-ice/test_meta.parquet', \n        path_sensor='/kaggle/input/icecube-neutrinos-in-deep-ice/sensor_geometry.csv',\n        target_mode=target_mode, \n        K=8, \n        features=['x', 'y', 'z', 'time', 'charge'], \n        threshold_events=250,\n        targets_test=True\n    )     \n\n    data_loader = DataLoader(dataset, batch_size=bs_pred, shuffle=False)\n    \n    angle_error_sum = 0\n    \n    for id_batch, data in enumerate(tqdm(data_loader)): # Iterate over batches \n        \n        # Load data and labels to device and predict\n        events, event_ids = data\n        x, edge_index, batch = events.x.to(device), events.edge_index.to(device), events.batch.to(device)\n        \n        # for big events, do not compute\n        labels = events.y.to(device).reshape(-1, model.dim_output) # reshape bc model returns (batchsize, dim_out), while loader (idiot!) returns (dim_out*batchsize)\n\n        out = model(x, edge_index, batch) # Perform a single forward pass\n\n        # Convert preds to angles - same for labels \n        if target_mode == 'angles':\n            az_true, zen_true, az_pred, zen_pred = pred_to_angles.from_angles(out, labels)\n        if target_mode == 'cossin':        \n            az_true, zen_true, az_pred, zen_pred = pred_to_angles.from_cossin(out, labels)\n        if target_mode == 'xyz':\n            az_true, zen_true, az_pred, zen_pred = pred_to_angles.from_xyz(out, labels)\n\n        # Detach from GPU and send to CPU - convert to np to be accepted by host metric function\n        az_pred = az_pred.detach().cpu().numpy()\n        zen_pred = zen_pred.detach().cpu().numpy()\n        az_true = az_true.detach().cpu().numpy()\n        zen_true = zen_true.detach().cpu().numpy()\n            \n        # Metrics\n        if metrics:\n            angle_error = angular_dist_score(az_true, zen_true, az_pred, zen_pred)\n            angle_error_sum += angle_error * events.num_graphs\n\n            if id_batch % 100 == 0:\n                print(f'Batch {id_batch}/{len(data_loader)} - Angle error {angle_error}') \n        \n        list_event_ids.append(event_ids)\n        list_azs.append(az_pred)\n        list_zens.append(zen_pred)\n\n    if metrics:\n        print(angle_error_sum / len(data_loader.dataset))\n    \n    del dataset, data_loader\n    gc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-03-08T09:24:58.640995Z","iopub.execute_input":"2023-03-08T09:24:58.641307Z","iopub.status.idle":"2023-03-08T09:24:58.890515Z","shell.execute_reply.started":"2023-03-08T09:24:58.641282Z","shell.execute_reply":"2023-03-08T09:24:58.889353Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Convert preds to csv","metadata":{}},{"cell_type":"code","source":"submission = pd.DataFrame(\n    {\n        'event_id': np.concatenate(list_event_ids, axis=0),  \n        'azimuth': np.concatenate(list_azs, axis=0),  \n        'zenith': np.concatenate(list_zens, axis=0)\n    }\n)\nsubmission","metadata":{"execution":{"iopub.status.busy":"2023-03-08T09:24:58.961601Z","iopub.execute_input":"2023-03-08T09:24:58.961876Z","iopub.status.idle":"2023-03-08T09:24:58.976227Z","shell.execute_reply.started":"2023-03-08T09:24:58.961851Z","shell.execute_reply":"2023-03-08T09:24:58.974815Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission.to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2023-03-08T09:24:59.125493Z","iopub.execute_input":"2023-03-08T09:24:59.125797Z","iopub.status.idle":"2023-03-08T09:24:59.132194Z","shell.execute_reply.started":"2023-03-08T09:24:59.125768Z","shell.execute_reply":"2023-03-08T09:24:59.130741Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}