{"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-27T15:16:03.410718Z","iopub.execute_input":"2023-04-27T15:16:03.411128Z","iopub.status.idle":"2023-04-27T15:20:56.134906Z","shell.execute_reply.started":"2023-04-27T15:16:03.411102Z","shell.execute_reply":"2023-04-27T15:20:56.133728Z"},"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-27T15:20:56.138648Z","iopub.execute_input":"2023-04-27T15:20:56.139424Z","iopub.status.idle":"2023-04-27T15:20:56.216636Z","shell.execute_reply.started":"2023-04-27T15:20:56.139389Z","shell.execute_reply":"2023-04-27T15:20:56.215368Z"},"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-27T15:21:45.248348Z","iopub.execute_input":"2023-04-27T15:21:45.248765Z","iopub.status.idle":"2023-04-27T15:21:45.293162Z","shell.execute_reply.started":"2023-04-27T15:21:45.248714Z","shell.execute_reply":"2023-04-27T15:21:45.292300Z"},"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-27T15:21:59.012549Z","iopub.execute_input":"2023-04-27T15:21:59.012920Z","iopub.status.idle":"2023-04-27T15:21:59.018992Z","shell.execute_reply.started":"2023-04-27T15:21:59.012887Z","shell.execute_reply":"2023-04-27T15:21:59.017856Z"},"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    \nclass Azimuth(Label):\n    \"\"\"Class for producing my label.\"\"\"\n    def __init__(self):\n        \"\"\"Construct `MyCustomLabel`.\"\"\"\n        # Base class constructor\n        super().__init__(key=\"azimuth\")\n\n    def __call__(self, graph: Data) -> torch.tensor:\n        \"\"\"Compute label for `graph`.\"\"\"\n        return graph.y[1].reshape(1) # assuming y is a pandas dataframe\n    \nclass Zenith(Label):\n    \"\"\"Class for producing my label.\"\"\"\n    def __init__(self):\n        \"\"\"Construct `MyCustomLabel`.\"\"\"\n        # Base class constructor\n        super().__init__(key=\"zenith\")\n\n    def __call__(self, graph: Data) -> torch.tensor:\n        \"\"\"Compute label for `graph`.\"\"\"\n        return graph.y[0].reshape(1) # assuming y is a pandas dataframe","metadata":{"execution":{"iopub.status.busy":"2023-04-27T15:56:42.952789Z","iopub.execute_input":"2023-04-27T15:56:42.953495Z","iopub.status.idle":"2023-04-27T15:56:42.966354Z","shell.execute_reply.started":"2023-04-27T15:56:42.953457Z","shell.execute_reply":"2023-04-27T15:56:42.965389Z"},"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-27T16:19:11.792884Z","iopub.execute_input":"2023-04-27T16:19:11.793409Z","iopub.status.idle":"2023-04-27T16:19:11.820440Z","shell.execute_reply.started":"2023-04-27T16:19:11.793373Z","shell.execute_reply":"2023-04-27T16:19:11.819385Z"},"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)","metadata":{"execution":{"iopub.status.busy":"2023-04-27T16:18:50.496948Z","iopub.execute_input":"2023-04-27T16:18:50.498032Z","iopub.status.idle":"2023-04-27T16:18:50.596046Z","shell.execute_reply.started":"2023-04-27T16:18:50.497991Z","shell.execute_reply":"2023-04-27T16:18:50.595068Z"},"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-27T15:56:53.977993Z","iopub.execute_input":"2023-04-27T15:56:53.978357Z","iopub.status.idle":"2023-04-27T15:56:57.465772Z","shell.execute_reply.started":"2023-04-27T15:56:53.978328Z","shell.execute_reply":"2023-04-27T15:56:57.464769Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset.add_label(Direction())\ndataset.add_label(Zenith())\ndataset.add_label(Azimuth())","metadata":{"execution":{"iopub.status.busy":"2023-04-27T15:57:02.116495Z","iopub.execute_input":"2023-04-27T15:57:02.116905Z","iopub.status.idle":"2023-04-27T15:57:02.122547Z","shell.execute_reply.started":"2023-04-27T15:57:02.116874Z","shell.execute_reply":"2023-04-27T15:57:02.121594Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset.get(0)","metadata":{"execution":{"iopub.status.busy":"2023-04-27T15:57:06.677014Z","iopub.execute_input":"2023-04-27T15:57:06.677379Z","iopub.status.idle":"2023-04-27T15:57:07.909926Z","shell.execute_reply.started":"2023-04-27T15:57:06.677348Z","shell.execute_reply":"2023-04-27T15:57:07.908736Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import random\nfrom torch.utils.data import Subset\n\ndef random_dataset_subset(dataset_, num):\n    random.seed(10)\n    idxs = random.sample(range(0,dataset.len()), num)\n    subset_lst = []\n    for idx in idxs:\n        subset_lst.append(dataset_.get(idx))\n    return subset_lst\n    \nsubset = random_dataset_subset(dataset, 10000)\n","metadata":{"execution":{"iopub.status.busy":"2023-04-27T15:57:14.612412Z","iopub.execute_input":"2023-04-27T15:57:14.612795Z","iopub.status.idle":"2023-04-27T15:57:53.348365Z","shell.execute_reply.started":"2023-04-27T15:57:14.612753Z","shell.execute_reply":"2023-04-27T15:57:53.347356Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_subset = subset[0:8000]\ntest_subset = subset[8000:]\n\n# for test_subset:\ntrue_azimuth = []\ntrue_zenith = []\ntrue_dirs = []\nfor graph in test_subset:\n    true_zenith.append(graph['zenith'])\n    true_azimuth.append(graph['azimuth'])\n    true_dirs.append(graph['direction'])","metadata":{"execution":{"iopub.status.busy":"2023-04-27T16:20:53.370948Z","iopub.execute_input":"2023-04-27T16:20:53.371368Z","iopub.status.idle":"2023-04-27T16:20:53.384900Z","shell.execute_reply.started":"2023-04-27T16:20:53.371337Z","shell.execute_reply":"2023-04-27T16:20:53.383906Z"},"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_model',\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-27T15:58:13.127594Z","iopub.execute_input":"2023-04-27T15:58:13.128487Z","iopub.status.idle":"2023-04-27T15:58:13.135934Z","shell.execute_reply.started":"2023-04-27T15:58:13.128438Z","shell.execute_reply":"2023-04-27T15:58:13.134803Z"},"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-27T15:58:44.597676Z","iopub.execute_input":"2023-04-27T15:58:44.598398Z","iopub.status.idle":"2023-04-27T15:58:44.610465Z","shell.execute_reply.started":"2023-04-27T15:58:44.598360Z","shell.execute_reply":"2023-04-27T15:58:44.609381Z"},"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    additional_attributes = ['zenith', 'azimuth']\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-27T16:01:52.378301Z","iopub.execute_input":"2023-04-27T16:01:52.378725Z","iopub.status.idle":"2023-04-27T16:01:52.394357Z","shell.execute_reply.started":"2023-04-27T16:01:52.378693Z","shell.execute_reply":"2023-04-27T16:01:52.392942Z"},"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['truth']}\")\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    percent_test = 0.7\n    if isinstance(dataset_, List):\n        length = len(dataset_)\n    else:\n        length = dataset_.len()\n    \n    # do train-test split as:\n    train_len = int(percent_test*length)\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-27T16:01:58.717203Z","iopub.execute_input":"2023-04-27T16:01:58.717594Z","iopub.status.idle":"2023-04-27T16:01:58.726261Z","shell.execute_reply.started":"2023-04-27T16:01:58.717561Z","shell.execute_reply":"2023-04-27T16:01:58.725279Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = training_step(config, train_subset)","metadata":{"execution":{"iopub.status.busy":"2023-04-27T16:02:02.571863Z","iopub.execute_input":"2023-04-27T16:02:02.572231Z","iopub.status.idle":"2023-04-27T16:03:38.286179Z","shell.execute_reply.started":"2023-04-27T16:02:02.572200Z","shell.execute_reply":"2023-04-27T16:03:38.285235Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_loader = DataLoader(test_subset, batch_size=config['batch_size'], shuffle=False)\nresults = model.predict_as_dataframe(\n        gpus = [0],\n        dataloader = test_loader,\n        prediction_columns=model.prediction_columns,\n        additional_attributes=model.additional_attributes\n    )","metadata":{"execution":{"iopub.status.busy":"2023-04-27T16:29:14.751305Z","iopub.execute_input":"2023-04-27T16:29:14.751667Z","iopub.status.idle":"2023-04-27T16:29:16.926944Z","shell.execute_reply.started":"2023-04-27T16:29:14.751633Z","shell.execute_reply":"2023-04-27T16:29:16.926034Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"results","metadata":{"execution":{"iopub.status.busy":"2023-04-27T16:29:22.191603Z","iopub.execute_input":"2023-04-27T16:29:22.192667Z","iopub.status.idle":"2023-04-27T16:29:22.208954Z","shell.execute_reply.started":"2023-04-27T16:29:22.192614Z","shell.execute_reply":"2023-04-27T16:29:22.207895Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def convert_to_3d(df: pd.DataFrame) -> pd.DataFrame:\n    \"\"\"Converts zenith and azimuth to 3D direction vectors\"\"\"\n    df['true_x'] = np.cos(df['azimuth']) * np.sin(df['zenith'])\n    df['true_y'] = np.sin(df['azimuth'])*np.sin(df['zenith'])\n    df['true_z'] = np.cos(df['zenith'])\n    return df\n\ndef calculate_angular_error(df : pd.DataFrame) -> pd.DataFrame:\n    \"\"\"Calcualtes the opening angle (angular error) between true and reconstructed direction vectors\"\"\"\n    df['angular_error'] = np.arccos(df['true_x']*df['direction_x'] + df['true_y']*df['direction_y'] + df['true_z']*df['direction_z'])\n    return df","metadata":{"execution":{"iopub.status.busy":"2023-04-27T16:29:25.670187Z","iopub.execute_input":"2023-04-27T16:29:25.670559Z","iopub.status.idle":"2023-04-27T16:29:25.678454Z","shell.execute_reply.started":"2023-04-27T16:29:25.670525Z","shell.execute_reply":"2023-04-27T16:29:25.676978Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"results = convert_to_3d(results)\nresults = calculate_angular_error(results)\nresults.head()","metadata":{"execution":{"iopub.status.busy":"2023-04-27T16:29:29.045278Z","iopub.execute_input":"2023-04-27T16:29:29.045654Z","iopub.status.idle":"2023-04-27T16:29:29.068076Z","shell.execute_reply.started":"2023-04-27T16:29:29.045621Z","shell.execute_reply":"2023-04-27T16:29:29.067013Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"score = results['angular_error'].mean()\nscore","metadata":{"execution":{"iopub.status.busy":"2023-04-27T16:36:13.866876Z","iopub.execute_input":"2023-04-27T16:36:13.867276Z","iopub.status.idle":"2023-04-27T16:36:13.875829Z","shell.execute_reply.started":"2023-04-27T16:36:13.867244Z","shell.execute_reply":"2023-04-27T16:36:13.874729Z"},"trusted":true},"execution_count":null,"outputs":[]}]}