{"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":"# inference code for  [IceCube - Neutrinos in Deep Ice](https://www.kaggle.com/competitions/icecube-neutrinos-in-deep-ice/)\n\nThis notebook contains the inference code of winning solution for the Kaggle [IceCube - Neutrinos in Deep Ice](https://www.kaggle.com/competitions/icecube-neutrinos-in-deep-ice/) competition.\n\nThe source code is based on the [excellent GraphNet baseline](https://www.kaggle.com/code/rasmusrse/graphnet-example) provided by the competition host.  \nI am using some functions of GraphNet with slight modifications.\nThese modifications have not been made directly to the source files under software/graphnet/src; instead, the necessary classes have been copied into this notebook and then modified.\n\nDocumentation is available [here](https://www.kaggle.com/competitions/icecube-neutrinos-in-deep-ice/discussion/402976).\n\nThe training code is not planned to be shared, as it is executed through multiple commands and complicated.","metadata":{}},{"cell_type":"code","source":"import os\nimport sys\nif os.path.exists('/home/tito'):\n    sys.path.append('/home/tito/kaggle/icecube-neutrinos-in-deep-ice/notebook_graphnet/software/graphnet/src')\n    KAGGLE_ENV = False\nelse:\n    KAGGLE_ENV = True\n    sys.path.append('/kaggle/working/software/graphnet/src')\n\nif KAGGLE_ENV:\n    # 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    # Install GraphNeT\n    !cd software/graphnet;pip install --no-index --find-links=\"/kaggle/working/software/dependencies\" -e .[torch]","metadata":{"lines_to_next_cell":2,"papermill":{"duration":279.435011,"end_time":"2023-04-21T01:06:47.206817","exception":false,"start_time":"2023-04-21T01:02:07.771806","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-22T23:49:06.826876Z","iopub.execute_input":"2023-04-22T23:49:06.827220Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from abc import ABC, abstractmethod\nfrom collections import OrderedDict\nfrom copy import deepcopy\nfrom graphnet.constants import GRAPHNET_ROOT_DIR\nfrom graphnet.data.constants import FEATURES, TRUTH\nfrom graphnet.data.parquet import ParquetDataset\nfrom graphnet.data.utilities.string_selection_resolver import StringSelectionResolver\nfrom graphnet.models import Model\nfrom graphnet.models import StandardModel\nfrom graphnet.models.coarsening import Coarsening\nfrom graphnet.models.components.layers import DynEdgeConv\nfrom graphnet.models.detector.detector import Detector\nfrom graphnet.models.detector.icecube import IceCubeKaggle\nfrom graphnet.models.gnn.gnn import GNN\nfrom graphnet.models.graph_builders import KNNGraphBuilder\nfrom graphnet.models.model import Model\nfrom graphnet.models.task import Task\nfrom graphnet.models.task.reconstruction import DirectionReconstructionWithKappa, ZenithReconstructionWithKappa, AzimuthReconstructionWithKappa\nfrom graphnet.models.utils import calculate_distance_matrix\nfrom graphnet.models.utils import calculate_xyzt_homophily\nfrom graphnet.training.callbacks import ProgressBar, PiecewiseLinearLR\nfrom graphnet.training.labels import Direction\nfrom graphnet.training.loss_functions import VonMisesFisher3DLoss, VonMisesFisher2DLoss, LossFunction, VonMisesFisherLoss\nfrom graphnet.training.utils import make_dataloader\nfrom graphnet.utilities.config import Configurable, DatasetConfig, save_dataset_config\nfrom graphnet.utilities.config import save_model_config\nfrom graphnet.utilities.logging import LoggerMixin\nfrom graphnet.utilities.logging import get_logger\nfrom pytorch_lightning import LightningModule\nfrom pytorch_lightning import Trainer\nfrom pytorch_lightning import loggers as pl_loggers\nfrom pytorch_lightning.callbacks import EarlyStopping\nfrom pytorch_lightning.callbacks import LearningRateMonitor\nfrom pytorch_lightning.callbacks import EarlyStopping, ModelCheckpoint\nfrom pytorch_lightning.callbacks import GradientAccumulationScheduler\nfrom pytorch_lightning.profiler import PyTorchProfiler\nfrom torch import Tensor\nfrom torch import Tensor, LongTensor\nfrom torch import nn\nfrom torch.functional import Tensor\nfrom torch.nn import ModuleList\nfrom torch.nn.modules import TransformerEncoder, TransformerEncoderLayer\nfrom torch.nn.modules.normalization import LayerNorm\nfrom torch.optim import Adam\nfrom torch.optim.adam import Adam\nfrom torch.utils.data import ConcatDataset\nfrom torch.utils.data import DataLoader\nfrom torch.utils.data import Dataset\nfrom torch_geometric.data import Batch, Data\nfrom torch_geometric.data import Data\nfrom torch_geometric.nn import knn_graph, radius_graph\nfrom torch_geometric.nn.conv import MessagePassing\nfrom torch_geometric.nn.inits import reset\nfrom torch_geometric.nn.pool import knn_graph\nfrom torch_geometric.typing import Adj\nfrom torch_geometric.typing import Adj, OptTensor, PairOptTensor, PairTensor\nfrom torch_geometric.utils import to_dense_batch\nfrom torch_scatter import scatter_max, scatter_mean, scatter_min, scatter_sum\nfrom tqdm import tqdm\nfrom typing import cast, Any, Callable, Optional, Sequence, Union, Dict, List, Optional, Union, Tuple\nimport gc\nimport graphnet\nimport numpy as np\nimport os\nimport pandas as pd\nimport random\nimport socket\nimport sys\nimport torch\n\ntry:\n    from torch_cluster import knn\nexcept ImportError:\n    knn = None","metadata":{"execution":{"iopub.status.busy":"2023-04-22T13:23:14.771500Z","iopub.execute_input":"2023-04-22T13:23:14.771940Z","iopub.status.idle":"2023-04-22T13:23:20.327153Z","shell.execute_reply.started":"2023-04-22T13:23:14.771879Z","shell.execute_reply":"2023-04-22T13:23:20.326095Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nclass ColumnMissingException(Exception):\n    \"\"\"Exception to indicate a missing column in a dataset.\"\"\"\n\n\nclass Dataset2(torch.utils.data.Dataset, Configurable, LoggerMixin, ABC):\n    \"\"\"Base Dataset class for reading from any intermediate file format.\"\"\"\n\n    # Class method(s)\n    @classmethod\n    def from_config(  # type: ignore[override]\n        cls,\n        source: Union[DatasetConfig, str],\n    ) -> Union[\n        \"Dataset\",\n        ConcatDataset,\n        Dict[str, \"Dataset\"],\n        Dict[str, ConcatDataset],\n    ]:\n        \"\"\"Construct `Dataset` instance from `source` configuration.\"\"\"\n        if isinstance(source, str):\n            source = DatasetConfig.load(source)\n\n        assert isinstance(source, DatasetConfig), (\n            f\"Argument `source` of type ({type(source)}) is not a \"\n            \"`DatasetConfig\"\n        )\n\n        # Parse set of `selection``.\n        if isinstance(source.selection, dict):\n            return cls._construct_datasets_from_dict(source)\n        elif (\n            isinstance(source.selection, list)\n            and len(source.selection)\n            and isinstance(source.selection[0], str)\n        ):\n            return cls._construct_dataset_from_list_of_strings(source)\n\n        return source._dataset_class(**source.dict())\n\n    @classmethod\n    def concatenate(\n        cls,\n        datasets: List[\"Dataset\"],\n    ) -> ConcatDataset:\n        \"\"\"Concatenate multiple `Dataset`s into one instance.\"\"\"\n        return ConcatDataset(datasets)\n\n    @classmethod\n    def _construct_datasets_from_dict(\n        cls, config: DatasetConfig\n    ) -> Dict[str, \"Dataset\"]:\n        \"\"\"Construct `Dataset` for each entry in dict `self.selection`.\"\"\"\n        assert isinstance(config.selection, dict)\n        datasets: Dict[str, \"Dataset\"] = {}\n        selections: Dict[str, Union[str, List]] = deepcopy(config.selection)\n        for key, selection in selections.items():\n            config.selection = selection\n            dataset = Dataset.from_config(config)\n            assert isinstance(dataset, (Dataset, ConcatDataset))\n            datasets[key] = dataset\n\n        # Reset `selections`.\n        config.selection = selections\n\n        return datasets\n\n    @classmethod\n    def _construct_dataset_from_list_of_strings(\n        cls, config: DatasetConfig\n    ) -> \"Dataset\":\n        \"\"\"Construct `Dataset` for each entry in list `self.selection`.\"\"\"\n        assert isinstance(config.selection, list)\n        datasets: List[\"Dataset\"] = []\n        selections: List[str] = deepcopy(cast(List[str], config.selection))\n        for selection in selections:\n            config.selection = selection\n            dataset = Dataset.from_config(config)\n            assert isinstance(dataset, Dataset)\n            datasets.append(dataset)\n\n        # Reset `selections`.\n        config.selection = selections\n\n        return cls.concatenate(datasets)\n\n    @classmethod\n    def _resolve_graphnet_paths(\n        cls, path: Union[str, List[str]]\n    ) -> Union[str, List[str]]:\n        if isinstance(path, list):\n            return [cast(str, cls._resolve_graphnet_paths(p)) for p in path]\n\n        assert isinstance(path, str)\n        return (\n            path.replace(\"$graphnet\", GRAPHNET_ROOT_DIR)\n            .replace(\"$GRAPHNET\", GRAPHNET_ROOT_DIR)\n            .replace(\"${graphnet}\", GRAPHNET_ROOT_DIR)\n            .replace(\"${GRAPHNET}\", GRAPHNET_ROOT_DIR)\n        )\n\n    @save_dataset_config\n    def __init__(\n        self,\n        path: Union[str, List[str]],\n        pulsemaps: Union[str, List[str]],\n        features: List[str],\n        truth: List[str],\n        *,\n        node_truth: Optional[List[str]] = None,\n        index_column: str = \"event_no\",\n        truth_table: str = \"truth\",\n        node_truth_table: Optional[str] = None,\n        string_selection: Optional[List[int]] = None,\n        selection: Optional[Union[str, List[int], List[List[int]]]] = None,\n        dtype: torch.dtype = torch.float32,\n        loss_weight_table: Optional[str] = None,\n        loss_weight_column: Optional[str] = None,\n        loss_weight_default_value: Optional[float] = None,\n        seed: Optional[int] = None,\n    ):\n        \"\"\"Construct Dataset.\n\n        Args:\n            path: Path to the file(s) from which this `Dataset` should read.\n            pulsemaps: Name(s) of the pulse map series that should be used to\n                construct the nodes on the individual graph objects, and their\n                features. Multiple pulse series maps can be used, e.g., when\n                different DOM types are stored in different maps.\n            features: List of columns in the input files that should be used as\n                node features on the graph objects.\n            truth: List of event-level columns in the input files that should\n                be used added as attributes on the  graph objects.\n            node_truth: List of node-level columns in the input files that\n                should be used added as attributes on the graph objects.\n            index_column: Name of the column in the input files that contains\n                unique indicies to identify and map events across tables.\n            truth_table: Name of the table containing event-level truth\n                information.\n            node_truth_table: Name of the table containing node-level truth\n                information.\n            string_selection: Subset of strings for which data should be read\n                and used to construct graph objects. Defaults to None, meaning\n                all strings for which data exists are used.\n            selection: The events that should be read. This can be given either\n                as list of indicies (in `index_column`); or a string-based\n                selection used to query the `Dataset` for events passing the\n                selection. Defaults to None, meaning that all events in the\n                input files are read.\n            dtype: Type of the feature tensor on the graph objects returned.\n            loss_weight_table: Name of the table containing per-event loss\n                weights.\n            loss_weight_column: Name of the column in `loss_weight_table`\n                containing per-event loss weights. This is also the name of the\n                corresponding attribute assigned to the graph object.\n            loss_weight_default_value: Default per-event loss weight.\n                NOTE: This default value is only applied when\n                `loss_weight_table` and `loss_weight_column` are specified, and\n                in this case to events with no value in the corresponding\n                table/column. That is, if no per-event loss weight table/column\n                is provided, this value is ignored. Defaults to None.\n            seed: Random number generator seed, used for selecting a random\n                subset of events when resolving a string-based selection (e.g.,\n                `\"10000 random events ~ event_no % 5 > 0\"` or `\"20% random\n                events ~ event_no % 5 > 0\"`).\n        \"\"\"\n        # Check(s)\n        if isinstance(pulsemaps, str):\n            pulsemaps = [pulsemaps]\n\n        assert isinstance(features, (list, tuple))\n        assert isinstance(truth, (list, tuple))\n\n        # Resolve reference to `$GRAPHNET` in path(s)\n        path = self._resolve_graphnet_paths(path)\n\n        # Member variable(s)\n        self._path = path\n        self._selection = None\n        self._pulsemaps = pulsemaps\n        self._features = [index_column] + features\n        self._truth = [index_column] + truth\n        self._index_column = index_column\n        self._truth_table = truth_table\n        self._loss_weight_default_value = loss_weight_default_value\n\n        if node_truth is not None:\n            assert isinstance(node_truth_table, str)\n            if isinstance(node_truth, str):\n                node_truth = [node_truth]\n\n        self._node_truth = node_truth\n        self._node_truth_table = node_truth_table\n\n        if string_selection is not None:\n            self.warning(\n                (\n                    \"String selection detected.\\n \"\n                    f\"Accepted strings: {string_selection}\\n \"\n                    \"All other strings are ignored!\"\n                )\n            )\n            if isinstance(string_selection, int):\n                string_selection = [string_selection]\n\n        self._string_selection = string_selection\n\n        self._selection = None\n        if self._string_selection:\n            self._selection = f\"string in {str(tuple(self._string_selection))}\"\n\n        self._loss_weight_column = loss_weight_column\n        self._loss_weight_table = loss_weight_table\n        if (self._loss_weight_table is None) and (\n            self._loss_weight_column is not None\n        ):\n            self.warning(\"Error: no loss weight table specified\")\n            assert isinstance(self._loss_weight_table, str)\n        if (self._loss_weight_table is not None) and (\n            self._loss_weight_column is None\n        ):\n            self.warning(\"Error: no loss weight column specified\")\n            assert isinstance(self._loss_weight_column, str)\n\n        self._dtype = dtype\n\n        self._label_fns: Dict[str, Callable[[Data], Any]] = {}\n\n        self._string_selection_resolver = StringSelectionResolver(\n            self,\n            index_column=index_column,\n            seed=seed,\n        )\n\n        # Implementation-specific initialisation.\n        self._init()\n\n        # Set unique indices\n        self._indices: Union[List[int], List[List[int]]]\n        if selection is None:\n            self._indices = self._get_all_indices()\n        elif isinstance(selection, str):\n            self._indices = self._resolve_string_selection_to_indices(\n                selection\n            )\n        else:\n            self._indices = selection\n\n        # Purely internal member variables\n        self._missing_variables: Dict[str, List[str]] = {}\n        self._remove_missing_columns()\n\n        # Implementation-specific post-init code.\n        self._post_init()\n\n        # Base class constructor\n        super().__init__()\n\n    # Properties\n    @property\n    def path(self) -> Union[str, List[str]]:\n        \"\"\"Path to the file(s) from which this `Dataset` reads.\"\"\"\n        return self._path\n\n    @property\n    def truth_table(self) -> str:\n        \"\"\"Name of the table containing event-level truth information.\"\"\"\n        return self._truth_table\n\n    # Abstract method(s)\n    @abstractmethod\n    def _init(self) -> None:\n        \"\"\"Set internal representation needed to read data from input file.\"\"\"\n\n    def _post_init(self) -> None:\n        \"\"\"Implemenation-specific code to be run after the main constructor.\"\"\"\n\n    @abstractmethod\n    def _get_all_indices(self) -> List[int]:\n        \"\"\"Return a list of all available values in `self._index_column`.\"\"\"\n\n    @abstractmethod\n    def _get_event_index(\n        self, sequential_index: Optional[int]\n    ) -> Optional[int]:\n        \"\"\"Return a the event index corresponding to a `sequential_index`.\"\"\"\n\n    @abstractmethod\n    def query_table(\n        self,\n        table: str,\n        columns: Union[List[str], str],\n        sequential_index: Optional[int] = None,\n        selection: Optional[str] = None,\n    ) -> List[Tuple[Any, ...]]:\n        \"\"\"Query a table at a specific index, optionally with some selection.\n\n        Args:\n            table: Table to be queried.\n            columns: Columns to read out.\n            sequential_index: Sequentially numbered index\n                (i.e. in [0,len(self))) of the event to query. This _may_\n                differ from the indexation used in `self._indices`. If no value\n                is provided, the entire column is returned.\n            selection: Selection to be imposed before reading out data.\n                Defaults to None.\n\n        Returns:\n            List of tuples containing the values in `columns`. If the `table`\n                contains only scalar data for `columns`, a list of length 1 is\n                returned\n\n        Raises:\n            ColumnMissingException: If one or more element in `columns` is not\n                present in `table`.\n        \"\"\"\n\n    # Public method(s)\n    def add_label(self, key: str, fn: Callable[[Data], Any]) -> None:\n        \"\"\"Add custom graph label define using function `fn`.\"\"\"\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 __len__(self) -> int:\n        \"\"\"Return number of graphs in `Dataset`.\"\"\"\n        return len(self._indices)\n\n    def __getitem__(self, sequential_index: int) -> Data:\n        \"\"\"Return graph `Data` object at `index`.\"\"\"\n        if not (0 <= sequential_index < len(self)):\n            raise IndexError(\n                f\"Index {sequential_index} not in range [0, {len(self) - 1}]\"\n            )\n        #import pdb;pdb.set_trace()\n        features, truth, node_truth, loss_weight = self._query(\n            sequential_index\n        )\n        graph = self._create_graph(features, truth, node_truth, loss_weight)\n        return graph\n\n    # Internal method(s)\n    def _resolve_string_selection_to_indices(\n        self, selection: str\n    ) -> List[int]:\n        \"\"\"Resolve selection as string to list of indicies.\n\n        Selections are expected to have pandas.DataFrame.query-compatible\n        syntax, e.g., ``` \"event_no % 5 > 0\" ``` Selections may also specify a\n        fixed number of events to randomly sample, e.g., ``` \"10000 random\n        events ~ event_no % 5 > 0\" \"20% random events ~ event_no % 5 > 0\" ```\n        \"\"\"\n        return self._string_selection_resolver.resolve(selection)\n\n    def _remove_missing_columns(self) -> None:\n        \"\"\"Remove columns that are not present in the input file.\n\n        Columns are removed from `self._features` and `self._truth`.\n        \"\"\"\n        # Check if table is completely empty\n        if len(self) == 0:\n            self.warning(\"Dataset is empty.\")\n            return\n\n        # Find missing features\n        missing_features_set = set(self._features)\n        for pulsemap in self._pulsemaps:\n            missing = self._check_missing_columns(self._features, pulsemap)\n            missing_features_set = missing_features_set.intersection(missing)\n\n        missing_features = list(missing_features_set)\n\n        # Find missing truth variables\n        missing_truth_variables = self._check_missing_columns(\n            self._truth, self._truth_table\n        )\n\n        # Remove missing features\n        if missing_features:\n            self.warning(\n                \"Removing the following (missing) features: \"\n                + \", \".join(missing_features)\n            )\n            for missing_feature in missing_features:\n                self._features.remove(missing_feature)\n\n        # Remove missing truth variables\n        if missing_truth_variables:\n            self.warning(\n                (\n                    \"Removing the following (missing) truth variables: \"\n                    + \", \".join(missing_truth_variables)\n                )\n            )\n            for missing_truth_variable in missing_truth_variables:\n                self._truth.remove(missing_truth_variable)\n\n    def _check_missing_columns(\n        self,\n        columns: List[str],\n        table: str,\n    ) -> List[str]:\n        \"\"\"Return a list missing columns in `table`.\"\"\"\n        for column in columns:\n            try:\n                self.query_table(table, [column], 0)\n            except ColumnMissingException:\n                if table not in self._missing_variables:\n                    self._missing_variables[table] = []\n                self._missing_variables[table].append(column)\n            except IndexError:\n                self.warning(f\"Dataset contains no entries for {column}\")\n            except:\n                if table not in self._missing_variables:\n                    self._missing_variables[table] = []\n                self._missing_variables[table].append(column)\n\n        return self._missing_variables.get(table, [])\n\n    def _query(\n        self, sequential_index: int\n    ) -> Tuple[\n        List[Tuple[float, ...]],\n        Tuple[Any, ...],\n        Optional[List[Tuple[Any, ...]]],\n        Optional[float],\n    ]:\n        \"\"\"Query file for event features and truth information.\n\n        The returned lists have lengths correspondings to the number of pulses\n        in the event. Their constituent tuples have lengths corresponding to\n        the number of features/attributes in each output\n\n        Args:\n            sequential_index: Sequentially numbered index\n                (i.e. in [0,len(self))) of the event to query. This _may_\n                differ from the indexation used in `self._indices`.\n\n        Returns:\n            Tuple containing pulse-level event features; event-level truth\n                information; pulse-level truth information; and event-level\n                loss weights, respectively.\n        \"\"\"\n        #import pdb;pdb.set_trace()\n        features = []\n        for pulsemap in self._pulsemaps:\n            features_pulsemap = self.query_table(\n                pulsemap, self._features, sequential_index, self._selection\n            )\n            features.extend(features_pulsemap)\n\n        truth: Tuple[Any, ...] = self.query_table(\n            self._truth_table, self._truth, sequential_index\n        )# [0] # updated\n        if self._node_truth:\n            assert self._node_truth_table is not None\n            node_truth = self.query_table(\n                self._node_truth_table,\n                self._node_truth,\n                sequential_index,\n                self._selection,\n            )\n        else:\n            node_truth = None\n\n        loss_weight: Optional[float] = None  # Default\n        if self._loss_weight_column is not None:\n            assert self._loss_weight_table is not None\n            loss_weight_list = self.query_table(\n                self._loss_weight_table,\n                self._loss_weight_column,\n                sequential_index,\n            )\n            if len(loss_weight_list):\n                loss_weight = loss_weight_list[0][0]\n            else:\n                loss_weight = -1.0\n\n        return features, truth, node_truth, loss_weight\n\n    def _create_graph(\n        self,\n        features: List[Tuple[float, ...]],\n        truth: Tuple[Any, ...],\n        node_truth: Optional[List[Tuple[Any, ...]]] = None,\n        loss_weight: Optional[float] = None,\n    ) -> Data:\n        \"\"\"Create Pytorch Data (i.e. graph) object.\n\n        No preprocessing is performed at this stage, just as no node adjancency\n        is imposed. This means that the `edge_attr` and `edge_weight`\n        attributes are not set.\n\n        Args:\n            features: List of tuples, containing event features.\n            truth: List of tuples, containing truth information.\n            node_truth: List of tuples, containing node-level truth.\n            loss_weight: A weight associated with the event for weighing the\n                loss.\n\n        Returns:\n            Graph object.\n        \"\"\"\n        # Convert nested list to simple dict\n        truth_dict = {\n            key: truth[index] for index, key in enumerate(self._truth)\n        }\n\n        # Define custom labels\n        labels_dict = self._get_labels(truth_dict)\n\n        # Convert nested list to simple dict\n        if node_truth is not None:\n            node_truth_array = np.asarray(node_truth)\n            assert self._node_truth is not None\n            node_truth_dict = {\n                key: node_truth_array[:, index]\n                for index, key in enumerate(self._node_truth)\n            }\n\n        # updated\n        # Catch cases with no reconstructed pulses\n        if len(features):\n            data = np.asarray(features)[:, 1:]\n        else:\n            data = np.array([]).reshape((0, len(self._features) - 1))\n        #data = features[:, 1:]\n\n        # Construct graph data object\n        x = torch.tensor(data.astype('float32'), dtype=self._dtype)  # pylint: disable=C0103\n        n_pulses = torch.tensor(len(x), dtype=torch.int32)\n        graph = Data(x=x, edge_index=None)\n        graph.n_pulses = n_pulses\n        graph.features = self._features[1:]\n\n        # Add loss weight to graph.\n        if loss_weight is not None and self._loss_weight_column is not None:\n            # No loss weight was retrieved, i.e., it is missing for the current\n            # event.\n            if loss_weight < 0:\n                if self._loss_weight_default_value is None:\n                    raise ValueError(\n                        \"At least one event is missing an entry in \"\n                        f\"{self._loss_weight_column} \"\n                        \"but loss_weight_default_value is None.\"\n                    )\n                graph[self._loss_weight_column] = torch.tensor(\n                    self._loss_weight_default_value, dtype=self._dtype\n                ).reshape(-1, 1)\n            else:\n                graph[self._loss_weight_column] = torch.tensor(\n                    loss_weight, dtype=self._dtype\n                ).reshape(-1, 1)\n\n        # Write attributes, either target labels, truth info or original\n        # features.\n        add_these_to_graph = [labels_dict, truth_dict]\n        if node_truth is not None:\n            add_these_to_graph.append(node_truth_dict)\n        for write_dict in add_these_to_graph:\n            for key, value in write_dict.items():\n                try:\n                    graph[key] = torch.tensor(value)\n                except TypeError:\n                    # Cannot convert `value` to Tensor due to its data type,\n                    # e.g. `str`.\n                    self.debug(\n                        (\n                            f\"Could not assign `{key}` with type \"\n                            f\"'{type(value).__name__}' as attribute to graph.\"\n                        )\n                    )\n\n        # Additionally add original features as (static) attributes\n        for index, feature in enumerate(graph.features):\n            if feature not in [\"x\"]:\n                graph[feature] = graph.x[:, index].detach()\n\n        # Add custom labels to the graph\n        for key, fn in self._label_fns.items():\n            graph[key] = fn(graph)\n        return graph\n\n    def _get_labels(self, truth_dict: Dict[str, Any]) -> Dict[str, Any]:\n        \"\"\"Return dictionary of  labels, to be added as graph attributes.\"\"\"\n        if \"pid\" in truth_dict.keys():\n            abs_pid = abs(truth_dict[\"pid\"])\n            sim_type = truth_dict[\"sim_type\"]\n\n            labels_dict = {\n                self._index_column: truth_dict[self._index_column],\n                \"muon\": int(abs_pid == 13),\n                \"muon_stopped\": int(truth_dict.get(\"stopped_muon\") == 1),\n                \"noise\": int((abs_pid == 1) & (sim_type != \"data\")),\n                \"neutrino\": int(\n                    (abs_pid != 13) & (abs_pid != 1)\n                ),  # @TODO: `abs_pid in [12,14,16]`?\n                \"v_e\": int(abs_pid == 12),\n                \"v_u\": int(abs_pid == 14),\n                \"v_t\": int(abs_pid == 16),\n                \"track\": int(\n                    (abs_pid == 14) & (truth_dict[\"interaction_type\"] == 1)\n                ),\n                \"dbang\": self._get_dbang_label(truth_dict),\n                \"corsika\": int(abs_pid > 20),\n            }\n        else:\n            labels_dict = {\n                self._index_column: truth_dict[self._index_column],\n                \"muon\": -1,\n                \"muon_stopped\": -1,\n                \"noise\": -1,\n                \"neutrino\": -1,\n                \"v_e\": -1,\n                \"v_u\": -1,\n                \"v_t\": -1,\n                \"track\": -1,\n                \"dbang\": -1,\n                \"corsika\": -1,\n            }\n        return labels_dict\n\n    def _get_dbang_label(self, truth_dict: Dict[str, Any]) -> int:\n        \"\"\"Get label for double-bang classification.\"\"\"\n        try:\n            label = int(truth_dict[\"dbang_decay_length\"] > -1)\n            return label\n        except KeyError:\n            return -1","metadata":{"papermill":{"duration":0.085089,"end_time":"2023-04-21T01:06:47.338128","exception":false,"start_time":"2023-04-21T01:06:47.253039","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-22T13:23:20.329173Z","iopub.execute_input":"2023-04-22T13:23:20.329965Z","iopub.status.idle":"2023-04-22T13:23:20.398038Z","shell.execute_reply.started":"2023-04-22T13:23:20.329904Z","shell.execute_reply":"2023-04-22T13:23:20.396977Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ParquetDataset2(Dataset2):\n    \"\"\"Pytorch dataset for reading from Parquet files.\"\"\"\n    \n    def __init__(\n        self,\n        path: Union[str, List[str]],\n        pulsemaps: Union[str, List[str]],\n        features: List[str],\n        truth: List[str],\n        batch_ids = List[int],\n        *,\n        node_truth: Optional[List[str]] = None,\n        index_column: str = \"event_no\",\n        truth_table: str = \"truth\",\n        node_truth_table: Optional[str] = None,\n        string_selection: Optional[List[int]] = None,\n        selection: Optional[Union[str, List[int], List[List[int]]]] = None,\n        dtype: torch.dtype = torch.float32,\n        loss_weight_table: Optional[str] = None,\n        loss_weight_column: Optional[str] = None,\n        loss_weight_default_value: Optional[float] = None,\n        seed: Optional[int] = None,\n        max_len:  Optional[int] = 0,\n        max_pulse:  Optional[int] = 200,\n        min_pulse:  Optional[int] = 0,\n        \n    ):\n        self.batch_ids = batch_ids\n        self.max_len = max_len\n        self.max_pulse = max_pulse\n        self.min_pulse = min_pulse\n        self.this_batch_id = 0\n        \n        super().__init__(\n        path=path,\n        pulsemaps=pulsemaps,\n        features=features,\n        truth=truth,\n        selection=selection,\n        node_truth=node_truth,\n        truth_table=truth_table,\n        node_truth_table=node_truth_table,\n        string_selection=string_selection,\n        loss_weight_table=loss_weight_table,\n        loss_weight_column=loss_weight_column,\n        index_column=index_column,\n        )\n        \n    def reset_epoch(self) -> None:\n        self.this_batch_idx += 1\n        if self.this_batch_idx >= len(self.batch_ids):\n            self.this_batch_idx = 0\n            \n        if self.this_batch_id == self.batch_ids[self.this_batch_idx]:\n            print('skip reset epoch ', self.this_batch_id, self.this_batch_idx)\n            return\n        else:\n            self.this_batch_id = self.batch_ids[self.this_batch_idx]\n            print('reset epoch to batch_id:', self.this_batch_id, self.this_batch_idx)\n\n        #print('reading meta', f'../input/train/meta_{self.this_batch_id}.parquet')\n        meta = pd.read_parquet(f'{META_DIR}/meta_{self.this_batch_id}.parquet')\n            \n        if self.max_pulse > 0:\n            pulse_count = meta.last_pulse_index - meta.first_pulse_index +1\n            meta = meta[pulse_count<self.max_pulse].reset_index(drop=True)\n        if self.min_pulse > 0:\n            pulse_count = meta.last_pulse_index - meta.first_pulse_index +1\n            meta = meta[pulse_count>=self.min_pulse].reset_index(drop=True)\n\n            \n        self.meta_name_index = {n:i for i,n in enumerate(meta.columns)}\n        self.this_meta_arr = meta.values\n        \n        batch = pd.read_parquet(f\"{BATCH_DIR}/batch_{self.this_batch_id}.parquet\").reset_index()\n        batch['x'] = batch['sensor_id'].map(self.sensor_geometry_dict['x'])\n        batch['y'] = batch['sensor_id'].map(self.sensor_geometry_dict['y'])\n        batch['z'] = batch['sensor_id'].map(self.sensor_geometry_dict['z'])\n        self.batch_name_index = {n:i for i,n in enumerate(batch.columns)}\n        self.this_batch = batch\n        \n    def _init(self) -> None:\n        self.this_batch_idx = -1\n        sensor_geometry = pd.read_csv(f'../input/icecube-neutrinos-in-deep-ice/sensor_geometry.csv')\n        self.sensor_geometry_dict = sensor_geometry.set_index('sensor_id').to_dict()\n\n        self.reset_epoch()\n        #print('done')\n        \n    def __len__(self) -> int:\n        \"\"\"Return number of graphs in `Dataset`.\"\"\"\n        if self.max_len > 0:\n            return self.max_len\n        else:\n            return len(self.this_meta_arr)\n\n    def _get_all_indices(self) -> List[int]:\n        return range(len(self.this_meta_arr))\n\n    def _get_event_index(sequential_index):\n        return self.this_meta_arr.iloc[sequential_index,self.meta_name_index['event_id']]\n\n    def query_table(\n        self,\n        table: str,\n        columns: Union[List[str], str],\n        sequential_index: Optional[int] = None,\n        selection: Optional[str] = None,\n    ) -> List[Tuple[Any, ...]]:\n        if table == 'pulse_table':\n            #columns = [self.batch_name_index[c] for c in columns]\n            first_pulse_index = int(self.this_meta_arr[sequential_index,self.meta_name_index['first_pulse_index']])\n            last_pulse_index = int(self.this_meta_arr[sequential_index,self.meta_name_index['last_pulse_index']])\n            last_pulse_index = min(first_pulse_index+FORCE_MAX_PULSE, last_pulse_index)\n            this_batch = self.this_batch[first_pulse_index:last_pulse_index+1]\n            if ONLY_AUX_FALSE:\n                this_batch = this_batch[this_batch.auxiliary == False]\n            if len(this_batch)==0:\n                new_sequential_index = np.random.randint(self.__len__())\n                print('Warning: batch len is 0 for sequential_index, new_sequential_index', sequential_index, new_sequential_index)\n                return self.query_table(table, columns, new_sequential_index, selection)\n            return this_batch[columns].values\n        else:\n            columns = [self.meta_name_index[c] for c in columns]\n            return self.this_meta_arr[sequential_index, columns]","metadata":{"papermill":{"duration":0.0446,"end_time":"2023-04-21T01:06:52.120289","exception":false,"start_time":"2023-04-21T01:06:52.075689","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-22T13:23:20.399671Z","iopub.execute_input":"2023-04-22T13:23:20.400452Z","iopub.status.idle":"2023-04-22T13:23:20.422853Z","shell.execute_reply.started":"2023-04-22T13:23:20.400414Z","shell.execute_reply":"2023-04-22T13:23:20.421783Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n\ndef collate_fn(graphs: List[Data]) -> Batch:\n    \"\"\"Remove graphs with less than two DOM hits.\n\n    Should not occur in \"production.\n    \"\"\"\n    graphs = [g for g in graphs if g.n_pulses > 1]\n    return Batch.from_data_list(graphs)\n\n# @TODO: Remove in favour of DataLoader{,.from_dataset_config}\ndef make_dataloader2(\n    db: str,\n    pulsemaps: Union[str, List[str]],\n    features: List[str],\n    truth: List[str],\n    batch_ids = List[int],\n    max_len = 0,\n    max_pulse = 200,\n    min_pulse = 200,\n    *,\n    batch_size: int,\n    shuffle: bool,\n    selection: Optional[List[int]] = None,\n    num_workers: int = 10,\n    persistent_workers: bool = False,\n    node_truth: List[str] = None,\n    truth_table: str = \"truth\",\n    node_truth_table: Optional[str] = None,\n    string_selection: List[int] = None,\n    loss_weight_table: Optional[str] = None,\n    loss_weight_column: Optional[str] = None,\n    index_column: str = \"event_no\",\n    labels: Optional[Dict[str, Callable]] = None,\n) -> DataLoader:\n    \"\"\"Construct `DataLoader` instance.\"\"\"\n    # Check(s)\n    if isinstance(pulsemaps, str):\n        pulsemaps = [pulsemaps]\n\n    dataset = ParquetDataset2(\n        path=db,\n        pulsemaps=pulsemaps,\n        features=features,\n        truth=truth,\n        batch_ids=batch_ids,\n        selection=selection,\n        node_truth=node_truth,\n        truth_table=truth_table,\n        node_truth_table=node_truth_table,\n        string_selection=string_selection,\n        loss_weight_table=loss_weight_table,\n        loss_weight_column=loss_weight_column,\n        index_column=index_column,\n        max_len=max_len,\n        max_pulse=max_pulse,\n        min_pulse=min_pulse,\n    )\n\n    # adds custom labels to dataset\n    if isinstance(labels, dict):\n        for label in labels.keys():\n            dataset.add_label(key=label, fn=labels[label])\n\n    dataloader = DataLoader(\n        dataset,\n        batch_size=batch_size,\n        shuffle=shuffle,\n        num_workers=num_workers,\n        collate_fn=collate_fn,\n        persistent_workers=persistent_workers,\n        prefetch_factor=2,\n    )\n\n    return dataloader, dataset","metadata":{"papermill":{"duration":0.170413,"end_time":"2023-04-21T01:06:52.350348","exception":false,"start_time":"2023-04-21T01:06:52.179935","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-22T13:23:20.428095Z","iopub.execute_input":"2023-04-22T13:23:20.428397Z","iopub.status.idle":"2023-04-22T13:23:20.441449Z","shell.execute_reply.started":"2023-04-22T13:23:20.428372Z","shell.execute_reply":"2023-04-22T13:23:20.440169Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n\nclass StandardModel2(Model):\n    \"\"\"Main class for standard models in graphnet.\n\n    This class chains together the different elements of a complete GNN-based\n    model (detector read-in, GNN architecture, and task-specific read-outs).\n    \"\"\"\n\n    @save_model_config\n    def __init__(\n        self,\n        *,\n        detector: Detector,\n        gnn: GNN,\n        tasks: Union[Task, List[Task]],\n        dataset: Dataset2,\n        max_epochs: 0,\n        coarsening: Optional[Coarsening] = None,\n        optimizer_class: type = Adam,\n        optimizer_kwargs: Optional[Dict] = None,\n        scheduler_class: Optional[type] = None,\n        scheduler_kwargs: Optional[Dict] = None,\n        scheduler_config: Optional[Dict] = None,\n    ) -> None:\n        \"\"\"Construct `StandardModel`.\"\"\"\n        # Base class constructor\n        super().__init__()\n\n        # Check(s)\n        if isinstance(tasks, Task):\n            tasks = [tasks]\n        assert isinstance(tasks, (list, tuple))\n        assert all(isinstance(task, Task) for task in tasks)\n        assert isinstance(detector, Detector)\n        assert isinstance(gnn, GNN)\n        assert coarsening is None or isinstance(coarsening, Coarsening)\n\n        # Member variable(s)\n        self._detector = detector\n        self._gnn = gnn\n        self._tasks = ModuleList(tasks)\n        self._coarsening = coarsening\n        self._optimizer_class = optimizer_class\n        self._optimizer_kwargs = optimizer_kwargs or dict()\n        self._scheduler_class = scheduler_class\n        self._scheduler_kwargs = scheduler_kwargs or dict()\n        self._scheduler_config = scheduler_config or dict()\n        self._dataset = dataset\n        self._max_epochs = max_epochs\n\n    def configure_optimizers(self) -> Dict[str, Any]:\n        \"\"\"Configure the model's optimizer(s).\"\"\"\n        optimizer = self._optimizer_class(\n            self.parameters(), **self._optimizer_kwargs\n        )\n        config = {\n            \"optimizer\": optimizer,\n        }\n        if self._scheduler_class is not None:\n            scheduler = self._scheduler_class(\n                optimizer, **self._scheduler_kwargs\n            )\n            config.update(\n                {\n                    \"lr_scheduler\": {\n                        \"scheduler\": scheduler,\n                        **self._scheduler_config,\n                    },\n                }\n            )\n        return config\n\n    def forward(self, data: Data) -> List[Union[Tensor, Data]]:\n        \"\"\"Forward pass, chaining model components.\"\"\"\n        #import pdb;pdb.set_trace()\n        if self._coarsening:\n            data = self._coarsening(data)\n        data = self._detector(data)\n        x = self._gnn(data)\n        #preds = [task(x) for task in self._tasks]\n        if USE_ALL_FEA_IN_PRED:\n            preds = [task(x) for task in self._tasks] + [x]\n        else:\n            preds = [task(x) for task in self._tasks]\n        #preds = [self._tasks[0](x)]\n        return preds\n\n    def training_step(self, train_batch: Data, batch_idx: int) -> Tensor:\n        \"\"\"Perform training step.\"\"\"\n        preds = self(train_batch)\n        vlosses = self._tasks[1].compute_loss(preds[1], train_batch)\n        vloss = torch.sum(vlosses)\n        \n        tlosses = self._tasks[0].compute_loss(preds[0], train_batch)\n        tloss = torch.sum(tlosses)\n\n        #x = self.current_epoch/self._max_epochs\n        #x = 0.5 + x/2\n        #y = 1-x\n        #loss = vloss*y + tloss*x\n        loss = vloss + tloss\n        return {\"loss\": loss, 'vloss': vloss, 'tloss': tloss}\n\n    def validation_step(self, val_batch: Data, batch_idx: int) -> Tensor:\n        \"\"\"Perform validation step.\"\"\"\n        preds = self(val_batch)\n        vlosses = self._tasks[1].compute_loss(preds[1], val_batch)\n        vloss = torch.sum(vlosses)\n        \n        tlosses = self._tasks[0].compute_loss(preds[0], val_batch)\n        tloss = torch.sum(tlosses)\n        loss = vloss + tloss\n        return {\"loss\": loss, 'vloss': vloss, 'tloss': tloss}\n\n    def _get_batch_size(self, data: Data) -> int:\n        return torch.numel(torch.unique(data.batch))\n\n    def inference(self) -> None:\n        \"\"\"Activate inference mode.\"\"\"\n        for task in self._tasks:\n            task.inference()\n\n    def train(self, mode: bool = True) -> \"Model\":\n        \"\"\"Deactivate inference mode.\"\"\"\n        super().train(mode)\n        if mode:\n            for task in self._tasks:\n                task.train_eval()\n        return self\n\n    def predict(\n        self,\n        dataloader: DataLoader,\n        gpus: Optional[Union[List[int], int]] = None,\n        distribution_strategy: Optional[str] = None,\n    ) -> List[Tensor]:\n        \"\"\"Return predictions for `dataloader`.\"\"\"\n        self.inference()\n        return super().predict(\n            dataloader=dataloader,\n            gpus=gpus,\n            distribution_strategy=distribution_strategy,\n        )\n    \n    def training_epoch_end(self, training_step_outputs):\n        loss = torch.stack([x[\"loss\"] for x in training_step_outputs]).mean()\n        vloss = torch.stack([x[\"vloss\"] for x in training_step_outputs]).mean()\n        tloss = torch.stack([x[\"tloss\"] for x in training_step_outputs]).mean()\n        self.log_dict(\n            {\"trn_loss\": loss, \"trn_vloss\": vloss, \"trn_tloss\": tloss},\n            prog_bar=True,\n            sync_dist=True,\n        )\n        print(f'epoch:{self.current_epoch}, train loss:{loss.item()}, tloss:{tloss.item()}, vloss:{vloss.item()}')\n        self._dataset.reset_epoch()\n        \n    def validation_epoch_end(self, validation_step_outputs):\n        loss = torch.stack([x[\"loss\"] for x in validation_step_outputs]).mean()\n        vloss = torch.stack([x[\"vloss\"] for x in validation_step_outputs]).mean()\n        tloss = torch.stack([x[\"tloss\"] for x in validation_step_outputs]).mean()\n        self.log_dict(\n            {\"val_loss\": loss, \"val_vloss\": vloss, \"val_tloss\": tloss},\n            prog_bar=True,\n            sync_dist=True,\n        )\n        print(f'epoch:{self.current_epoch}, valid loss:{loss.item()}, tloss:{tloss.item()}, vloss:{vloss.item()}')","metadata":{"lines_to_next_cell":2,"papermill":{"duration":0.851662,"end_time":"2023-04-21T01:06:53.215169","exception":false,"start_time":"2023-04-21T01:06:52.363507","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-22T13:23:20.443249Z","iopub.execute_input":"2023-04-22T13:23:20.443653Z","iopub.status.idle":"2023-04-22T13:23:20.470513Z","shell.execute_reply.started":"2023-04-22T13:23:20.443618Z","shell.execute_reply":"2023-04-22T13:23:20.469432Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n\n\nclass EdgeConv0(MessagePassing):\n    r\"\"\"The edge convolutional operator from the `\"Dynamic Graph CNN for\n    Learning on Point Clouds\" <https://arxiv.org/abs/1801.07829>`_ paper\n\n    .. math::\n        \\mathbf{x}^{\\prime}_i = \\sum_{j \\in \\mathcal{N}(i)}\n        h_{\\mathbf{\\Theta}}(\\mathbf{x}_i \\, \\Vert \\,\n        \\mathbf{x}_j - \\mathbf{x}_i),\n\n    where :math:`h_{\\mathbf{\\Theta}}` denotes a neural network, *.i.e.* a MLP.\n\n    Args:\n        nn (torch.nn.Module): A neural network :math:`h_{\\mathbf{\\Theta}}` that\n            maps pair-wise concatenated node features :obj:`x` of shape\n            :obj:`[-1, 2 * in_channels]` to shape :obj:`[-1, out_channels]`,\n            *e.g.*, defined by :class:`torch.nn.Sequential`.\n        aggr (string, optional): The aggregation scheme to use\n            (:obj:`\"add\"`, :obj:`\"mean\"`, :obj:`\"max\"`).\n            (default: :obj:`\"max\"`)\n        **kwargs (optional): Additional arguments of\n            :class:`torch_geometric.nn.conv.MessagePassing`.\n\n    Shapes:\n        - **input:**\n          node features :math:`(|\\mathcal{V}|, F_{in})` or\n          :math:`((|\\mathcal{V}|, F_{in}), (|\\mathcal{V}|, F_{in}))`\n          if bipartite,\n          edge indices :math:`(2, |\\mathcal{E}|)`\n        - **output:** node features :math:`(|\\mathcal{V}|, F_{out})` or\n          :math:`(|\\mathcal{V}_t|, F_{out})` if bipartite\n    \"\"\"\n    def __init__(self, nn: Callable, aggr: str = 'max', **kwargs):\n        super().__init__(aggr=aggr, **kwargs)\n        self.nn = nn\n        self.reset_parameters()\n\n    def reset_parameters(self):\n        reset(self.nn)\n\n    def forward(self, x: Union[Tensor, PairTensor], edge_index: Adj) -> Tensor:\n        \"\"\"\"\"\"\n        if isinstance(x, Tensor):\n            x: PairTensor = (x, x)\n        # propagate_type: (x: PairTensor)\n        return self.propagate(edge_index, x=x, size=None)\n\n    def message(self, x_i: Tensor, x_j: Tensor) -> Tensor:\n        return self.nn(torch.cat([x_i, x_j - x_i, x_j], dim=-1)) ##edgeConv0\n\n    def __repr__(self) -> str:\n        return f'{self.__class__.__name__}(nn={self.nn})'\n    \nNODE_EDGE_FEA_RATIO = 0.7\n\nclass EdgeConv1(MessagePassing):\n    r\"\"\"The edge convolutional operator from the `\"Dynamic Graph CNN for\n    Learning on Point Clouds\" <https://arxiv.org/abs/1801.07829>`_ paper\n\n    .. math::\n        \\mathbf{x}^{\\prime}_i = \\sum_{j \\in \\mathcal{N}(i)}\n        h_{\\mathbf{\\Theta}}(\\mathbf{x}_i \\, \\Vert \\,\n        \\mathbf{x}_j - \\mathbf{x}_i),\n\n    where :math:`h_{\\mathbf{\\Theta}}` denotes a neural network, *.i.e.* a MLP.\n\n    Args:\n        nn (torch.nn.Module): A neural network :math:`h_{\\mathbf{\\Theta}}` that\n            maps pair-wise concatenated node features :obj:`x` of shape\n            :obj:`[-1, 2 * in_channels]` to shape :obj:`[-1, out_channels]`,\n            *e.g.*, defined by :class:`torch.nn.Sequential`.\n        aggr (string, optional): The aggregation scheme to use\n            (:obj:`\"add\"`, :obj:`\"mean\"`, :obj:`\"max\"`).\n            (default: :obj:`\"max\"`)\n        **kwargs (optional): Additional arguments of\n            :class:`torch_geometric.nn.conv.MessagePassing`.\n\n    Shapes:\n        - **input:**\n          node features :math:`(|\\mathcal{V}|, F_{in})` or\n          :math:`((|\\mathcal{V}|, F_{in}), (|\\mathcal{V}|, F_{in}))`\n          if bipartite,\n          edge indices :math:`(2, |\\mathcal{E}|)`\n        - **output:** node features :math:`(|\\mathcal{V}|, F_{out})` or\n          :math:`(|\\mathcal{V}_t|, F_{out})` if bipartite\n    \"\"\"\n    def __init__(self, nn: Callable, aggr: str = 'max', **kwargs):\n        super().__init__(aggr=aggr, **kwargs)\n        self.nn = nn\n        self.reset_parameters()\n\n    def reset_parameters(self):\n        reset(self.nn)\n\n    def forward(self, x: Union[Tensor, PairTensor], edge_index: Adj) -> Tensor:\n        \"\"\"\"\"\"\n        if isinstance(x, Tensor):\n            x: PairTensor = (x, x)\n        # propagate_type: (x: PairTensor)\n        return self.propagate(edge_index, x=x, size=None)\n\n    def message(self, x_i: Tensor, x_j: Tensor) -> Tensor:\n        edge_cnt = round(x_i.shape[-1]*NODE_EDGE_FEA_RATIO)\n        if edge_cnt < 20:\n            edge_cnt = 4                   #TODO　通常edge素性の数は20以上、それより少ないときは最初のレイヤーなのでXYZTの4つ\n        edge_ij = torch.cat([(x_j - x_i)[:,:edge_cnt],x_j[:,edge_cnt:]], axis=-1)\n        return self.nn(torch.cat([x_i, edge_ij], dim=-1)) ##edgeConv1\n\n    def __repr__(self) -> str:\n        return f'{self.__class__.__name__}(nn={self.nn})'","metadata":{"papermill":{"duration":0.032024,"end_time":"2023-04-21T01:06:53.259996","exception":false,"start_time":"2023-04-21T01:06:53.227972","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-22T13:23:20.472292Z","iopub.execute_input":"2023-04-22T13:23:20.472724Z","iopub.status.idle":"2023-04-22T13:23:20.488667Z","shell.execute_reply.started":"2023-04-22T13:23:20.472685Z","shell.execute_reply":"2023-04-22T13:23:20.487681Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"USE_TRANS_IN_DYN1=True\n\nclass dynTrans1(EdgeConv0, LightningModule):\n    \"\"\"Dynamical edge convolution layer.\"\"\"\n\n    def __init__(\n        self,\n        layer_sizes,\n        aggr: str = \"max\",\n        nb_neighbors: int = 8,\n        features_subset: Optional[Union[Sequence[int], slice]] = None,\n        **kwargs: Any,\n    ):\n        \"\"\"Construct `DynEdgeConv`.\n\n        Args:\n            nn: The MLP/torch.Module to be used within the `EdgeConv`.\n            aggr: Aggregation method to be used with `EdgeConv`.\n            nb_neighbors: Number of neighbours to be clustered after the\n                `EdgeConv` operation.\n            features_subset: Subset of features in `Data.x` that should be used\n                when dynamically performing the new graph clustering after the\n                `EdgeConv` operation. Defaults to all features.\n            **kwargs: Additional features to be passed to `EdgeConv`.\n        \"\"\"\n        # Check(s)\n        if features_subset is None:\n            features_subset = slice(None)  # Use all features\n        assert isinstance(features_subset, (list, slice))\n                \n        layers = []\n        for ix, (nb_in, nb_out) in enumerate(\n            zip(layer_sizes[:-1], layer_sizes[1:])\n        ):\n            if ix == 0:\n                nb_in *= 3 # edgeConv1\n            layers.append(torch.nn.Linear(nb_in, nb_out))\n            layers.append(torch.nn.LeakyReLU())\n        d_model = nb_out\n        # Base class constructor\n        super().__init__(nn=torch.nn.Sequential(*layers), aggr=aggr, **kwargs)\n\n        # Additional member variables\n        self.nb_neighbors = nb_neighbors\n        self.features_subset = features_subset\n        \n\n        self.norm_first=False\n        \n        self.norm1 = LayerNorm(d_model, eps=1e-5) #lNorm\n        \n        # Transformer layer(s)\n        if USE_TRANS_IN_DYN1:\n            encoder_layer = TransformerEncoderLayer(d_model=d_model, nhead=8, batch_first=True, dropout=DROPOUT, norm_first=self.norm_first)\n            self._transformer_encoder = TransformerEncoder(encoder_layer, num_layers=1)        \n        \n        \n\n    def forward(\n        self, x: Tensor, edge_index: Adj, batch: Optional[Tensor] = None\n    ) -> Tensor:\n        \"\"\"Forward pass.\"\"\"\n        # Standard EdgeConv forward pass\n        \n        if self.norm_first:\n            x = self.norm1(x) # lNorm\n            \n        x_out = super().forward(x, edge_index)\n        \n        if x_out.shape[-1] == x.shape[-1] and SERIAL_CONNECTION:\n            x = x + x_out\n        else:\n            x = x_out\n            \n        if not self.norm_first:\n            x = self.norm1(x) # lNorm\n\n        # Recompute adjacency\n        edge_index = None\n\n        # Transformer layer\n        if USE_TRANS_IN_DYN1:\n            x, mask = to_dense_batch(x, batch)\n            x = self._transformer_encoder(x, src_key_padding_mask=~mask)\n            x = x[mask]\n\n        return x, edge_index","metadata":{"lines_to_next_cell":0,"papermill":{"duration":0.030686,"end_time":"2023-04-21T01:06:53.302889","exception":false,"start_time":"2023-04-21T01:06:53.272203","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-22T13:23:20.490240Z","iopub.execute_input":"2023-04-22T13:23:20.490783Z","iopub.status.idle":"2023-04-22T13:23:20.505273Z","shell.execute_reply.started":"2023-04-22T13:23:20.490744Z","shell.execute_reply":"2023-04-22T13:23:20.504255Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"GLOBAL_POOLINGS = {\n    \"min\": scatter_min,\n    \"max\": scatter_max,\n    \"sum\": scatter_sum,\n    \"mean\": scatter_mean,\n}\n\nclass DynEdge(GNN):\n    \"\"\"DynEdge (dynamical edge convolutional) model.\"\"\"\n\n    @save_model_config\n    def __init__(\n        self,\n        nb_inputs: int,\n        *,\n        nb_neighbours: int = 8,\n        features_subset: Optional[Union[List[int], slice]] = None,\n        dynedge_layer_sizes: Optional[List[Tuple[int, ...]]] = None,\n        post_processing_layer_sizes: Optional[List[int]] = None,\n        readout_layer_sizes: Optional[List[int]] = None,\n        global_pooling_schemes: Optional[Union[str, List[str]]] = None,\n        add_global_variables_after_pooling: bool = False,\n    ):\n        \"\"\"Construct `DynEdge`.\n\n        Args:\n            nb_inputs: Number of input features on each node.\n            nb_neighbours: Number of neighbours to used in the k-nearest\n                neighbour clustering which is performed after each (dynamical)\n                edge convolution.\n            features_subset: The subset of latent features on each node that\n                are used as metric dimensions when performing the k-nearest\n                neighbours clustering. Defaults to [0,1,2].\n            dynedge_layer_sizes: The layer sizes, or latent feature dimenions,\n                used in the `DynEdgeConv` layer. Each entry in\n                `dynedge_layer_sizes` corresponds to a single `DynEdgeConv`\n                layer; the integers in the corresponding tuple corresponds to\n                the layer sizes in the multi-layer perceptron (MLP) that is\n                applied within each `DynEdgeConv` layer. That is, a list of\n                size-two tuples means that all `DynEdgeConv` layers contain a\n                two-layer MLP.\n                Defaults to [(128, 256), (336, 256), (336, 256), (336, 256)].\n            post_processing_layer_sizes: Hidden layer sizes in the MLP\n                following the skip-concatenation of the outputs of each\n                `DynEdgeConv` layer. Defaults to [336, 256].\n            readout_layer_sizes: Hidden layer sizes in the MLP following the\n                post-processing _and_ optional global pooling. As this is the\n                last layer(s) in the model, the last layer in the read-out\n                yields the output of the `DynEdge` model. Defaults to [128,].\n            global_pooling_schemes: The list global pooling schemes to use.\n                Options are: \"min\", \"max\", \"mean\", and \"sum\".\n            add_global_variables_after_pooling: Whether to add global variables\n                after global pooling. The alternative is to  added (distribute)\n                them to the individual nodes before any convolutional\n                operations.\n        \"\"\"\n        # Latent feature subset for computing nearest neighbours in DynEdge.\n        if features_subset is None:\n            features_subset = slice(0, 4) #4D\n\n        # DynEdge layer sizes\n        if dynedge_layer_sizes is None: #nb_nearest_neighboursと合わせて変更\n            dynedge_layer_sizes = DYNEDGE_LAYER_SIZE\n\n        assert isinstance(dynedge_layer_sizes, list)\n        assert len(dynedge_layer_sizes)\n        assert all(isinstance(sizes, tuple) for sizes in dynedge_layer_sizes)\n        assert all(len(sizes) > 0 for sizes in dynedge_layer_sizes)\n        assert all(\n            all(size > 0 for size in sizes) for sizes in dynedge_layer_sizes\n        )\n\n        self._dynedge_layer_sizes = dynedge_layer_sizes\n\n        # Post-processing layer sizes\n        if post_processing_layer_sizes is None:\n            post_processing_layer_sizes = [\n                336,\n                256,\n            ]\n\n        assert isinstance(post_processing_layer_sizes, list)\n        assert len(post_processing_layer_sizes)\n        assert all(size > 0 for size in post_processing_layer_sizes)\n\n        self._post_processing_layer_sizes = post_processing_layer_sizes\n\n        # Read-out layer sizes\n        if readout_layer_sizes is None:\n            readout_layer_sizes = [\n                256,\n                128,\n            ]\n\n        assert isinstance(readout_layer_sizes, list)\n        assert len(readout_layer_sizes)\n        assert all(size > 0 for size in readout_layer_sizes)\n\n        self._readout_layer_sizes = readout_layer_sizes\n        \n\n\n        # Global pooling scheme(s)\n        if isinstance(global_pooling_schemes, str):\n            global_pooling_schemes = [global_pooling_schemes]\n\n        if isinstance(global_pooling_schemes, list):\n            for pooling_scheme in global_pooling_schemes:\n                assert (\n                    pooling_scheme in GLOBAL_POOLINGS\n                ), f\"Global pooling scheme {pooling_scheme} not supported.\"\n        else:\n            assert global_pooling_schemes is None\n\n        self._global_pooling_schemes = global_pooling_schemes\n\n        if add_global_variables_after_pooling:\n            assert self._global_pooling_schemes, (\n                \"No global pooling schemes were request, so cannot add global\"\n                \" variables after pooling.\"\n            )\n        self._add_global_variables_after_pooling = (\n            add_global_variables_after_pooling\n        )\n\n        # Base class constructor\n        super().__init__(nb_inputs, self._readout_layer_sizes[-1])\n\n        # Remaining member variables()\n        self._activation = torch.nn.LeakyReLU()\n        self._nb_inputs = nb_inputs\n        self._nb_global_variables = 5 + nb_inputs\n        self._nb_neighbours = nb_neighbours\n        self._features_subset = features_subset\n\n        self._construct_layers()\n\n    def _construct_layers(self) -> None:\n        \"\"\"Construct layers (torch.nn.Modules).\"\"\"\n        # Convolutional operations\n        nb_input_features = self._nb_inputs\n        if USE_G:\n            if not self._add_global_variables_after_pooling:\n                nb_input_features += self._nb_global_variables\n\n        self._conv_layers = torch.nn.ModuleList()\n        nb_latent_features = nb_input_features\n        for sizes in self._dynedge_layer_sizes:\n            conv_layer = dynTrans1(\n                [nb_latent_features] + list(sizes),\n                aggr=\"max\",\n                nb_neighbors=self._nb_neighbours,\n                features_subset=self._features_subset,\n            )\n            self._conv_layers.append(conv_layer)\n            nb_latent_features = sizes[-1]\n\n        # Post-processing operations\n        if SERIAL_CONNECTION:\n            nb_latent_features = self._dynedge_layer_sizes[-1][-1]\n        else:\n            nb_latent_features = (\n                sum(sizes[-1] for sizes in self._dynedge_layer_sizes)\n                + nb_input_features\n            )\n\n        if USE_PP:\n            post_processing_layers = []\n            layer_sizes = [nb_latent_features] + list(\n                self._post_processing_layer_sizes\n            )\n            for nb_in, nb_out in zip(layer_sizes[:-1], layer_sizes[1:]):\n                post_processing_layers.append(torch.nn.Linear(nb_in, nb_out))\n                post_processing_layers.append(self._activation)\n            last_posting_layer_output_dim = nb_out\n\n            self._post_processing = torch.nn.Sequential(*post_processing_layers)\n        else:\n            last_posting_layer_output_dim = nb_latent_features\n\n        # Read-out operations\n        nb_poolings = (\n            len(self._global_pooling_schemes)\n            if self._global_pooling_schemes\n            else 1\n        )\n        nb_latent_features = last_posting_layer_output_dim * nb_poolings\n        if USE_G:\n            if self._add_global_variables_after_pooling:  \n                nb_latent_features += self._nb_global_variables\n\n        readout_layers = []\n        layer_sizes = [nb_latent_features] + list(self._readout_layer_sizes)\n        for nb_in, nb_out in zip(layer_sizes[:-1], layer_sizes[1:]):\n            readout_layers.append(torch.nn.Linear(nb_in, nb_out))\n            readout_layers.append(self._activation)\n\n        self._readout = torch.nn.Sequential(*readout_layers)\n        \n\n        # Transformer layer(s)\n        if USE_TRANS_IN_LAST:\n            encoder_layer = TransformerEncoderLayer(d_model=last_posting_layer_output_dim, nhead=8, batch_first=True, dropout=DROPOUT, norm_first=False)\n            self._transformer_encoder = TransformerEncoder(encoder_layer, num_layers=USE_TRANS_IN_LAST)        \n\n\n    def _global_pooling(self, x: Tensor, batch: LongTensor) -> Tensor:\n        \"\"\"Perform global pooling.\"\"\"\n        assert self._global_pooling_schemes\n        pooled = []\n        for pooling_scheme in self._global_pooling_schemes:\n            pooling_fn = GLOBAL_POOLINGS[pooling_scheme]\n            pooled_x = pooling_fn(x, index=batch, dim=0)\n            if isinstance(pooled_x, tuple) and len(pooled_x) == 2:\n                # `scatter_{min,max}`, which return also an argument, vs.\n                # `scatter_{mean,sum}`\n                pooled_x, _ = pooled_x\n            pooled.append(pooled_x)\n\n        return torch.cat(pooled, dim=1)\n\n    def _calculate_global_variables(\n        self,\n        x: Tensor,\n        edge_index: LongTensor,\n        batch: LongTensor,\n        *additional_attributes: Tensor,\n    ) -> Tensor:\n        \"\"\"Calculate global variables.\"\"\"\n        # Calculate homophily (scalar variables)\n        h_x, h_y, h_z, h_t = calculate_xyzt_homophily(x, edge_index, batch)\n\n        # Calculate mean features\n        global_means = scatter_mean(x, batch, dim=0)\n\n        # Add global variables\n        global_variables = torch.cat(\n            [\n                global_means,\n                h_x,\n                h_y,\n                h_z,\n                h_t,\n            ]\n            + [attr.unsqueeze(dim=1) for attr in additional_attributes],\n            dim=1,\n        )\n\n        return global_variables\n\n    def forward(self, data: Data) -> Tensor:\n        \"\"\"Apply learnable forward pass.\"\"\"\n        # Convenience variables\n        x, edge_index, batch = data.x, data.edge_index, data.batch\n\n        \n        if USE_G:\n            global_variables = self._calculate_global_variables(\n                x,\n                edge_index[0],\n                batch,\n                torch.log10(data.n_pulses),\n            )\n\n            # Distribute global variables out to each node\n            if not self._add_global_variables_after_pooling:\n                distribute = (\n                    batch.unsqueeze(dim=1) == torch.unique(batch).unsqueeze(dim=0)\n                ).type(torch.float)\n\n                global_variables_distributed = torch.sum(\n                    distribute.unsqueeze(dim=2)\n                    * global_variables.unsqueeze(dim=0),\n                    dim=1,\n                )\n\n                x = torch.cat((x, global_variables_distributed), dim=1)\n\n\n        if SERIAL_CONNECTION:\n            for conv_layer_index, conv_layer in enumerate(self._conv_layers):\n                x, _edge_index = conv_layer(x, data.edge_index[0], batch)\n        else:\n            skip_connections = [x]\n            for conv_layer_index, conv_layer in enumerate(self._conv_layers):\n                x, _edge_index = conv_layer(x, data.edge_index[0], batch)\n                #print('    dynEdge output skip_connections', x.shape)\n                skip_connections.append(x)\n\n            # Skip-cat\n            x = torch.cat(skip_connections, dim=1)\n\n        if USE_TRANS_IN_LAST:\n            x, mask = to_dense_batch(x, batch)\n            x = self._transformer_encoder(x, src_key_padding_mask=~mask)\n            x = x[mask]\n        \n        if USE_PP:\n            x = self._post_processing(x)\n        \n        if self._global_pooling_schemes:\n            x = self._global_pooling(x, batch=batch)\n            if USE_G:\n                if self._add_global_variables_after_pooling:\n                    x = torch.cat(\n                        [\n                            x,\n                            global_variables,\n                        ],\n                        dim=1,\n                    )\n\n        x = self._readout(x)\n\n        return x","metadata":{"lines_to_next_cell":0,"papermill":{"duration":0.054201,"end_time":"2023-04-21T01:06:53.369198","exception":false,"start_time":"2023-04-21T01:06:53.314997","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-22T13:23:20.508193Z","iopub.execute_input":"2023-04-22T13:23:20.509125Z","iopub.status.idle":"2023-04-22T13:23:20.545541Z","shell.execute_reply.started":"2023-04-22T13:23:20.509082Z","shell.execute_reply":"2023-04-22T13:23:20.544578Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nclass GraphBuilder(Model):  # pylint: disable=too-few-public-methods\n    \"\"\"Base class for graph building.\"\"\"\n\n    pass\n\n\n\nclass KNNGraphBuilderMulti(GraphBuilder):  # pylint: disable=too-few-public-methods\n    \"\"\"Builds graph from the k-nearest neighbours.\"\"\"\n\n    @save_model_config\n    def __init__(\n        self,\n        nb_nearest_neighbours,\n        columns,\n    ):\n        \"\"\"Construct `KNNGraphBuilder`.\"\"\"\n        # Base class constructor\n        super().__init__()\n\n        # Member variable(s)\n        assert len(nb_nearest_neighbours) == len(columns)\n        self._nb_nearest_neighbours = nb_nearest_neighbours\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        \n        if data.edge_index is not None:\n            self.info(\n                \"WARNING: GraphBuilder received graph with pre-existing \"\n                \"structure. Will overwrite.\"\n            )\n        edge_index_list = []\n        x = data.x\n        TIME_PARAM_FOR_DIST = 1/10\n        x[:,3] = x[:,3]*TIME_PARAM_FOR_DIST\n        for idx in range(len(self._nb_nearest_neighbours)):\n            nb_nearest_neighbour = self._nb_nearest_neighbours[idx]\n            if type(nb_nearest_neighbour) == str:\n                nb_nearest_neighbour_min, nb_nearest_neighbour_max = nb_nearest_neighbour.split('-')\n                nb_nearest_neighbour = torch.randint(int(nb_nearest_neighbour_min), int(nb_nearest_neighbour_max), (1,))[0]\n            elif type(nb_nearest_neighbour) == list:\n                nb_nearest_neighbour = random.choice(nb_nearest_neighbour)\n            edge_index = knn_graph(\n                x[:, self._columns[idx]]/1000,\n                nb_nearest_neighbour,\n                data.batch,\n            ).to(self.device)\n            edge_index_list.append(edge_index)\n        x[:,3] = x[:,3]/TIME_PARAM_FOR_DIST\n\n        data.edge_index = edge_index_list\n        return data","metadata":{"papermill":{"duration":0.027006,"end_time":"2023-04-21T01:06:53.408671","exception":false,"start_time":"2023-04-21T01:06:53.381665","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-22T13:23:20.549947Z","iopub.execute_input":"2023-04-22T13:23:20.551161Z","iopub.status.idle":"2023-04-22T13:23:20.561944Z","shell.execute_reply.started":"2023-04-22T13:23:20.551122Z","shell.execute_reply":"2023-04-22T13:23:20.560950Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nlogger = get_logger()\n\nclass DistanceLoss2(LossFunction):\n\n    def _forward(self, prediction: Tensor, target: Tensor) -> Tensor:\n        target = target.reshape(-1, 3)\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        eps = 1e-4\n        prediction_length = torch.linalg.vector_norm(prediction[:, [0, 1, 2]], dim=1)\n        prediction_length = torch.clamp(prediction_length, min=eps)\n        prediction =  prediction[:, [0, 1, 2]]/prediction_length.unsqueeze(1)\n        cosLoss = prediction[:, 0] * target[:, 0] + prediction[:, 1] * target[:, 1] + prediction[:, 2] * target[:, 2]    \n        cosLoss = torch.clamp(cosLoss, min=-1+eps, max=1-eps)\n        thetaLoss = torch.arccos(cosLoss)\n        return thetaLoss\n    \n    \n\nclass IceCubeKaggle2(Detector):\n    \"\"\"`Detector` class for Kaggle Competition.\"\"\"\n\n    # Implementing abstract class attribute\n    features = FEATURES.KAGGLE\n\n    def _forward(self, data: Data) -> Data:\n        \"\"\"Ingest data, build graph, and preprocess features.\n\n        Args:\n            data: Input graph data.\n\n        Returns:\n            Connected and preprocessed graph data.\n        \"\"\"\n        # Check(s)\n        self._validate_features(data)\n\n        # Preprocessing\n        data.x[:, 0] /= 500.0  # x\n        data.x[:, 1] /= 500.0  # y\n        data.x[:, 2] /= 500.0  # z\n        data.x[:, 3] = (data.x[:, 3] - 1.0e04) / (500.0*0.23)  # 0.23 is speed of light in ice\n        data.x[:, 4] = torch.log10(data.x[:, 4]) / 3.0  # charge\n\n        return data\n    \nclass DirectionReconstructionWithKappa2(Task):\n    \"\"\"Reconstructs direction with kappa from the 3D-vMF distribution.\"\"\"\n\n    # Requires three features: untransformed points in (x,y,z)-space.\n    nb_inputs = 3\n\n    def _forward(self, x: Tensor) -> Tensor:\n        # Transform outputs to angle and prepare prediction\n        kappa = torch.linalg.vector_norm(x, dim=1)# + eps_like(x)\n        kappa = torch.clamp(kappa, min=torch.finfo(x.dtype).eps)\n        vec_x = x[:, 0] / kappa\n        vec_y = x[:, 1] / kappa\n        vec_z = x[:, 2] / kappa\n        return torch.stack((vec_x, vec_y, vec_z, kappa), dim=1)\n\ndef build_model2(config: Dict[str,Any], train_dataloader: Any, train_dataset: Any) -> StandardModel2:\n    \"\"\"Builds GNN from config\"\"\"\n    # Building model\n    detector = IceCubeKaggle2(\n        graph_builder=KNNGraphBuilderMulti(nb_nearest_neighbours=NB_NEAREST_NEIGHBOURS, columns=COLUMNS_NEAREST_NEIGHBOURS) #dynedge_layer_sizes と合わせて変更\n    )\n    gnn = DynEdge(\n        nb_inputs=detector.nb_outputs,\n        #global_pooling_schemes=[\"min\", \"max\", \"mean\"],\n        global_pooling_schemes=[\"max\"],\n        add_global_variables_after_pooling=True\n    )\n\n    if config[\"target\"] == 'direction':\n        task = DirectionReconstructionWithKappa2(\n            hidden_size=gnn.nb_outputs,\n            target_labels=config[\"target\"],\n            loss_function=VonMisesFisher3DLoss(),\n        )\n        task2 = DirectionReconstructionWithKappa2(\n            hidden_size=gnn.nb_outputs,\n            target_labels=config[\"target\"],\n            loss_function=DistanceLoss2(),\n        )\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 = StandardModel2(\n        detector=detector,\n        gnn=gnn,\n        tasks=[task2, task],\n        dataset=train_dataset,\n        max_epochs=config[\"fit\"][\"max_epochs\"],\n        optimizer_class=Adam,\n        optimizer_kwargs={\"lr\": 1e-03, \"eps\": 1e-03},\n        #optimizer_class=Lion,\n        #optimizer_kwargs={\"lr\": 1e-04},\n        scheduler_class=PiecewiseLinearLR,\n        scheduler_kwargs={\n            \"milestones\": [\n                0,\n                10  * len(train_dataloader)//(len(config['gpus'])*config['accumulate_grad_batches'][0]),\n                len(train_dataloader)*config[\"fit\"][\"max_epochs\"]//(len(config['gpus'])*config['accumulate_grad_batches'][0]*2),\n                len(train_dataloader)*config[\"fit\"][\"max_epochs\"]//(len(config['gpus'])*config['accumulate_grad_batches'][0]),                \n            ],\n            \"factors\": [1e-03, 1, 1, 1e-03],\n            \"verbose\": config[\"scheduler_verbose\"],\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\n\ndef load_pretrained_model(config: Dict[str,Any], state_dict_path: str = '/home/tito/kaggle/icecube-neutrinos-in-deep-ice/input/dynedge-pretrained/dynedge_pretrained_batch_1_to_50/state_dict.pth') -> StandardModel2:\n    train_dataloader, _, train_dataset, _ = make_dataloaders2(config = config)\n    model = build_model2(config = config, \n                        train_dataloader = train_dataloader,\n                        train_dataset = train_dataset,\n                        )\n    #model._inference_trainer = Trainer(config['fit'])\n    model.load_state_dict(state_dict_path)\n    model.prediction_columns = [config[\"target\"] + \"_x\", \n                              config[\"target\"] + \"_y\", \n                              config[\"target\"] + \"_z\", \n                              config[\"target\"] + \"_kappa\" ]\n    model.additional_attributes = ['zenith', 'azimuth', 'event_id']\n    return model\n\ndef make_dataloaders2(config: Dict[str, Any]) -> List[Any]:\n    \"\"\"Constructs training and validation dataloaders for training with early stopping.\"\"\"\n    train_dataloader, train_dataset = make_dataloader2(db = \"dummy\",\n                                            selection = None,\n                                            pulsemaps = config['pulsemap'],\n                                            features = features,\n                                            truth = truth,\n                                            batch_ids = config['train_batch_ids'],\n                                            batch_size = config['batch_size'],\n                                            num_workers = config['num_workers'],\n                                            shuffle = True,\n                                            labels = {'direction': Direction()},\n                                            index_column = config['index_column'],\n                                            truth_table = config['truth_table'],\n                                            max_len = config['train_len'],\n                                            max_pulse = config['train_max_pulse'],\n                                            min_pulse = config['train_min_pulse'],\n                                            )\n    \n    validate_dataloader, validate_dataset = make_dataloader2(db = \"dummy\",\n                                            selection = None,\n                                            pulsemaps = config['pulsemap'],\n                                            features = features,\n                                            truth = truth,\n                                            batch_ids = config['valid_batch_ids'],\n                                            batch_size = config['batch_size'],\n                                            num_workers = config['num_workers'],\n                                            shuffle = False,\n                                            labels = {'direction': Direction()},\n                                            index_column = config['index_column'],\n                                            truth_table = config['truth_table'],\n                                            max_len = config['valid_len'],\n                                            max_pulse = config['valid_max_pulse'],\n                                            min_pulse = config['valid_min_pulse'],\n                                          \n                                            )\n    return train_dataloader, validate_dataloader,  train_dataset, validate_dataset\n\n\ndef inference(model, config: Dict[str, Any]) -> pd.DataFrame:\n    \"\"\"Applies model to the database specified in config['inference_database_path'] and saves results to disk.\"\"\"\n    # Make Dataloader\n    test_dataloader = make_dataloader(db = config['inference_database_path'],\n                                            selection = None, # Entire database\n                                            pulsemaps = config['pulsemap'],\n                                            features = features,\n                                            truth = truth,\n                                            batch_size = config['batch_size'],\n                                            num_workers = config['num_workers'],\n                                            shuffle = False,\n                                            labels = {'direction': Direction()},\n                                            index_column = config['index_column'],\n                                            truth_table = config['truth_table'],\n                                            )\n    \n    # Get predictions\n    results = model.predict_as_dataframe(\n        gpus = config['gpus'],\n        dataloader = test_dataloader,\n        prediction_columns=model.prediction_columns,\n        additional_attributes=model.additional_attributes,\n    )\n    return results","metadata":{"lines_to_next_cell":0,"papermill":{"duration":0.061019,"end_time":"2023-04-21T01:06:53.481729","exception":false,"start_time":"2023-04-21T01:06:53.420710","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-22T13:23:20.563473Z","iopub.execute_input":"2023-04-22T13:23:20.564113Z","iopub.status.idle":"2023-04-22T13:23:20.596391Z","shell.execute_reply.started":"2023-04-22T13:23:20.564074Z","shell.execute_reply":"2023-04-22T13:23:20.595329Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def is_env_notebook():\n    \"\"\"Determine wheather is the environment Jupyter Notebook\"\"\"\n    if 'get_ipython' not in globals():\n        # Python shell\n        return False\n    env_name = get_ipython().__class__.__name__\n    if env_name == 'TerminalInteractiveShell':\n        # IPython shell\n        return False\n    # Jupyter Notebook\n    return True","metadata":{"papermill":{"duration":0.01995,"end_time":"2023-04-21T01:06:53.514039","exception":false,"start_time":"2023-04-21T01:06:53.494089","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-22T13:23:20.598152Z","iopub.execute_input":"2023-04-22T13:23:20.598585Z","iopub.status.idle":"2023-04-22T13:23:20.606399Z","shell.execute_reply.started":"2023-04-22T13:23:20.598537Z","shell.execute_reply":"2023-04-22T13:23:20.605469Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training\nThis notebook dose not support training part.","metadata":{}},{"cell_type":"code","source":"BATCH_DIR = '../input/icecube-neutrinos-in-deep-ice/train'\nMETA_DIR = '../input/icecube-neutrinos-in-deep-ice/train'\nFILTER_BY_KAPPA_THRE = 0.5","metadata":{"lines_to_next_cell":2,"papermill":{"duration":0.020154,"end_time":"2023-04-21T01:06:53.546137","exception":false,"start_time":"2023-04-21T01:06:53.525983","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-22T13:23:20.608638Z","iopub.execute_input":"2023-04-22T13:23:20.609421Z","iopub.status.idle":"2023-04-22T13:23:20.614902Z","shell.execute_reply.started":"2023-04-22T13:23:20.609386Z","shell.execute_reply":"2023-04-22T13:23:20.614051Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# training setting\n\nTRAIN_MODE = False\nhostName = socket.gethostname()\nDROPOUT=0.0\nNB_NEAREST_NEIGHBOURS = [6]\nCOLUMNS_NEAREST_NEIGHBOURS = [slice(0,4)]\nUSE_G = True\nONLY_AUX_FALSE = False\nSERIAL_CONNECTION = True\nUSE_PP = True\nUSE_TRANS_IN_LAST=0\nDYNEDGE_LAYER_SIZE = [\n                (\n                    256,\n                    256,\n                ),\n                (\n                    256,\n                    256,\n                ),\n                (\n                    256,\n                    256,\n                ),\n            ]\n\n# Constants\nfeatures = FEATURES.KAGGLE\ntruth = TRUTH.KAGGLE\n\n\nrunName = 'base1-3l250p4n-batch400-650x30-infer' #TODO\n\n    \n# Configuration\nFORCE_MAX_PULSE = 6000 # 強制的にこの件数以上のパルスは返さない\n\nproject = \"650x20\"\nn_batch = 650\nn_round = 10\n\nbatch_size = 400                # multi GPU\ngpus = [0]                     # multi GPU\naccumulate_grad_batches = {0: 5}\n\nif len(gpus) > 1:\n    if is_env_notebook():\n        distribution_strategy = 'ddp_notebook'\n    else:\n        distribution_strategy = 'ddp'\nelse:\n    distribution_strategy = None\n\n\n\nconfig = {\n        \"path\": 'dummy',\n        \"inference_database_path\": '/home/tito/kaggle/icecube-neutrinos-in-deep-ice/notebook_graphnet/batch_660.db', #dummy\n        \"pulsemap\": 'pulse_table', #dummy\n        \"truth_table\": 'meta_table', #dummy\n        \"features\": features,\n        \"truth\": truth,\n        \"index_column\": 'event_id',\n        \"run_name_tag\": 'my_example',\n        \"batch_size\": batch_size,\n        \"num_workers\": 2, #todo\n        \"target\": 'direction',\n        \"early_stopping_patience\": n_batch,\n        \"gpus\": gpus,\n        \"fit\": {\n                \"max_epochs\": n_batch*n_round,\n                \"gpus\": gpus,\n                \"distribution_strategy\": distribution_strategy,\n                \"check_val_every_n_epoch\":10,\n                \"precision\": 16,\n                #\"gradient_clip_val\": 0.9,\n                \"reload_dataloaders_every_n_epochs\": 1,\n                },\n\n        \"accumulate_grad_batches\": accumulate_grad_batches,\n        'runName': runName,\n        'project': project,\n        'scheduler_verbose': False,\n        'train_batch_ids': list(range(1,n_batch+1)),\n        'valid_batch_ids': [660], # only suport one batch\n        'test_selection': None,\n        'base_dir': 'training',\n        'train_len': 0,                    #not using anymore\n        'valid_len': 0,                    #not using anymore\n        'train_max_pulse': 300,\n        'valid_max_pulse': 200,\n        'train_min_pulse': 0,\n        'valid_min_pulse': 0,\n}\n\n\ndebug = False # bbb\nif debug:\n    runName = runName + '_debug'\n    config[\"project\"] = 'debug'\n    #config[\"num_workers\"] = 0\n    config[\"batch_size\"] = 2\n    config[\"train_len\"] = 2\n    config[\"valid_len\"] = 2","metadata":{"papermill":{"duration":0.028971,"end_time":"2023-04-21T01:06:53.587377","exception":false,"start_time":"2023-04-21T01:06:53.558406","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-22T13:23:20.620613Z","iopub.execute_input":"2023-04-22T13:23:20.621303Z","iopub.status.idle":"2023-04-22T13:23:20.634061Z","shell.execute_reply.started":"2023-04-22T13:23:20.621263Z","shell.execute_reply":"2023-04-22T13:23:20.632983Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if TRAIN_MODE:\n\n    train_dataloader, validate_dataloader, train_dataset, validate_dataset = make_dataloaders2(config = config)\n    model = build_model2(config, train_dataloader, train_dataset)\n\n\n    # Training model\n    callbacks = [\n        ModelCheckpoint(\n            dirpath='../model_checkpoint_graphnet/',\n            filename=runName+'-{epoch:02d}-{val_tloss:.6f}',\n            monitor= 'val_tloss',\n            save_top_k = 30,\n            every_n_epochs = 10,\n            save_weights_only=False,\n        ),\n        ProgressBar(),\n    ]\n\n    if 'accumulate_grad_batches' in config and len(config['accumulate_grad_batches']) > 0:\n        callbacks.append(GradientAccumulationScheduler(scheduling=config['accumulate_grad_batches']))\n\n    if debug == False:\n        config[\"fit\"][\"logger\"] = pl_loggers.WandbLogger(project=config[\"project\"], name=runName)\n        callbacks.append(LearningRateMonitor(logging_interval='epoch'))\n\n    #config[\"fit\"][\"profiler\"] = PyTorchProfiler( output_filename='profiler_results.txt', trace_every_n_steps=1)\n\n\n    model.fit(\n        train_dataloader,\n        validate_dataloader,\n        callbacks=callbacks,\n        **config[\"fit\"],\n    )\n    model.save_state_dict(f'../model_graphnet/{runName}-last.pth')","metadata":{"execution":{"iopub.status.busy":"2023-04-22T13:23:20.636617Z","iopub.execute_input":"2023-04-22T13:23:20.637367Z","iopub.status.idle":"2023-04-22T13:23:20.646227Z","shell.execute_reply.started":"2023-04-22T13:23:20.637329Z","shell.execute_reply":"2023-04-22T13:23:20.645289Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Inference & Evaluation","metadata":{"papermill":{"duration":0.011855,"end_time":"2023-04-21T01:06:53.611597","exception":false,"start_time":"2023-04-21T01:06:53.599742","status":"completed"},"tags":[]}},{"cell_type":"code","source":"validateMode = False\nUSE_ALL_FEA_IN_PRED=True\nINF_DEVICES = 0\nresult1_files = []","metadata":{"papermill":{"duration":0.019866,"end_time":"2023-04-21T01:06:53.643528","exception":false,"start_time":"2023-04-21T01:06:53.623662","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-22T13:23:20.647758Z","iopub.execute_input":"2023-04-22T13:23:20.648138Z","iopub.status.idle":"2023-04-22T13:23:20.657501Z","shell.execute_reply.started":"2023-04-22T13:23:20.648102Z","shell.execute_reply":"2023-04-22T13:23:20.656771Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def is_env_notebook():\n    \"\"\"Determine wheather is the environment Jupyter Notebook\"\"\"\n    if 'get_ipython' not in globals():\n        # Python shell\n        return False\n    env_name = get_ipython().__class__.__name__\n    if env_name == 'TerminalInteractiveShell':\n        # IPython shell\n        return False\n    # Jupyter Notebook\n    return True\n\n\ndef 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\n\ndef infer(min_pulse, max_pulse, batch_size, this_batch_id):\n    if validateMode:\n        labels = {'direction': Direction()}\n    else:\n        labels = None\n    print('labels', labels)\n    test_dataloader, test_dataset = make_dataloader2(db = \"dummy\",\n                                                selection = None,\n                                                pulsemaps = config['pulsemap'],\n                                                features = features,\n                                                truth = truth,\n                                                batch_ids = [this_batch_id],\n                                                batch_size = batch_size,\n                                                num_workers = config['num_workers'],\n                                                shuffle = False,\n                                                labels = labels,\n                                                index_column = config['index_column'],\n                                                truth_table = config['truth_table'],\n                                                max_len = 0,\n                                                max_pulse = max_pulse,\n                                                min_pulse = min_pulse,\n                                                )\n\n    if len(test_dataset) == 0:\n        print('skip inference')\n        return pd.DataFrame()\n    \n    model = build_model2(config, test_dataloader, test_dataset)\n\n    state_dict =  torch.load(CKPT, torch.device('cpu'))\n    if 'state_dict' in state_dict:\n        state_dict = state_dict['state_dict']\n    model.load_state_dict(state_dict)\n\n    USE_ORIG_PRED = True\n    if USE_ORIG_PRED:\n        event_ids = []\n        zenith = []\n        azimuth = []\n        preds = []\n        print('start predict')\n        with torch.no_grad():\n            model.eval()\n            model.to(f'cuda:{INF_DEVICES}')\n            for batch in tqdm(test_dataloader):\n                pred = model(batch.to(f'cuda:{INF_DEVICES}'))\n                #preds.append(pred[0])\n                if USE_ALL_FEA_IN_PRED:\n                    preds.append(torch.cat(pred, axis=-1))\n                else:\n                    preds.append(pred[0])\n                event_ids.append(batch.event_id)\n                if validateMode:\n                    zenith.append(batch.zenith)\n                    azimuth.append(batch.azimuth)\n        preds = torch.cat(preds).to('cpu').detach().numpy()\n        #results = pd.DataFrame(preds, columns=model.prediction_columns)\n        if USE_ALL_FEA_IN_PRED:\n            if preds.shape[1] == 128+8:\n                columns = ['direction_x','direction_y','direction_z','direction_kappa1','direction_x1','direction_y1','direction_z1','direction_kappa'] + [f'idx{i}' for i in range(128)]\n            else:\n                columns = ['direction_x','direction_y','direction_z','direction_kappa'] + [f'idx{i}' for i in range(128)]\n        else:\n            columns=model.prediction_columns\n        results = pd.DataFrame(preds, columns=columns)\n        results['event_id'] = torch.cat(event_ids).to('cpu').detach().numpy()\n        if validateMode:\n            results['zenith'] = torch.cat(zenith).to('cpu').numpy()\n            results['azimuth'] = torch.cat(azimuth).to('cpu').numpy()\n            \n        del zenith, azimuth, event_ids, preds\n    else:\n        results = model.predict_as_dataframe(\n            gpus = config['gpus'],\n            dataloader = test_dataloader,\n            prediction_columns=model.prediction_columns,\n            additional_attributes=model.additional_attributes,\n        )\n    gc.collect()\n    if validateMode:\n        results = convert_to_3d(results)\n        results = calculate_angular_error(results)\n        print('angular_error',results[\"angular_error\"].mean())\n    return results\n","metadata":{"papermill":{"duration":0.034826,"end_time":"2023-04-21T01:06:53.690478","exception":false,"start_time":"2023-04-21T01:06:53.655652","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-22T13:23:20.658900Z","iopub.execute_input":"2023-04-22T13:23:20.659570Z","iopub.status.idle":"2023-04-22T13:23:20.679797Z","shell.execute_reply.started":"2023-04-22T13:23:20.659524Z","shell.execute_reply":"2023-04-22T13:23:20.678627Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if validateMode:\n    BATCH_DIR = '../input/icecube-neutrinos-in-deep-ice/train'\n    #meta = pd.read_parquet('../input/icecubedatas/train_meta_651to660.parquet')\n    #META_DIR = '../work/test_valid10'\n    meta = pd.read_parquet('../input/icecubedatas/train_meta_656to660.parquet')\n    META_DIR = '../work/test_valid5'\n    #meta = pd.read_parquet('../input/icecubedatas/train_meta_660.parquet')\n    #META_DIR = '../work/test_valid1'\nelse:\n    BATCH_DIR = '../input/icecube-neutrinos-in-deep-ice/test'\n    META_DIR = '../work/test'\n    meta = pd.read_parquet('../input/icecube-neutrinos-in-deep-ice/test_meta.parquet')\n    \nWORK_DIR = META_DIR\nos.makedirs(META_DIR, exist_ok = True)\nCSV_META = f'{META_DIR}/test_meta.csv'\nif not os.path.exists(CSV_META):\n    meta.to_csv(CSV_META, index=False)\nbatch_ids = meta.batch_id.unique()\nfor batch_id in batch_ids:\n    out_file = f'{META_DIR}/meta_{batch_id}.parquet'\n    if os.path.exists(out_file):\n        continue\n    meta[meta.batch_id == batch_id].reset_index(drop=False).to_parquet(out_file, index=False)\ndel meta\n_ = gc.collect()","metadata":{"lines_to_next_cell":0,"papermill":{"duration":0.315953,"end_time":"2023-04-21T01:06:54.018551","exception":false,"start_time":"2023-04-21T01:06:53.702598","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-22T13:23:20.681040Z","iopub.execute_input":"2023-04-22T13:23:20.682102Z","iopub.status.idle":"2023-04-22T13:23:20.967947Z","shell.execute_reply.started":"2023-04-22T13:23:20.682062Z","shell.execute_reply":"2023-04-22T13:23:20.966855Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"lines_to_next_cell":2,"papermill":{"duration":0.012087,"end_time":"2023-04-21T01:06:54.043263","exception":false,"start_time":"2023-04-21T01:06:54.031176","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ONLY_AUX_FALSE = False\nrunName = '4l-ph3'\nCKPT = '../input/icecubedatas/base1-4l-lr2-batch1000-splitModel2-650x14-retryFromStart2-last.pth'\nCOLUMNS_NEAREST_NEIGHBOURS = [slice(0,3)]\nUSE_G = True\nUSE_PP = True\nDYNEDGE_LAYER_SIZE = [\n                (\n                    256,\n                    256,\n                ),\n                (\n                    256,\n                    256,\n                ),\n                (\n                    256,\n                    256,\n                ),\n                (\n                    256,\n                    256,\n                ),\n            ]\n\nresults_all_batch = []\nfor this_batch_id in batch_ids: #TODO\n    result0 = infer(0, 96, 400, this_batch_id); gc.collect()\n    result1 = infer(96, 140, 200, this_batch_id); gc.collect()\n    result2 = infer(140, 300, 100, this_batch_id); gc.collect()\n    result3 = infer(300, 1000, 20, this_batch_id); gc.collect()\n    result4 = infer(1000, 2000, 2, this_batch_id); gc.collect() # left all with FORCE_MAX_PULSE\n    result5 = infer(2000, 0, 1, this_batch_id); gc.collect() # left all with FORCE_MAX_PULSE\n    results = pd.concat([result0, result1, result2, result3, result4, result5]).sort_values('event_id')\n    results_all_batch.append(results)\n\n    if validateMode:\n        print('angular_error',this_batch_id, results[\"angular_error\"].mean())\n\n    del result0, result1, result2, result3, result4, result5, results\n    gc.collect()\n\nresults_all_batch = pd.concat(results_all_batch).sort_values('event_id').reset_index(drop=True)\nresults_all_batch['event_id'] = results_all_batch['event_id'].astype(int)\nresults_all_batch.to_csv(f'{WORK_DIR}/{runName}_{validateMode}.csv', index=False)\nresult1_files.append(f'{WORK_DIR}/{runName}_{validateMode}.csv')\n\nif validateMode:\n    print('angular_error',batch_ids, results_all_batch[\"angular_error\"].mean())\n\ndel results_all_batch\ngc.collect()","metadata":{"papermill":{"duration":0.026513,"end_time":"2023-04-21T01:06:54.081884","exception":false,"start_time":"2023-04-21T01:06:54.055371","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-22T13:23:20.969861Z","iopub.execute_input":"2023-04-22T13:23:20.970268Z","iopub.status.idle":"2023-04-22T13:23:26.759106Z","shell.execute_reply.started":"2023-04-22T13:23:20.970230Z","shell.execute_reply":"2023-04-22T13:23:26.757979Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ONLY_AUX_FALSE = False\nrunName = '4l-retry1-e2t10-ph3'\nCKPT = '../input/icecubedatas/base1-4l-splitModel2-650x2-theta10-retryFromStart3-last.pth'\nCOLUMNS_NEAREST_NEIGHBOURS = [slice(0,3)]\nUSE_G = True\nUSE_PP = True\nDYNEDGE_LAYER_SIZE = [\n                (\n                    256,\n                    256,\n                ),\n                (\n                    256,\n                    256,\n                ),\n                (\n                    256,\n                    256,\n                ),\n                (\n                    256,\n                    256,\n                ),\n            ]\n\nresults_all_batch = []\nfor this_batch_id in batch_ids: #TODO\n    result0 = infer(0, 96, 400, this_batch_id); gc.collect()\n    result1 = infer(96, 140, 200, this_batch_id); gc.collect()\n    result2 = infer(140, 300, 100, this_batch_id); gc.collect()\n    result3 = infer(300, 1000, 20, this_batch_id); gc.collect()\n    result4 = infer(1000, 2000, 2, this_batch_id); gc.collect() # left all with FORCE_MAX_PULSE\n    result5 = infer(2000, 0, 1, this_batch_id); gc.collect() # left all with FORCE_MAX_PULSE\n    results = pd.concat([result0, result1, result2, result3, result4, result5]).sort_values('event_id')\n    results_all_batch.append(results)\n\n    if validateMode:\n        print('angular_error',this_batch_id, results[\"angular_error\"].mean())\n\n    del result0, result1, result2, result3, result4, result5, results\n    gc.collect()\n\nresults_all_batch = pd.concat(results_all_batch).sort_values('event_id').reset_index(drop=True)\nresults_all_batch['event_id'] = results_all_batch['event_id'].astype(int)\nresults_all_batch.to_csv(f'{WORK_DIR}/{runName}_{validateMode}.csv', index=False)\nresult1_files.append(f'{WORK_DIR}/{runName}_{validateMode}.csv')\n\nif validateMode:\n    print('angular_error',batch_ids, results_all_batch[\"angular_error\"].mean())\n\ndel results_all_batch\ngc.collect()","metadata":{"papermill":{"duration":0.029224,"end_time":"2023-04-21T01:06:54.123115","exception":false,"start_time":"2023-04-21T01:06:54.093891","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-22T13:23:26.760895Z","iopub.execute_input":"2023-04-22T13:23:26.761681Z","iopub.status.idle":"2023-04-22T13:23:29.987644Z","shell.execute_reply.started":"2023-04-22T13:23:26.761635Z","shell.execute_reply":"2023-04-22T13:23:29.986516Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ONLY_AUX_FALSE = False\nrunName = '3lnoPP-ph3'\nCKPT = '../input/icecubedatas/base1-3l300p500b-noPP-650x2-retryFromStart2-epoch1299-val_tloss0.999595.ckpt' # 0.9689275622367859, 0.9656946659088135\nCOLUMNS_NEAREST_NEIGHBOURS = [slice(0,3)]\nUSE_G = False\nUSE_PP = False\nDYNEDGE_LAYER_SIZE = [\n                (\n                    256,\n                    256,\n                ),\n                (\n                    256,\n                    256,\n                ),\n                (\n                    256,\n                    256,\n                ),\n            ]\n\nresults_all_batch = []\nfor this_batch_id in batch_ids: #TODO\n    result0 = infer(0, 96, 400, this_batch_id); gc.collect()\n    result1 = infer(96, 140, 200, this_batch_id); gc.collect()\n    result2 = infer(140, 300, 100, this_batch_id); gc.collect()\n    result3 = infer(300, 1000, 20, this_batch_id); gc.collect()\n    result4 = infer(1000, 2000, 2, this_batch_id); gc.collect() # left all with FORCE_MAX_PULSE\n    result5 = infer(2000, 0, 1, this_batch_id); gc.collect() # left all with FORCE_MAX_PULSE\n    results = pd.concat([result0, result1, result2, result3, result4, result5]).sort_values('event_id')\n    results_all_batch.append(results)\n\n    if validateMode:\n        print('angular_error',this_batch_id, results[\"angular_error\"].mean())\n\n    del result0, result1, result2, result3, result4, result5, results\n    gc.collect()\n\nresults_all_batch = pd.concat(results_all_batch).sort_values('event_id').reset_index(drop=True)\nresults_all_batch['event_id'] = results_all_batch['event_id'].astype(int)\nresults_all_batch.to_csv(f'{WORK_DIR}/{runName}_{validateMode}.csv', index=False)\nresult1_files.append(f'{WORK_DIR}/{runName}_{validateMode}.csv')\n\nif validateMode:\n    print('angular_error',batch_ids, results_all_batch[\"angular_error\"].mean())\n\ndel results_all_batch\ngc.collect()","metadata":{"papermill":{"duration":0.026198,"end_time":"2023-04-21T01:06:54.161413","exception":false,"start_time":"2023-04-21T01:06:54.135215","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-22T13:23:29.989503Z","iopub.execute_input":"2023-04-22T13:23:29.990189Z","iopub.status.idle":"2023-04-22T13:23:33.517141Z","shell.execute_reply.started":"2023-04-22T13:23:29.990145Z","shell.execute_reply":"2023-04-22T13:23:33.515958Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ONLY_AUX_FALSE = False\nrunName = '3l4n-avg'\nCKPT = '../input/icecubedatas/base1-3l250p4n-batch1000-650x8-retryFromStart1-avg.pth' # 0.9705609679222107\nCOLUMNS_NEAREST_NEIGHBOURS = [slice(0,4)]\nUSE_G = True\nUSE_PP = True\nDYNEDGE_LAYER_SIZE = [\n                (\n                    256,\n                    256,\n                ),\n                (\n                    256,\n                    256,\n                ),\n                (\n                    256,\n                    256,\n                ),\n            ]\n\nresults_all_batch = []\nfor this_batch_id in batch_ids: #TODO\n    result0 = infer(0, 96, 400, this_batch_id); gc.collect()\n    result1 = infer(96, 140, 200, this_batch_id); gc.collect()\n    result2 = infer(140, 300, 100, this_batch_id); gc.collect()\n    result3 = infer(300, 1000, 20, this_batch_id); gc.collect()\n    result4 = infer(1000, 2000, 2, this_batch_id); gc.collect() # left all with FORCE_MAX_PULSE\n    result5 = infer(2000, 0, 1, this_batch_id); gc.collect() # left all with FORCE_MAX_PULSE\n    results = pd.concat([result0, result1, result2, result3, result4, result5]).sort_values('event_id')\n    results_all_batch.append(results)\n\n    if validateMode:\n        print('angular_error',this_batch_id, results[\"angular_error\"].mean())\n\n    del result0, result1, result2, result3, result4, result5, results\n    gc.collect()\n\nresults_all_batch = pd.concat(results_all_batch).sort_values('event_id').reset_index(drop=True)\nresults_all_batch['event_id'] = results_all_batch['event_id'].astype(int)\nresults_all_batch.to_csv(f'{WORK_DIR}/{runName}_{validateMode}.csv', index=False)\nresult1_files.append(f'{WORK_DIR}/{runName}_{validateMode}.csv')\n\nif validateMode:\n    print('angular_error',batch_ids, results_all_batch[\"angular_error\"].mean())\n\ndel results_all_batch\ngc.collect()","metadata":{"papermill":{"duration":0.025639,"end_time":"2023-04-21T01:06:54.199080","exception":false,"start_time":"2023-04-21T01:06:54.173441","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-22T13:23:33.518658Z","iopub.execute_input":"2023-04-22T13:23:33.519405Z","iopub.status.idle":"2023-04-22T13:23:36.207730Z","shell.execute_reply.started":"2023-04-22T13:23:33.519360Z","shell.execute_reply":"2023-04-22T13:23:36.206692Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ONLY_AUX_FALSE = False\nrunName = '4l-splitModel2-650x12-ph3'\nCKPT = '../input/icecubedatas/base1-4l-splitModel2-650x12-retryFromStart4-retry1-last.pth' #0.9661051034927368\nCOLUMNS_NEAREST_NEIGHBOURS = [slice(0,3)]\nUSE_G = True\nUSE_PP = True\nDYNEDGE_LAYER_SIZE = [\n                (\n                    256,\n                    256,\n                ),\n                (\n                    256,\n                    256,\n                ),\n                (\n                    256,\n                    256,\n                ),\n                (\n                    256,\n                    256,\n                ),\n            ]\n\nresults_all_batch = []\nfor this_batch_id in batch_ids: #TODO\n    result0 = infer(0, 96, 400, this_batch_id); gc.collect()\n    result1 = infer(96, 140, 200, this_batch_id); gc.collect()\n    result2 = infer(140, 300, 100, this_batch_id); gc.collect()\n    result3 = infer(300, 1000, 20, this_batch_id); gc.collect()\n    result4 = infer(1000, 2000, 2, this_batch_id); gc.collect() # left all with FORCE_MAX_PULSE\n    result5 = infer(2000, 0, 1, this_batch_id); gc.collect() # left all with FORCE_MAX_PULSE\n    results = pd.concat([result0, result1, result2, result3, result4, result5]).sort_values('event_id')\n    results_all_batch.append(results)\n\n    if validateMode:\n        print('angular_error',this_batch_id, results[\"angular_error\"].mean())\n\n    del result0, result1, result2, result3, result4, result5, results\n    gc.collect()\n\nresults_all_batch = pd.concat(results_all_batch).sort_values('event_id').reset_index(drop=True)\nresults_all_batch['event_id'] = results_all_batch['event_id'].astype(int)\nresults_all_batch.to_csv(f'{WORK_DIR}/{runName}_{validateMode}.csv', index=False)\nresult1_files.append(f'{WORK_DIR}/{runName}_{validateMode}.csv')\n\nif validateMode:\n    print('angular_error',batch_ids, results_all_batch[\"angular_error\"].mean())\n\ndel results_all_batch\ngc.collect()","metadata":{"papermill":{"duration":0.025795,"end_time":"2023-04-21T01:06:54.236921","exception":false,"start_time":"2023-04-21T01:06:54.211126","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-22T13:23:36.210580Z","iopub.execute_input":"2023-04-22T13:23:36.210935Z","iopub.status.idle":"2023-04-22T13:23:39.155464Z","shell.execute_reply.started":"2023-04-22T13:23:36.210888Z","shell.execute_reply":"2023-04-22T13:23:39.154316Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ONLY_AUX_FALSE = False\nrunName = '4l4D500p-batch500-splitModel2-650x20-ph3'\nCKPT = '../input/icecubedatas/base1-4l4D500p-batch500-splitModel2-650x20-retry4-last.pth' # 0.9676517248153687\nCOLUMNS_NEAREST_NEIGHBOURS = [slice(0,4)]\nUSE_G = True\nUSE_PP = True\nDYNEDGE_LAYER_SIZE = [\n                (\n                    256,\n                    256,\n                ),\n                (\n                    256,\n                    256,\n                ),\n                (\n                    256,\n                    256,\n                ),\n                (\n                    256,\n                    256,\n                ),\n            ]\n\nresults_all_batch = []\nfor this_batch_id in batch_ids: #TODO\n    result0 = infer(0, 96, 400, this_batch_id); gc.collect()\n    result1 = infer(96, 140, 200, this_batch_id); gc.collect()\n    result2 = infer(140, 300, 100, this_batch_id); gc.collect()\n    result3 = infer(300, 1000, 20, this_batch_id); gc.collect()\n    result4 = infer(1000, 2000, 2, this_batch_id); gc.collect() # left all with FORCE_MAX_PULSE\n    result5 = infer(2000, 0, 1, this_batch_id); gc.collect() # left all with FORCE_MAX_PULSE\n    results = pd.concat([result0, result1, result2, result3, result4, result5]).sort_values('event_id')\n    results_all_batch.append(results)\n\n    if validateMode:\n        print('angular_error',this_batch_id, results[\"angular_error\"].mean())\n\n    del result0, result1, result2, result3, result4, result5, results\n    gc.collect()\n\nresults_all_batch = pd.concat(results_all_batch).sort_values('event_id').reset_index(drop=True)\nresults_all_batch['event_id'] = results_all_batch['event_id'].astype(int)\nresults_all_batch.to_csv(f'{WORK_DIR}/{runName}_{validateMode}.csv', index=False)\nresult1_files.append(f'{WORK_DIR}/{runName}_{validateMode}.csv')\n\nif validateMode:\n    print('angular_error',batch_ids, results_all_batch[\"angular_error\"].mean())\n\ndel results_all_batch\ngc.collect()","metadata":{"lines_to_next_cell":0,"papermill":{"duration":0.033722,"end_time":"2023-04-21T01:06:54.283258","exception":false,"start_time":"2023-04-21T01:06:54.249536","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-22T13:23:39.157331Z","iopub.execute_input":"2023-04-22T13:23:39.157730Z","iopub.status.idle":"2023-04-22T13:23:42.155318Z","shell.execute_reply.started":"2023-04-22T13:23:39.157685Z","shell.execute_reply":"2023-04-22T13:23:42.154145Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## stacking part","metadata":{}},{"cell_type":"code","source":"def zeaz2xyz(az, ze):\n    z = np.cos(ze)\n    rz = np.sin(ze)\n    x = rz*np.cos(az)\n    y = rz*np.sin(az)\n    return x, y, z\n\ndef xyz2zeaz(x,y,z):\n    r = np.sqrt(x**2+y**2+z**2)\n    #print('R', r)\n    x = x/r\n    y = y/r\n    z = z/r\n    ze = np.arccos(z)\n    rz = np.sin(ze)\n    az = np.arccos(x/rz)\n    az = np.where(y < 0, np.pi*2-az, az)\n    az = np.nan_to_num(az,0)\n    return az, ze\n\nPRED_DIR = '../work/'\nMETA_DIR = '../input/icecube-neutrinos-in-deep-ice/train'\nclass MyDatasetFile(Dataset):\n    def __init__(self, \n                 runNames, \n                 batch_ids,\n                 tgt_cols = ['direction_x','direction_y','direction_z','direction_kappa', 'direction_x1','direction_y1','direction_z1','direction_kappa1'],\n                 use_mid_fea = True,\n                ):\n        self.batch_ids = batch_ids\n        self.runNames = runNames\n        self.tgt_cols = tgt_cols\n        self.use_mid_fea = use_mid_fea\n        self.this_batch_id = -1\n        self.this_batch_idx = -1\n        self.reset_epoch()\n        \n    def reset_epoch(self) -> None:\n        self.this_batch_idx += 1\n        if self.this_batch_idx >= len(self.batch_ids):\n            self.this_batch_idx = 0\n            \n        if self.this_batch_id == self.batch_ids[self.this_batch_idx]:\n            print('skip reset epoch ', self.this_batch_id, self.this_batch_idx)\n            return\n        else:\n            self.this_batch_id = self.batch_ids[self.this_batch_idx]\n            print('reset epoch to batch_id:', self.this_batch_id, self.this_batch_idx)\n\n        meta = pd.read_parquet(f'{META_DIR}/meta_{self.this_batch_id}.parquet').reset_index(drop=True)\n        meta['pulse_count'] = np.log1p(meta.last_pulse_index - meta.first_pulse_index + 1)\n        self.meta = meta\n        \n        result_list = []\n        for runName in self.runNames:\n            df = pd.read_pickle(f'{WORK_DIR}/{runName}_{self.this_batch_id}.pkl')\n            if 'direction_kappa' in df:\n                df['direction_kappa'] = np.log1p(df['direction_kappa'])\n            if 'direction_kappa1' in df:\n                df['direction_kappa1'] = np.log1p(df['direction_kappa1'])\n            columns = self.tgt_cols\n            if self.use_mid_fea:\n                columns = columns + [f'idx{i}' for i in range(128)]\n            result_list.append(df[columns].reset_index(drop=True))\n        self.X = pd.concat(result_list, axis=1).values.astype('float32')\n\n        self.event_ids = meta['event_id'].values\n        if 'azimuth' in meta:\n            #self.Y = np.stack(zeaz2xyz(results[0]['azimuth'], results[0]['zenith'])).T.astype('float32')\n            self.Y = np.stack(zeaz2xyz(meta['azimuth'], meta['zenith'])).T.astype('float32')\n            self.with_gt = True\n        else:\n            self.with_gt = False\n\n    def __len__(self):\n        return self.X.shape[0]\n    \n    def n_columns(self):\n        return self.X.shape[1]\n\n    def __getitem__(self, index):\n        x = self.X[index]        \n        event_id = self.event_ids[index]\n        if self.with_gt:\n            y = self.Y[index]\n        else:\n            y = None\n        return x, y, event_id\n\n\nTGT_COLS = ['direction_x','direction_y','direction_z','direction_kappa', 'direction_x1','direction_y1','direction_z1','direction_kappa1']\nclass MyDataset(Dataset):\n    def __init__(self, results, meta, mode='train'):\n        meta = meta.reset_index(drop=True)\n        self.mode = mode\n        result_list = []\n        for result in results:\n            if len(TGT_COLS):\n                columns = TGT_COLS\n            elif 'direction_x1' in result:\n                print('use 8 fea')\n                columns = ['direction_x','direction_y','direction_z','direction_kappa', 'direction_x1','direction_y1','direction_z1','direction_kappa1']\n            else:\n                columns = ['direction_x','direction_y','direction_z','direction_kappa']\n            if USE_MID_FEA:\n                columns = columns + [f'idx{i}' for i in range(128)]\n            result_list.append(result[columns].reset_index(drop=True))\n        self.X = pd.concat(result_list, axis=1).values.astype('float32')\n        self.event_ids = results[0]['event_id'].values\n        if 'azimuth' in meta:\n            #self.Y = np.stack(zeaz2xyz(results[0]['azimuth'], results[0]['zenith'])).T.astype('float32')\n            self.Y = np.stack(zeaz2xyz(meta['azimuth'], meta['zenith'])).T.astype('float32')\n            self.with_gt = True\n        else:\n            self.with_gt = False\n\n    def __len__(self):\n        return self.X.shape[0]\n    \n    def n_columns(self):\n        return self.X.shape[1]\n\n    def __getitem__(self, index):\n        x = self.X[index]\n        event_id = self.event_ids[index]\n        if self.mode == 'test':\n            return x,0,event_id\n        \n        if self.with_gt:\n            y = self.Y[index]\n        else:\n            y = None\n        return x, y, event_id        ","metadata":{"execution":{"iopub.status.busy":"2023-04-22T13:23:42.157613Z","iopub.execute_input":"2023-04-22T13:23:42.158015Z","iopub.status.idle":"2023-04-22T13:23:42.184956Z","shell.execute_reply.started":"2023-04-22T13:23:42.157968Z","shell.execute_reply":"2023-04-22T13:23:42.183883Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class StandardModel3(Model):\n    @save_model_config\n    def __init__(\n        self,\n        *,\n        tasks: Union[Task, List[Task]],\n        n_input_fea, \n        dataset,\n        optimizer_class: type = Adam,\n        optimizer_kwargs: Optional[Dict] = None,\n        scheduler_class: Optional[type] = None,\n        scheduler_kwargs: Optional[Dict] = None,\n        scheduler_config: Optional[Dict] = None,\n    ) -> None:\n        \"\"\"Construct `StandardModel`.\"\"\"\n        # Base class constructor\n        super().__init__()\n\n        # Check(s)\n        if isinstance(tasks, Task):\n            tasks = [tasks]\n        assert isinstance(tasks, (list, tuple))\n        assert all(isinstance(task, Task) for task in tasks)\n\n        # Member variable(s)\n        self._tasks = ModuleList(tasks)\n        self._optimizer_class = optimizer_class\n        self._optimizer_kwargs = optimizer_kwargs or dict()\n        self._scheduler_class = scheduler_class\n        self._scheduler_kwargs = scheduler_kwargs or dict()\n        self._scheduler_config = scheduler_config or dict()\n        self._n_input_fea = n_input_fea\n        self._dataset = dataset\n        \n        mlp_layers = []\n        layer_sizes = [n_input_fea, HIDDEN_SIZE, HIDDEN_SIZE, HIDDEN_SIZE] # todo1\n        for nb_in, nb_out in zip(layer_sizes[:-1], layer_sizes[1:]):\n            mlp_layers.append(torch.nn.Linear(nb_in, nb_out))\n            mlp_layers.append(torch.nn.LeakyReLU())\n            mlp_layers.append(torch.nn.Dropout(DROPOUT_PH2_MODEL))\n        last_posting_layer_output_dim = nb_out\n\n        self._mlp = torch.nn.Sequential(*mlp_layers)\n\n            \n\n    def configure_optimizers(self) -> Dict[str, Any]:\n        \"\"\"Configure the model's optimizer(s).\"\"\"\n        optimizer = self._optimizer_class(\n            self.parameters(), **self._optimizer_kwargs\n        )\n        config = {\n            \"optimizer\": optimizer,\n        }\n        if self._scheduler_class is not None:\n            scheduler = self._scheduler_class(\n                optimizer, **self._scheduler_kwargs\n            )\n            config.update(\n                {\n                    \"lr_scheduler\": {\n                        \"scheduler\": scheduler,\n                        **self._scheduler_config,\n                    },\n                }\n            )\n        return config\n\n    def forward(self, x):\n        x = self._mlp(x)\n        x = [task(x) for task in self._tasks]\n        return x\n\n    def training_step(self, xye, idx) -> Tensor:\n        \"\"\"Perform training step.\"\"\"\n        x,y,event_ids = xye\n        preds = self(x)\n        batch = Data(x=x, direction=y)\n        vlosses = self._tasks[1].compute_loss(preds[1], batch)\n        vloss = torch.sum(vlosses)\n        \n        tlosses = self._tasks[0].compute_loss(preds[0], batch)\n        tloss = torch.sum(tlosses)\n\n        loss = vloss*0.1 + tloss\n        return {\"loss\": loss, 'vloss': vloss, 'tloss': tloss}\n\n    def validation_step(self, xye, idx) -> Tensor:\n        \"\"\"Perform validation step.\"\"\"\n        x,y,event_ids = xye\n        preds = self(x)\n        batch = Data(x=x, direction=y)\n        vlosses = self._tasks[1].compute_loss(preds[1], batch)\n        vloss = torch.sum(vlosses)\n        \n        tlosses = self._tasks[0].compute_loss(preds[0], batch)\n        tloss = torch.sum(tlosses)\n        loss = vloss*0.1 + tloss\n        return {\"loss\": loss, 'vloss': vloss, 'tloss': tloss}\n\n    def inference(self) -> None:\n        \"\"\"Activate inference mode.\"\"\"\n        for task in self._tasks:\n            task.inference()\n\n    def train(self, mode: bool = True) -> \"Model\":\n        \"\"\"Deactivate inference mode.\"\"\"\n        super().train(mode)\n        if mode:\n            for task in self._tasks:\n                task.train_eval()\n        return self\n\n    def predict(\n        self,\n        dataloader: DataLoader,\n        gpus: Optional[Union[List[int], int]] = None,\n        distribution_strategy: Optional[str] = None,\n    ) -> List[Tensor]:\n        \"\"\"Return predictions for `dataloader`.\"\"\"\n        self.inference()\n        return super().predict(\n            dataloader=dataloader,\n            gpus=gpus,\n            distribution_strategy=distribution_strategy,\n        )\n    \n    def training_epoch_end(self, training_step_outputs):\n        loss = torch.stack([x[\"loss\"] for x in training_step_outputs]).mean()\n        vloss = torch.stack([x[\"vloss\"] for x in training_step_outputs]).mean()\n        tloss = torch.stack([x[\"tloss\"] for x in training_step_outputs]).mean()\n        self.log_dict(\n            {\"trn_loss\": loss, \"trn_vloss\": vloss, \"trn_tloss\": tloss},\n            prog_bar=True,\n            sync_dist=True,\n        )\n        print(f'epoch:{self.current_epoch}, train loss:{loss.item()}, tloss:{tloss.item()}, vloss:{vloss.item()}')\n        self._dataset.reset_epoch()\n        \n    def validation_epoch_end(self, validation_step_outputs):\n        loss = torch.stack([x[\"loss\"] for x in validation_step_outputs]).mean()\n        vloss = torch.stack([x[\"vloss\"] for x in validation_step_outputs]).mean()\n        tloss = torch.stack([x[\"tloss\"] for x in validation_step_outputs]).mean()\n        self.log_dict(\n            {\"val_loss\": loss, \"val_vloss\": vloss, \"val_tloss\": tloss},\n            prog_bar=True,\n            sync_dist=True,\n        )\n        print(f'epoch:{self.current_epoch}, valid loss:{loss.item()}, tloss:{tloss.item()}, vloss:{vloss.item()}')\n\ndef build_model3(config, dataloader, dataset) -> StandardModel2:\n    \"\"\"Builds GNN from config\"\"\"\n    # Building model\n\n    if config[\"target\"] == 'direction':\n        task = DirectionReconstructionWithKappa2(\n            hidden_size=HIDDEN_SIZE,\n            target_labels=config[\"target\"],\n            loss_function=VonMisesFisher3DLoss(),\n        )\n        task2 = DirectionReconstructionWithKappa2(\n            hidden_size=HIDDEN_SIZE,\n            target_labels=config[\"target\"],\n            loss_function=DistanceLoss2(),\n        )\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 = StandardModel3(\n        tasks=[task2, task],\n        n_input_fea=N_INPUT_FEA,\n        dataset=dataset,\n        optimizer_class=Adam,\n        optimizer_kwargs={\"lr\": 1e-03, \"eps\": 1e-03},\n        #optimizer_class=Lion,\n        #optimizer_kwargs={\"lr\": 1e-04},\n        scheduler_class=PiecewiseLinearLR,\n        scheduler_kwargs={\n            \"milestones\": [\n                0,\n                10  * len(dataloader)//(len(config['gpus'])*config['accumulate_grad_batches'][0]),\n                len(dataloader)*config[\"fit\"][\"max_epochs\"]//(len(config['gpus'])*config['accumulate_grad_batches'][0]*2),\n                len(dataloader)*config[\"fit\"][\"max_epochs\"]//(len(config['gpus'])*config['accumulate_grad_batches'][0]),                \n            ],\n            \"factors\": [1e-03, 1, 1, 1e-03],\n            \"verbose\": config[\"scheduler_verbose\"],\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-22T13:23:42.186855Z","iopub.execute_input":"2023-04-22T13:23:42.187338Z","iopub.status.idle":"2023-04-22T13:23:42.218958Z","shell.execute_reply.started":"2023-04-22T13:23:42.187299Z","shell.execute_reply":"2023-04-22T13:23:42.217996Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def infer2(test_dataloader, model):\n    USE_ORIG_PRED = True\n    if USE_ORIG_PRED:\n        event_ids = []\n        zenith = []\n        azimuth = []\n        preds = []\n        print('start predict')\n        with torch.no_grad():\n            model.eval()\n            model.to(f'cuda:{INF_DEVICES}')\n            for batch in tqdm(test_dataloader):\n                pred = model(batch[0].to(f'cuda:{INF_DEVICES}'))\n                #preds.append(pred[0])\n                if USE_ALL_FEA_IN_PRED:\n                    preds.append(torch.cat(pred, axis=-1))\n                else:\n                    preds.append(pred[0])\n                event_ids.append(batch[2])\n                if validateMode:\n                    azimuth.append(batch[1][0])\n                    zenith.append(batch[1][1])\n        preds = torch.cat(preds).to('cpu').detach().numpy()\n        #results = pd.DataFrame(preds, columns=model.prediction_columns)\n        if USE_ALL_FEA_IN_PRED:\n            if preds.shape[1] == 128+8:\n                columns = ['direction_x','direction_y','direction_z','direction_kappa1','direction_x1','direction_y1','direction_z1','direction_kappa'] + [f'idx{i}' for i in range(128)]\n            else:\n                columns = ['direction_x','direction_y','direction_z','direction_kappa'] + [f'idx{i}' for i in range(128)]\n        else:\n            if preds.shape[1] == 8:\n                columns=['direction_x','direction_y','direction_z','direction_kappa1','direction_x1','direction_y1','direction_z1','direction_kappa']\n            else:\n                columns=['direction_x','direction_y','direction_z','direction_kappa']\n        results = pd.DataFrame(preds, columns=columns)\n        results['event_id'] = torch.cat(event_ids).to('cpu').detach().numpy()\n            \n        del zenith, azimuth, event_ids, preds\n    else:\n        results = model.predict_as_dataframe(\n            gpus = config['gpus'],\n            dataloader = test_dataloader,\n            prediction_columns=model.prediction_columns,\n            additional_attributes=model.additional_attributes,\n        )\n    gc.collect()\n    return results\n\n\ndef calculate_angular_error2(df_pred, df_gt):\n    df_gt = df_gt.reset_index(drop=True)\n    df_pred = df_pred.reset_index(drop=True)\n    df_pred['angular_error'] = np.arccos(df_gt['true_x']*df_pred['direction_x'] + df_gt['true_y']*df_pred['direction_y'] + df_gt['true_z']*df_pred['direction_z'])\n    return df_pred\n\ndef 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\n","metadata":{"lines_to_next_cell":2,"execution":{"iopub.status.busy":"2023-04-22T13:23:42.220593Z","iopub.execute_input":"2023-04-22T13:23:42.221020Z","iopub.status.idle":"2023-04-22T13:23:42.237125Z","shell.execute_reply.started":"2023-04-22T13:23:42.220987Z","shell.execute_reply":"2023-04-22T13:23:42.236043Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import time\n\nHIDDEN_SIZE = 512\nDROPOUT_PH2_MODEL = 0.0\nN_INPUT_FEA = 136*len(result1_files)\n\nUSE_ALL_FEA_IN_PRED=False\nINF_DEVICES = 0\nUSE_MID_FEA = True\nCKPT = '../input/icecubedatas/base3-ens20-last.pth' # 0.9640\nCKPT = '../input/icecubedatas/base4-ens23-last.pth' # 0.9637\nCKPT = '../input/icecubedatas/base5-ens26-last.pth' # 0.9635\nCKPT = '../input/icecubedatas/base5-ens32-6model-3layer-659batch-last.pth' # 0.963? todo1 \n\nchunksize=50000\nmeta_chunks = pd.read_csv(CSV_META, chunksize=chunksize)\nchunks = [pd.read_csv(f, chunksize=chunksize) for f in result1_files]\nresult_list = []\n\nfirst_flg = True\nwhile True:\n    start_time = time.time()\n    try:\n        result_df = []\n        for c in chunks:\n            df = next(c)\n            if 'direction_kappa' in df:\n                df['direction_kappa'] = np.log1p(df['direction_kappa'])\n            if 'direction_kappa1' in df:\n                df['direction_kappa1'] = np.log1p(df['direction_kappa1'])\n            result_df.append(df)\n        this_meta = next(meta_chunks)\n        this_meta['pulse_count'] = np.log1p(this_meta.last_pulse_index - this_meta.first_pulse_index + 1)\n        test_dataset = MyDataset(result_df, this_meta, mode='test')\n        test_dataloader = DataLoader(test_dataset, batch_size=config['batch_size'], shuffle=False, num_workers=2)\n        if first_flg:\n            model = build_model3(config = config, dataloader = test_dataloader, dataset = test_dataset)\n            state_dict =  torch.load(CKPT, torch.device('cpu'))\n            if 'state_dict' in state_dict:\n                state_dict = state_dict['state_dict']\n            model.load_state_dict(state_dict)\n            first_flg = False\n\n        results2 = infer2(test_dataloader, model)\n        if validateMode:\n            this_meta = convert_to_3d(this_meta)\n            results2 = calculate_angular_error2(results2, this_meta)\n            print(results2.angular_error.mean())\n        result_list.append(results2)\n    except StopIteration:\n        break\n    del result_df\n    gc.collect()\n    print('total time', time.time()-start_time)\n    \nresults2 = pd.concat(result_list)\nif validateMode:\n    print('all', results2.angular_error.mean())\n    print('660', results2.angular_error[-153924:].mean())","metadata":{"execution":{"iopub.status.busy":"2023-04-22T13:23:42.240195Z","iopub.execute_input":"2023-04-22T13:23:42.240870Z","iopub.status.idle":"2023-04-22T13:23:42.903340Z","shell.execute_reply.started":"2023-04-22T13:23:42.240832Z","shell.execute_reply":"2023-04-22T13:23:42.902150Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def prepare_dataframe(df, angle_post_fix = '', vec_post_fix = '') -> pd.DataFrame:\n    r = np.sqrt(df['direction_x'+ vec_post_fix]**2 + df['direction_y'+ vec_post_fix]**2 + df['direction_z' + vec_post_fix]**2)\n    df['zenith' + angle_post_fix] = np.arccos(df['direction_z'+ vec_post_fix]/r)\n    df['azimuth'+ angle_post_fix] = np.arctan2(df['direction_y'+ vec_post_fix],df['direction_x' + vec_post_fix]) #np.sign(results['true_y'])*np.arccos((results['true_x'])/(np.sqrt(results['true_x']**2 + results['true_y']**2)))\n    df['azimuth'+ angle_post_fix][df['azimuth'  + angle_post_fix]<0] = df['azimuth'  + angle_post_fix][df['azimuth'  +  angle_post_fix]<0] + 2*np.pi \n    return df\n\nresults2 = prepare_dataframe(results2)\nresults2[['event_id','azimuth','zenith']].sort_values('event_id').to_csv(\"submission.csv\", index = False)\nresults2[['event_id','azimuth','zenith']].sort_values('event_id')","metadata":{"lines_to_next_cell":0,"execution":{"iopub.status.busy":"2023-04-22T13:23:42.905835Z","iopub.execute_input":"2023-04-22T13:23:42.906604Z","iopub.status.idle":"2023-04-22T13:23:42.938553Z","shell.execute_reply.started":"2023-04-22T13:23:42.906571Z","shell.execute_reply":"2023-04-22T13:23:42.937411Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"lines_to_next_cell":2},"execution_count":null,"outputs":[]}]}