{"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":"# Move software to working disk\n!rm  -r software\n!scp -r /kaggle/input/graphnet-and-dependencies/software .\n\n# Install dependencies\n!pip install /kaggle/working/software/dependencies/torch-1.11.0+cu115-cp37-cp37m-linux_x86_64.whl\n!pip install /kaggle/working/software/dependencies/torch_cluster-1.6.0-cp37-cp37m-linux_x86_64.whl\n!pip install /kaggle/working/software/dependencies/torch_scatter-2.0.9-cp37-cp37m-linux_x86_64.whl\n!pip install /kaggle/working/software/dependencies/torch_sparse-0.6.13-cp37-cp37m-linux_x86_64.whl\n!pip install /kaggle/working/software/dependencies/torch_geometric-2.0.4.tar.gz\n\n!cd software/graphnet;pip install --no-index --find-links=\"/kaggle/working/software/dependencies\" -e .[torch]","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-04-26T03:53:44.263986Z","iopub.execute_input":"2023-04-26T03:53:44.264473Z","iopub.status.idle":"2023-04-26T03:58:23.787956Z","shell.execute_reply.started":"2023-04-26T03:53:44.264358Z","shell.execute_reply":"2023-04-26T03:58:23.78659Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Install GraphNeT\nimport sys\n#sys.path.append('/kaggle/input/graphnet/graphnet/src')\nsys.path.append('/kaggle/working/software/graphnet/src')\nimport graphnet","metadata":{"execution":{"iopub.status.busy":"2023-04-26T03:59:22.909541Z","iopub.execute_input":"2023-04-26T03:59:22.910734Z","iopub.status.idle":"2023-04-26T03:59:23.033143Z","shell.execute_reply.started":"2023-04-26T03:59:22.910688Z","shell.execute_reply":"2023-04-26T03:59:23.031655Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport pyarrow \n\nDATA_PATH = '/kaggle/input/icecube-neutrinos-in-deep-ice/'\nSENSORS = DATA_PATH + 'sensor_geometry.csv'\nTRANSPERANCY = '/kaggle/input/icecube-additional/ice_transperancy.txt'\n\nimport sys\nsys.path.append('/kaggle/input/icecube-utils/')\nfrom prepare_sensors import prepare_sensors\nfrom ice_transparency import ice_transparency\nsensor_df = prepare_sensors(SENSORS)\nf_scattering, f_absorption = ice_transparency(TRANSPERANCY)","metadata":{"execution":{"iopub.status.busy":"2023-04-26T03:59:25.102636Z","iopub.execute_input":"2023-04-26T03:59:25.10308Z","iopub.status.idle":"2023-04-26T03:59:26.060986Z","shell.execute_reply.started":"2023-04-26T03:59:25.103028Z","shell.execute_reply":"2023-04-26T03:59:26.059817Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"META_PATH = '/kaggle/input/batched-metadata/'\ndef get_metadata_pd(batch, write=False):\n    if batch < 661:\n        return pd.read_parquet(META_PATH + f'train_meta_batches/batch_{batch}.parquet', \n                        engine=\"pyarrow\", use_threads=True)\n    elif batch == 661:\n        return pd.read_parquet(META_PATH + f'test_meta_batches/batch_{batch}.parquet', \n                        engine=\"pyarrow\", use_threads=True)\n    ","metadata":{"execution":{"iopub.status.busy":"2023-04-26T03:59:29.082491Z","iopub.execute_input":"2023-04-26T03:59:29.083096Z","iopub.status.idle":"2023-04-26T03:59:29.090306Z","shell.execute_reply.started":"2023-04-26T03:59:29.083053Z","shell.execute_reply":"2023-04-26T03:59:29.088913Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nfrom torch_geometric.data import Data\nfrom graphnet.training.labels import Label\n\nclass Direction(Label):\n    \"\"\"Class for producing my label.\"\"\"\n    def __init__(self):\n        \"\"\"Construct `MyCustomLabel`.\"\"\"\n        # Base class constructor\n        super().__init__(key=\"direction\")\n\n    def __call__(self, graph: Data) -> torch.tensor:\n        \"\"\"Compute label for `graph`.\"\"\"\n        zenith = graph.y[0]\n        azimuth = graph.y[1] # assuming y is a pandas dataframe\n               \n        dir_x = (torch.cos(azimuth) * torch.sin(zenith)).reshape(1)\n        dir_y = (torch.sin(azimuth) * torch.sin(zenith)).reshape(1)\n        dir_z = torch.cos(zenith).reshape(1)\n        direction = torch.cat([dir_x, dir_y, dir_z], dim=0)\n        return direction\n","metadata":{"execution":{"iopub.status.busy":"2023-04-26T03:59:31.693373Z","iopub.execute_input":"2023-04-26T03:59:31.694095Z","iopub.status.idle":"2023-04-26T03:59:36.209944Z","shell.execute_reply.started":"2023-04-26T03:59:31.694059Z","shell.execute_reply":"2023-04-26T03:59:36.208779Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nfrom torch_geometric.nn import knn_graph\nfrom torch_geometric.data import Data, Dataset\nfrom torch_geometric.loader import DataLoader\nfrom typing import (\n    cast,\n    Any,\n    Callable,\n    Dict,\n    List,\n    Optional,\n    Tuple,\n    Union,\n    Iterable,\n)\n\nclass IceCubeDataset(Dataset):\n    def __init__(self, event_ids, batch_id, PATH_TO_BATCH_FILES, \n                 f_scattering, f_absorption, sensor_df, y, x_features, y_features,\n                 pulse_limit=300, include_auxiliary=True, construct_graph=False,\n                 transform = None, pre_transform=None, pre_filter=None):\n        super().__init__(transform, pre_transform, pre_filter)\n        self.event_ids = event_ids\n        self.batch_df = pd.read_parquet(PATH_TO_BATCH_FILES + f\"batch_{batch_id}.parquet\")\n        self.sensor_df = sensor_df\n        self.pulse_limit = pulse_limit\n        self.f_scattering = f_scattering\n        self.f_absorption = f_absorption\n        self.y = y\n        self.x_features = x_features\n        if include_auxiliary == False and 'auxiliary' in self.x_features:\n            self.x_features.remove('auxiliary')\n        self.include_auxiliary = include_auxiliary\n        self.y_features = y_features\n        self.construct_graph = construct_graph\n        self._label_fns = dict()\n        \n        \n        # weird scaling...really don't get any of the scaling stuff\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_dir_vector(self, azimuth, zenith):\n        dir_x = np.cos(azimuth) * np.sin(zenith)\n        dir_y = np.sin(azimuth) * np.sin(zenith)\n        dir_z = np.cos(zenith)\n        directions = pd.Series({'direction_x':dir_x, 'direction_y':dir_y, 'direction_z':dir_z})\n        return directions\n    \n    def add_label(\n        self, fn: Callable[[Data], Any], key: Optional[str] = None\n    ) -> None:\n        \"\"\"Add custom graph label define using function `fn`.\"\"\"\n        if isinstance(fn, Label):\n            key = fn.key\n        assert isinstance(\n            key, str\n        ), \"Please specify a key for the custom label to be added.\"\n        assert (\n            key not in self._label_fns\n        ), f\"A custom label {key} has already been defined.\"\n        self._label_fns[key] = fn\n\n    def get(self, idx):\n        event_id = self.event_ids[idx]\n        event = self.batch_df.loc[event_id]\n        event = pd.merge(event, self.sensor_df, on=\"sensor_id\")\n        if self.include_auxiliary == False:\n            event.drop(event[event.auxiliary == 0.5].index)\n        \n        x_feats = self.x_features.copy()\n        if 'scattering' in self.x_features:\n            x_feats.remove('scattering')\n        if 'absorption' in self.x_features:\n            x_feats.remove('absorption')\n        x = event[x_feats].values\n        x = torch.tensor(x, dtype=torch.float32)\n        data = Data(x=x, n_pulses=torch.tensor(x.shape[0], dtype=torch.int32), features=x_feats)\n\n        # Add ice transparency data\n        z = data.x[:, 2].numpy()\n        if 'scattering' in self.x_features:\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        if 'absorption' in self.x_features:\n            absorption = torch.tensor(self.f_absorption(z), dtype=torch.float32).view(-1, 1)\n            data.x = torch.cat([data.x, absorption], 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        if self.construct_graph == True:\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        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            if self._label_fns:\n                for key in self._label_fns:\n                    data[key] = self._label_fns[key](data)\n            \n            '''\n            data.azimuth = self.y.loc[idx][self.y_features].azimuth\n            data.zenith = self.y.loc[idx][self.y_features].zenith\n            dirs = self.get_dir_vector(data.azimuth, data.zenith)\n            data.direction = torch.tensor(self.get_dir_vector(data.azimuth, data.zenith).values)\n            torch.reshape(data.direction, (1,3))\n            '''\n            \n        return data","metadata":{"execution":{"iopub.status.busy":"2023-04-26T04:07:29.014395Z","iopub.execute_input":"2023-04-26T04:07:29.014826Z","iopub.status.idle":"2023-04-26T04:07:29.042292Z","shell.execute_reply.started":"2023-04-26T04:07:29.014776Z","shell.execute_reply":"2023-04-26T04:07:29.041085Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BATCH_ID = 1\nTRAIN_PATH = DATA_PATH + 'train/'\nbatch_meta = get_metadata_pd(BATCH_ID, write=False)\nevent_ids = list(batch_meta['event_id'])\n#x_feats = ['x', 'y', 'z', 'time', \"charge\", \"qe\", \"auxiliary\", 'scattering', 'absorption']\n#x_feats = ['x', 'y', 'z', 'time', \"charge\", \"auxiliary\"]\nx_feats = ['x', 'y', 'z', 'time', \"charge\", \"qe\", \"auxiliary\", 'scattering']\ny_feats = ['zenith', 'azimuth']\ny = batch_meta[y_feats].reset_index(drop=True)\ndir_x = np.cos(y.azimuth) * np.sin(y.zenith)\ndir_y = np.sin(y.azimuth) * np.sin(y.zenith)\ndir_z = np.cos(y.zenith)\ndirections = pd.concat({'direction_x':dir_x, 'direction_y':dir_y, 'direction_z':dir_z}, axis = 1)\ndirections.head()","metadata":{"execution":{"iopub.status.busy":"2023-04-26T04:07:34.247374Z","iopub.execute_input":"2023-04-26T04:07:34.247801Z","iopub.status.idle":"2023-04-26T04:07:34.386724Z","shell.execute_reply.started":"2023-04-26T04:07:34.247766Z","shell.execute_reply":"2023-04-26T04:07:34.38544Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset = IceCubeDataset(event_ids, BATCH_ID, TRAIN_PATH, f_scattering, \n                         f_absorption, sensor_df, y, x_feats, y_feats)","metadata":{"execution":{"iopub.status.busy":"2023-04-26T04:07:36.792725Z","iopub.execute_input":"2023-04-26T04:07:36.793386Z","iopub.status.idle":"2023-04-26T04:07:42.573055Z","shell.execute_reply.started":"2023-04-26T04:07:36.793339Z","shell.execute_reply":"2023-04-26T04:07:42.571918Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset.add_label(Direction())","metadata":{"execution":{"iopub.status.busy":"2023-04-26T04:07:52.018325Z","iopub.execute_input":"2023-04-26T04:07:52.018749Z","iopub.status.idle":"2023-04-26T04:07:52.064561Z","shell.execute_reply.started":"2023-04-26T04:07:52.018714Z","shell.execute_reply":"2023-04-26T04:07:52.062516Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset.get(0)['direction']","metadata":{"execution":{"iopub.status.busy":"2023-04-26T04:06:20.212146Z","iopub.execute_input":"2023-04-26T04:06:20.212561Z","iopub.status.idle":"2023-04-26T04:06:21.364335Z","shell.execute_reply.started":"2023-04-26T04:06:20.212525Z","shell.execute_reply":"2023-04-26T04:06:21.363344Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"features = x_feats\ntruth = y_feats\n\nconfig = {\n        #\"path\": '/kaggle/working/batch_1.db',\n        #\"inference_database_path\": '/kaggle/working/batch_51.db',\n        #\"pulsemap\": 'pulse_table',\n        #\"truth_table\": 'meta_table',\n        'neighbours': 8,\n        'graph_builder_columns' : [0, 1, 2], # x, y, z\n        'global_pooling_schemes' : [\"min\", \"max\", \"mean\"],\n        \"features\": features,\n        #\"truth\": truth,\n        \"index_column\": 'event_id',\n        #\"run_name_tag\": 'my_example',\n        \"batch_size\": 32,\n        \"num_workers\": 2,\n        \"target\": 'direction',\n        \"early_stopping_patience\": 5,\n        \"fit\": {\n                \"max_epochs\": 10,\n                \"gpus\": [0],\n                \"distribution_strategy\": None,\n                },\n        #'train_selection': '/kaggle/working/train_selection_max_200_pulses.csv',\n        #'validate_selection': '/kaggle/working/validate_selection_max_200_pulses.csv',\n        #'test_selection': None,\n        #'base_dir': 'training'\n}","metadata":{"execution":{"iopub.status.busy":"2023-04-26T04:08:01.212391Z","iopub.execute_input":"2023-04-26T04:08:01.212834Z","iopub.status.idle":"2023-04-26T04:08:01.21983Z","shell.execute_reply.started":"2023-04-26T04:08:01.212802Z","shell.execute_reply":"2023-04-26T04:08:01.218546Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from abc import abstractmethod\nfrom typing import Any, Optional, Union, List, Dict\n\nimport numpy as np\nimport scipy.special\nimport torch\nfrom torch import Tensor\nfrom torch import nn\nfrom torch.nn.functional import (\n    one_hot,\n    cross_entropy,\n    binary_cross_entropy,\n    softplus,\n)\n\nfrom graphnet.utilities.config import save_model_config\nfrom graphnet.models.model import Model\nfrom graphnet.utilities.decorators import final\n\n\nclass LossFunction(Model):\n    \"\"\"Base class for loss functions in `graphnet`.\"\"\"\n\n    @save_model_config\n    def __init__(self, **kwargs: Any) -> None:\n        \"\"\"Construct `LossFunction`, saving model config.\"\"\"\n        super().__init__(**kwargs)\n\n    @final\n    def forward(  # type: ignore[override]\n        self,\n        prediction: Tensor,\n        target: Tensor,\n        weights: Optional[Tensor] = None,\n        return_elements: bool = False,\n    ) -> Tensor:\n        print('in lossfunction class forward------')\n        \"\"\"Forward pass for all loss functions.\n\n        Args:\n            prediction: Tensor containing predictions. Shape [N,P]\n            target: Tensor containing targets. Shape [N,T]\n            return_elements: Whether elementwise loss terms should be returned.\n                The alternative is to return the averaged loss across examples.\n\n        Returns:\n            Loss, either averaged to a scalar (if `return_elements = False`) or\n            elementwise terms with shape [N,] (if `return_elements = True`).\n        \"\"\"\n        elements = self._forward(prediction, target)\n        if weights is not None:\n            elements = elements * weights\n            \n        if (elements.size(dim=0) != target.size(dim=0)):\n            print(elements.shape)\n            print(target.shape)\n        #assert elements.size(dim=0) == target.size(dim=0), \n\nclass VonMisesFisherLoss(LossFunction):\n    \"\"\"General class for calculating von Mises-Fisher loss.\n\n    Requires implementation for specific dimension `m` in which the target and\n    prediction vectors need to be prepared.\n    \"\"\"\n    @classmethod\n    def log_cmk_exact(\n        cls, m: int, kappa: Tensor\n    ) -> Tensor:  # pylint: disable=invalid-name\n        \"\"\"Calculate $log C_{m}(k)$ term in von Mises-Fisher loss exactly.\"\"\"\n        return LogCMK.apply(m, kappa)\n\n\n    @classmethod\n    def log_cmk_approx(\n        cls, m: int, kappa: Tensor\n    ) -> Tensor:  # pylint: disable=invalid-name\n        \"\"\"Calculate $log C_{m}(k)$ term in von Mises-Fisher loss approx.\n\n        [https://arxiv.org/abs/1812.04616] Sec. 8.2 with additional minus sign.\n        \"\"\"\n        v = m / 2.0 - 0.5\n        a = torch.sqrt((v + 1) ** 2 + kappa**2)\n        b = v - 1\n        return -a + b * torch.log(b + a)\n    def log_cmk(\n        cls, m: int, kappa: Tensor, kappa_switch: float = 100.0\n    ) -> Tensor:  # pylint: disable=invalid-name\n        \"\"\"Calculate $log C_{m}(k)$ term in von Mises-Fisher loss.\n\n        Since `log_cmk_exact` is diverges for `kappa` >~ 700 (using float64\n        precision), and since `log_cmk_approx` is unaccurate for small `kappa`,\n        this method automatically switches between the two at `kappa_switch`,\n        ensuring continuity at this point.\n        \"\"\"\n        kappa_switch = torch.tensor([kappa_switch]).to(kappa.device)\n        mask_exact = kappa < kappa_switch\n\n        # Ensure continuity at `kappa_switch`\n        offset = cls.log_cmk_approx(m, kappa_switch) - cls.log_cmk_exact(\n            m, kappa_switch\n        )\n        ret = cls.log_cmk_approx(m, kappa) - offset\n        ret[mask_exact] = cls.log_cmk_exact(m, kappa[mask_exact])\n        return ret\n\n    def _evaluate(self, prediction: Tensor, target: Tensor) -> Tensor:\n        print('in von gneral evaluate------')\n        \"\"\"Calculate von Mises-Fisher loss for a vector in D dimensons.\n\n        This loss utilises the von Mises-Fisher distribution, which is a\n        probability distribution on the (D - 1) sphere in D-dimensional space.\n\n        Args:\n            prediction: Predicted vector, of shape [batch_size, D].\n            target: Target unit vector, of shape [batch_size, D].\n\n        Returns:\n            Elementwise von Mises-Fisher loss terms.\n        \"\"\"\n        # Check(s)\n        assert prediction.dim() == 2\n        assert target.dim() == 2\n        assert prediction.size() == target.size()\n\n        # Computing loss\n        m = target.size()[1]\n        k = torch.norm(prediction, dim=1)\n        dotprod = torch.sum(prediction * target, dim=1)\n        elements = -self.log_cmk(m, k) - dotprod\n        return elements\n\n    @abstractmethod\n    def _forward(self, prediction: Tensor, target: Tensor) -> Tensor:\n        raise NotImplementedError\n        \nclass VonMisesFisher3DLoss(VonMisesFisherLoss):\n    \"\"\"von Mises-Fisher loss function vectors in the 3D plane.\"\"\"\n\n    def _forward(self, prediction: Tensor, target: Tensor) -> Tensor:\n        \"\"\"Calculate von Mises-Fisher loss for a direction in the 3D.\n\n        Args:\n            prediction: Output of the model. Must have shape [N, 4] where\n                columns 0, 1, 2 are predictions of `direction` and last column\n                is an estimate of `kappa`.\n            target: Target tensor, extracted from graph object.\n\n        Returns:\n            Elementwise von Mises-Fisher loss terms. Shape [N,]\n        \"\"\"\n        \n        print('in Von3d ---')\n        print(target.shape)\n        target = target.reshape(-1, 3)\n        print('after reshape')\n        print(target.shape)\n        # Check(s)\n        assert prediction.dim() == 2 and prediction.size()[1] == 4\n        assert target.dim() == 2\n        assert prediction.size()[0] == target.size()[0]\n\n        kappa = prediction[:, 3]\n        p = kappa.unsqueeze(1) * prediction[:, [0, 1, 2]]\n        return self._evaluate(p, target)\n        \"`_forward` should return elementwise loss terms.\"\n\n        return elements if return_elements else torch.mean(elements)","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-04-26T04:08:11.198636Z","iopub.execute_input":"2023-04-26T04:08:11.199028Z","iopub.status.idle":"2023-04-26T04:08:11.227751Z","shell.execute_reply.started":"2023-04-26T04:08:11.198994Z","shell.execute_reply":"2023-04-26T04:08:11.226517Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from abc import abstractmethod\nfrom typing import Any, Optional, Union, List, Dict\n\nimport numpy as np\nimport scipy.special\nimport torch\nfrom torch import Tensor\nfrom torch import nn\nfrom torch.nn.functional import (\n    one_hot,\n    cross_entropy,\n    binary_cross_entropy,\n    softplus,\n)\n\nfrom graphnet.utilities.config import save_model_config\nfrom graphnet.models.model import Model\nfrom graphnet.utilities.decorators import final\n\n# overriding graphnet VonMisesFischer3dLoss and parent LossFunction\nclass vMF_Loss(Model):\n    \"\"\"Base class for loss functions in `graphnet`.\"\"\"\n\n    @save_model_config\n    def __init__(self, **kwargs: Any) -> None:\n        \"\"\"Construct `LossFunction`, saving model config.\"\"\"\n        super().__init__(**kwargs)\n\n    @final\n    def forward(  # type: ignore[override]\n        self,\n        prediction: Tensor,\n        target: Tensor,\n        weights: Optional[Tensor] = None,\n        return_elements: bool = False,\n    ) -> Tensor:\n\n        target = target.reshape(-1, 3)\n        \n        eps = 1e-8\n        kappa = prediction[:, 3]      \n        logC  = -kappa + torch.log( ( kappa+eps )/( 1-torch.exp(-2*kappa)+2*eps ) )\n        p = kappa.unsqueeze(1) * prediction[:, [0, 1, 2]]\n        return -( (target*p).sum(dim=1) + logC ).mean() ","metadata":{"execution":{"iopub.status.busy":"2023-04-26T04:15:35.298433Z","iopub.execute_input":"2023-04-26T04:15:35.298861Z","iopub.status.idle":"2023-04-26T04:15:35.311437Z","shell.execute_reply.started":"2023-04-26T04:15:35.298827Z","shell.execute_reply":"2023-04-26T04:15:35.310369Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from typing import Any, Callable, List, Optional, Sequence, Tuple, Union, Dict\nfrom pytorch_lightning.callbacks import EarlyStopping\nfrom torch.optim.adam import Adam\n#from graphnet.data.constants import FEATURES, TRUTH\nfrom graphnet.models.standard_model import StandardModel\n#from graphnet.models.detector.icecube import IceCubeKaggle\nfrom graphnet.models.gnn import DynEdge\nfrom graphnet.models.graph_builders import KNNGraphBuilder\nfrom graphnet.models.task.reconstruction import DirectionReconstructionWithKappa, ZenithReconstructionWithKappa, AzimuthReconstructionWithKappa\nfrom graphnet.training.callbacks import ProgressBar, PiecewiseLinearLR\n#from graphnet.training.loss_functions import VonMisesFisher3DLoss, VonMisesFisher2DLoss\n#from graphnet.training.labels import Direction\nfrom graphnet.training.utils import make_dataloader\nfrom graphnet.utilities.logging import Logger\nfrom pytorch_lightning import Trainer\nimport pandas as pd\nfrom graphnet.models.detector.detector import Detector\n\nlogger = Logger()\n\n#override graphnet class\nclass IceCubeKaggle(Detector):\n    \"\"\"`Detector` class for Kaggle Competition.\"\"\"\n\n    # Implementing abstract class attribute\n    features = features\n\n    def _forward(self, data: Data) -> Data:\n        \"\"\"Ingest data, build graph, and preprocess features.\n        Args:\n            data: Input graph data.\n        Returns:\n            Connected and preprocessed graph data.\n        \"\"\"\n        # Check(s) --- no we want to have flexible feature inputs\n        # Preprocessing was already done\n        data_features = [features[0] for features in data.features]\n        features = data_features\n\n        return data\n\ndef build_model(config: Dict[str,Any], train_dataloader: Any) -> StandardModel:\n    \"\"\"Builds GNN from config\"\"\"\n    # Building model\n    detector = IceCubeKaggle(\n        graph_builder=KNNGraphBuilder(nb_nearest_neighbours=config['neighbours'], \n                                     columns=config['graph_builder_columns']),\n    )\n    detector.features = config['features']\n    gnn = DynEdge(\n        nb_inputs=detector.nb_outputs,\n        global_pooling_schemes=config['global_pooling_schemes'],\n    )\n\n   # if config[\"target\"] == 'direction':\n    task = DirectionReconstructionWithKappa(\n            hidden_size=gnn.nb_outputs,\n            target_labels=config[\"target\"],\n            #loss_function=VonMisesFisher3DLoss(),\n            loss_function = vMF_Loss(),\n        )\n    prediction_columns = [config[\"target\"] + \"_x\", \n                              config[\"target\"] + \"_y\", \n                              config[\"target\"] + \"_z\", \n                              config[\"target\"] + \"_kappa\" ]\n    additional_attributes = ['zenith', 'azimuth', 'event_id']\n\n    model = StandardModel(\n        detector=detector,\n        gnn=gnn,\n        tasks=[task],\n        optimizer_class=Adam,\n        optimizer_kwargs={\"lr\": 1e-03, \"eps\": 1e-03},\n        scheduler_class=PiecewiseLinearLR,\n        scheduler_kwargs={\n            \"milestones\": [\n                0,\n                len(train_dataloader) / 2,\n                len(train_dataloader) * config[\"fit\"][\"max_epochs\"],\n            ],\n            \"factors\": [1e-02, 1, 1e-02],\n        },\n        scheduler_config={\n            \"interval\": \"step\",\n        },\n    )\n    model.prediction_columns = prediction_columns\n    model.additional_attributes = additional_attributes\n    \n    return model","metadata":{"execution":{"iopub.status.busy":"2023-04-26T04:15:41.157878Z","iopub.execute_input":"2023-04-26T04:15:41.158265Z","iopub.status.idle":"2023-04-26T04:15:41.174583Z","shell.execute_reply.started":"2023-04-26T04:15:41.15823Z","shell.execute_reply":"2023-04-26T04:15:41.173335Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def training_step(config: Dict[str, Any], dataset_) -> StandardModel:\n    \"\"\"Builds and trains GNN according to config.\"\"\"\n    logger.info(f\"features: {config['features']}\")\n    logger.info(f\"truth: {config['target']}\")\n    \n    #archive = os.path.join(config['base_dir'], \"train_model_without_configs\")\n    #run_name = f\"dynedge_{config['target']}_{config['run_name_tag']}\"\n    \n    # do train-test split as:\n    train_len = int(0.7*dataset_.len())\n    train_loader = DataLoader(dataset_[:train_len], batch_size=config['batch_size'], shuffle=True, follow_batch=config['target']) # shuffle data every epoch\n    val_loader = DataLoader(dataset_[train_len:], batch_size=config['batch_size'], shuffle=False, follow_batch=config['target'])\n    \n    model = build_model(config, train_loader)\n\n    # Training model\n    callbacks = [\n        EarlyStopping(\n            monitor=\"val_loss\",\n            patience=config[\"early_stopping_patience\"],\n        ),\n        ProgressBar(),\n    ]\n\n    model.fit(\n        train_loader,\n        val_loader,\n        callbacks=callbacks,\n        **config[\"fit\"],\n    )\n    return model","metadata":{"execution":{"iopub.status.busy":"2023-04-26T04:15:44.417404Z","iopub.execute_input":"2023-04-26T04:15:44.417795Z","iopub.status.idle":"2023-04-26T04:15:44.426588Z","shell.execute_reply.started":"2023-04-26T04:15:44.417763Z","shell.execute_reply":"2023-04-26T04:15:44.425221Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = training_step(config=config)","metadata":{"execution":{"iopub.status.busy":"2023-04-26T04:15:46.997211Z","iopub.execute_input":"2023-04-26T04:15:46.997975Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.cuda.memory_summary(device=None, abbreviated=False)","metadata":{"execution":{"iopub.status.busy":"2023-04-26T03:48:31.216967Z","iopub.execute_input":"2023-04-26T03:48:31.218023Z","iopub.status.idle":"2023-04-26T03:48:31.227571Z","shell.execute_reply.started":"2023-04-26T03:48:31.217969Z","shell.execute_reply":"2023-04-26T03:48:31.226402Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def make_dataset_for_batch(BATCH_ID):\n    TRAIN_PATH = DATA_PATH + 'train/'\n    batch_meta = get_metadata_for_batch(BATCH_ID, write=False)\n    event_ids = list(batch_meta['event_id'])\n    #x_feats = ['x', 'y', 'z', 'time', \"charge\", \"qe\", \"auxiliary\", 'scattering', 'absorption']\n    #x_feats = ['x', 'y', 'z', 'time', \"charge\", \"auxiliary\"]\n    x_feats = ['x', 'y', 'z', 'time', \"charge\", \"qe\", \"auxiliary\", 'scattering']\n    y_feats = ['zenith', 'azimuth']\n    y = batch_meta[y_feats].reset_index(drop=True)\n    return dataset\n\ninference_dataset = make_dataset_for_batch(2)\n# get first 10,000 samples only for batch 2 to do inference on\ninf_loader = DataLoader(inference_dataset[:10000], batch_size=config['batch_size'], shuffle=False)\nresults = model.predict_as_dataframe(\n        gpus = [0],\n        dataloader = inf_loader,\n        prediction_columns=model.prediction_columns,\n        additional_attributes=model.additional_attributes,\n    )","metadata":{},"execution_count":null,"outputs":[]}]}