{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# Transformer packages\n!pip install -q dm-haiku==0.0.9\n!pip install -q einops==0.6.0\n\n# GraphNet packages\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 -q /kaggle/working/software/dependencies/torch-1.11.0+cu115-cp37-cp37m-linux_x86_64.whl\n!pip install -q /kaggle/working/software/dependencies/torch_cluster-1.6.0-cp37-cp37m-linux_x86_64.whl\n!pip install -q /kaggle/working/software/dependencies/torch_scatter-2.0.9-cp37-cp37m-linux_x86_64.whl\n!pip install -q /kaggle/working/software/dependencies/torch_sparse-0.6.13-cp37-cp37m-linux_x86_64.whl\n!pip install -q /kaggle/working/software/dependencies/torch_geometric-2.0.4.tar.gz\n\n# Install GraphNeT\n!cd software/graphnet;pip install -q --no-index --find-links=\"/kaggle/working/software/dependencies\" -e .[torch]\n\n# Install dataset dependencies\n# !pip install -q /kaggle/input/pyarrow-1100/pyarrow-11.0.0-cp37-cp37m-manylinux_2_17_x86_64.manylinux2014_x86_64.whl\n!pip install -q polars==0.16.4","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-04-27T16:00:34.092984Z","iopub.execute_input":"2023-04-27T16:00:34.093714Z","iopub.status.idle":"2023-04-27T16:03:10.714692Z","shell.execute_reply.started":"2023-04-27T16:00:34.093663Z","shell.execute_reply":"2023-04-27T16:03:10.713431Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Append to PATH\nimport sys\nsys.path.append('/kaggle/working/software/graphnet/src')\n\n# Disable XLA preallocation so we can run both models on a single GPU\nimport os\nos.environ[\"XLA_PYTHON_CLIENT_PREALLOCATE\"] = \"false\"","metadata":{"execution":{"iopub.status.busy":"2023-04-27T16:03:10.718317Z","iopub.execute_input":"2023-04-27T16:03:10.718618Z","iopub.status.idle":"2023-04-27T16:03:10.724431Z","shell.execute_reply.started":"2023-04-27T16:03:10.718584Z","shell.execute_reply":"2023-04-27T16:03:10.723294Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Common imports\nimport gc\nimport math\nimport random\nimport torch\nfrom torch.utils import data\nfrom pathlib import Path\nfrom functools import reduce\nimport pandas as pd\nfrom operator import or_\nfrom collections import namedtuple\nimport polars\n\n# Transformer imports\nimport jax\nimport jax.numpy as jnp\nimport pickle\nimport haiku as hk\nfrom flax.jax_utils import prefetch_to_device, replicate\nfrom dataclasses import dataclass\nfrom tqdm import tqdm\nfrom functools import partial\nfrom itertools import islice\nfrom einops import rearrange, repeat\nfrom typing import Optional, Any\nfrom functools import wraps\n\n# GraphNet imports\nimport numpy as np\nfrom torch import nn\nfrom torch import LongTensor, Tensor\n\nfrom scipy.interpolate import interp1d\nfrom sklearn.preprocessing import RobustScaler\n\nimport pyarrow as pa\nimport pyarrow.parquet as pq\nfrom functools import reduce\nfrom typing import Any, Callable, List, Optional, Sequence, Tuple, Union\n\nimport pytorch_lightning as pl\n\nfrom torch_geometric.nn import EdgeConv\nfrom torch_geometric.nn.pool import knn_graph\nfrom torch_geometric.data import Data, Dataset\nfrom torch_geometric.loader import DataLoader\nfrom torch_geometric.typing import Adj\nfrom torch_scatter import scatter_max, scatter_mean, scatter_min, scatter_sum\n\nfrom graphnet.models.gnn.gnn import GNN\nfrom graphnet.models.utils import calculate_xyzt_homophily\nfrom graphnet.utilities.config import save_model_config\nfrom graphnet.models.task.reconstruction import (\n    AzimuthReconstructionWithKappa,\n    ZenithReconstruction,\n    DirectionReconstructionWithKappa,\n)\nfrom graphnet.training.loss_functions import VonMisesFisher2DLoss, VonMisesFisher3DLoss\nfrom graphnet.models.graph_builders import KNNGraphBuilder\n\nprint(\"All imports done.\")","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-04-27T16:03:10.728461Z","iopub.execute_input":"2023-04-27T16:03:10.729061Z","iopub.status.idle":"2023-04-27T16:03:11.183201Z","shell.execute_reply.started":"2023-04-27T16:03:10.729023Z","shell.execute_reply":"2023-04-27T16:03:11.181969Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# TEST BATCH SIZE & DEBUG\nDEBUG_DEBUG = False\nDEBUG_BATCHES = (660,)\n\nif DEBUG_DEBUG:\n    !rm -r valid-dataset/\n    !mkdir -p valid-dataset/test\n    for bidx in DEBUG_BATCHES:\n        !ln -s /kaggle/input/icecube-neutrinos-in-deep-ice/train/batch_{bidx}.parquet valid-dataset/test/\n    !ln -s /kaggle/input/icecube-neutrinos-in-deep-ice/train_meta.parquet valid-dataset/test_meta.parquet\n    !ln -s /kaggle/input/icecube-neutrinos-in-deep-ice/sensor_geometry.csv valid-dataset/\n    !ls valid-dataset/\n    !ls valid-dataset/test/","metadata":{"execution":{"iopub.status.busy":"2023-04-27T16:03:11.187262Z","iopub.execute_input":"2023-04-27T16:03:11.187561Z","iopub.status.idle":"2023-04-27T16:03:11.205368Z","shell.execute_reply.started":"2023-04-27T16:03:11.187533Z","shell.execute_reply":"2023-04-27T16:03:11.204011Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Transformer config\n@dataclass\nclass Config:\n    dataset: Path = Path(\"/kaggle/input/icecube-neutrinos-in-deep-ice\" if not DEBUG_DEBUG else \"/kaggle/working/valid-dataset\")\n    subset: str = \"test\"\n    seed: int = 1337\n    batch_size: int = 64\n    num_workers: int = 0\n    num_epochs: int = 1\n    iter_limit: int = 80000\n    learning_rate: float = 2e-6  # 1e-5\n    weight_decay: float = 0.0\n    grad_clip: Optional[float] = 10.0\n    num_bins: int = 256\n    log_dir: Path = Path(\"runs/\")\n    save_freq: int = 20000\n    log_freq: int = 50\n    infer_method: str = \"argmax\"\n    checkpoint: Optional[Path] = Path(\n        \"/kaggle/input/icecube-5th-place-models/transformer_model.pth\",\n    )\n\n\nconfig = Config()\nprint(config)\n\n\n# GraphNet config\nDATASET_PATH = \"/kaggle/input/icecube-neutrinos-in-deep-ice\" if not DEBUG_DEBUG else \"/kaggle/working/valid-dataset\" \nGEOMETRY_PATH = \"/kaggle/input/icecube-neutrinos-in-deep-ice/sensor_geometry.csv\"\nTRANSPARENCY_PATH = \"/kaggle/input/icecubetransparency/ice_transparency.txt\"\nCHPT_PATH = \"/kaggle/input/icecube-5th-place-models/graphnet_model.ckpt\"\nTTA_ENABLED = True\nNEM_ENABLED = True\nFEATURES_SUBSET_REST = 48\nREADOUT_FIRST_LAYER_SIZE = 256","metadata":{"execution":{"iopub.status.busy":"2023-04-27T16:03:11.206697Z","iopub.execute_input":"2023-04-27T16:03:11.207161Z","iopub.status.idle":"2023-04-27T16:03:11.223697Z","shell.execute_reply.started":"2023-04-27T16:03:11.207121Z","shell.execute_reply":"2023-04-27T16:03:11.222616Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset\n\nWe have two dataset implementations, one using Polars and other using PyArrow. PyArrow version works much nicer with PyTorch as it allows for multiple workers and reads parquet files much faster. However because we were unable to install PyArrow v11 in Kaggle environment (it works on TPU image but not on GPU one) we needed to fallback to Polars version, which forces us to use single worker for dataloading.","metadata":{}},{"cell_type":"code","source":"Event = namedtuple(\n    \"Event\",\n    [\n        \"event_id\",\n        \"azimuth\",\n        \"zenith\",\n        \"sensor_id\",\n        \"time\",\n        \"charge\",\n        \"auxiliary\",  # False when \"important\"\n    ],\n    defaults=[None, None, None, None, None, None, None],\n)\n\n\ndef identity(x):\n    return x\n\n\nclass IceCubeNeutrinosPolars(data.IterableDataset):\n    \"\"\"\n    PyTorch dataset for loading IceCube data. This implementation uses polars.\n    \"\"\"\n    def __init__(\n        self,\n        root_path,\n        subset=\"train\",\n        transform=identity,\n        shuffle_files=True,\n        index=None,\n        parts=None,\n    ):  \n        assert subset.lower() in (\"train\", \"test\", \"valid\")\n        self.root_path = Path(root_path)\n        self.subset = subset.lower()\n        self._filenames = sorted(list((self.root_path / subset).glob(\"*.parquet\")))\n        self.transform = transform\n        self._meta = None\n        self._meta_cols = [\"event_id\"]\n        if self.subset != \"test\":\n            self._meta_cols += [\"azimuth\", \"zenith\"]\n\n    def _iterate_over_batch_file(self, filename):\n        dataframe = (\n            polars.read_parquet(filename, parallel=\"none\")\n            .join(self._meta, on=\"event_id\", how=\"left\")\n            .shrink_to_fit(in_place=True)\n        )\n\n        for events in dataframe.partition_by(\"event_id\"):\n            yield Event(\n                **{col: events[col].unique().item() for col in self._meta_cols},\n                **{\n                    field: torch.from_numpy(events[field].to_numpy().copy())\n                    for field in Event._fields[-4:]\n                },\n            )\n\n        del dataframe\n\n    def __iter__(self):\n        worker_info = data.get_worker_info()\n\n        assert torch.multiprocessing.get_start_method() != \"fork\" or worker_info is None, (\n            'Using \"fork\" policy in multiprocessing will cause '\n            \"dataloader to hang due to polars not supporting it,\"\n            'please use torch.multiprocessing.set_start_method(\"spawn\")'\n        )\n\n        my_filenames = self._filenames\n        if worker_info is not None:\n            my_filenames = my_filenames[worker_info.id :: worker_info.num_workers]\n\n        for filename in my_filenames:\n            batch_id = int(filename.stem.split(\"_\")[-1])\n            self._meta = (\n                polars.scan_parquet(self.root_path / f\"{self.subset}_meta.parquet\", low_memory=True)\n                .filter(polars.col(\"batch_id\") == batch_id)\n                .select(self._meta_cols)\n                .collect()\n                .shrink_to_fit(in_place=True)\n            )\n            \n            for event in self._iterate_over_batch_file(filename):\n                yield self.transform(event)\n                \n\nIceCubeNeutrinos = IceCubeNeutrinosPolars","metadata":{"execution":{"iopub.status.busy":"2023-04-27T16:03:11.225432Z","iopub.execute_input":"2023-04-27T16:03:11.225788Z","iopub.status.idle":"2023-04-27T16:03:11.241255Z","shell.execute_reply.started":"2023-04-27T16:03:11.225751Z","shell.execute_reply":"2023-04-27T16:03:11.240189Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Transformer","metadata":{}},{"cell_type":"code","source":"class AzimuthZenithEncoder:\n    \"\"\"\n    This class discretizes `azimuth` and `zenith` angles.\n    Both `azimuth` and `zenith` bins are splitted equally.\n    \"\"\"\n\n    def __init__(self, azimuth_bins=128, zenith_bins=128):\n        self._azimuth = torch.linspace(0.0, 2.0 * torch.pi, azimuth_bins)\n        self._zenith = torch.linspace(0.0, torch.pi, zenith_bins)\n\n    def soft_targets(self, azimuth, zenith):\n        \"\"\"\n        Create a soft targets for discretized `azimuth` and `zenith`.\n        For `azimuth` it wraps around circle. So 0 and 2pi have same values.\n        \"\"\"\n        az = F.softmax(\n            (1.0 + torch.cos(self._azimuth - azimuth)) ** 7.2, dim=-1\n        )\n        zen = F.softmax(\n            (torch.pi - torch.abs(self._zenith - zenith)) * 40.0, dim=-1\n        )\n        return az, zen\n\n\nclass PositionalFeatures:\n    def __init__(\n        self,\n        sensor_path: str,\n        length: int = 128,\n        num_neighbours: int = 8,\n        num_vectors: int = 1000,\n        num_bins: int = 128,\n        training: bool = True,\n    ):\n        self.length = length\n        self.training = training\n        sensors_positions = torch.from_numpy(\n            pd.read_csv(sensor_path)[[\"x\", \"y\", \"z\"]].to_numpy()\n        )\n        self._sensors, self._t_valid = self._prepare_sensors(\n            sensors_positions.float()\n        )  # [S, 3]\n        self._az_encoder = AzimuthZenithEncoder(num_bins, num_bins)\n\n    def _prepare_sensors(self, sensors_positions: torch.Tensor):\n        \"\"\"\n        Normalize sensor positions and return valid time window.\n        \"\"\"\n        c_const = 0.299792458  # speed of light [m/ns]\n\n        # Sensor Min / Max Coordinates\n        xyz_min = sensors_positions.min(dim=0).values\n        xyz_max = sensors_positions.max(dim=0).values\n\n        detector_length = (xyz_max - xyz_min).pow(2).sum(-1).sqrt()\n        t_valid = detector_length / c_const\n\n        return sensors_positions / 600.0, t_valid\n\n    def _direction(self, azimuth, zenith):\n        dx = math.sin(zenith) * math.cos(azimuth)\n        dy = math.sin(zenith) * math.sin(azimuth)\n        dz = math.cos(zenith)\n        return dx, dy, dz\n\n    def _pad(self, tensor):\n        \"\"\"\n        Pad for \"ans\" token at the begining, and up to length at the end.\n        \"\"\"\n        pad = (0, 0) * (tensor.ndim - 1) + (1, self.length - tensor.size(0) - 1)\n        return torch.nn.functional.pad(tensor, pad=pad)\n\n    def __call__(self, event: Event) -> torch.Tensor:\n        time = event.time.double()\n        time = time - time.min()\n\n        t_peak = time[event.charge.argmax()]\n        t_valid_min = t_peak - self._t_valid * 2.0\n        t_valid_max = t_peak + self._t_valid * 2.0\n        t_valid_mask = (t_valid_min <= time) & (time <= t_valid_max)\n\n        # Rank\n        rank = 2 * (1 - event.auxiliary.long()) + t_valid_mask.long()\n\n        # Sort by Rank and Charge (important, goes backward)\n        charge_order = torch.sort(event.charge, dim=-1, stable=True).indices\n        rank_order = torch.sort(rank[charge_order], dim=-1, stable=True).indices\n        order = charge_order[rank_order][\n            -self.length + 1 :\n        ]  # +1 to account for \"ans\" token\n        order = order[torch.argsort(time[order])]  # sort by time (for lstm)\n\n        sensors_ids = event.sensor_id[order].long()\n        sensors_positions = self._sensors[sensors_ids]  # [L, 3]\n\n        is_core = sensors_ids >= 4680\n\n        time = time[order]\n        time = (time - time.min()) / 1000.0\n        time = time.float()\n\n        # for LSTM\n        time_diff = torch.diff(time, prepend=torch.tensor([0]))\n\n        charge = event.charge[order] / 300.0\n\n        mask = torch.zeros(self.length, dtype=torch.bool)\n        mask[: len(order) + 1] = True  # those will not be ignored\n\n        ans_mask = torch.ones(self.length, dtype=torch.float32)\n        ans_mask[0] = 0.0  # to clear out first token in transformer\n\n        target_data = {}\n        if self.training:\n            dx, dy, dz = self._direction(event.azimuth, event.zenith)\n            azimuth = torch.tensor(event.azimuth).float() / (torch.pi * 2.0)\n            zenith = torch.tensor(event.zenith).float() / torch.pi\n            azimuth_target, zenith_target = self._az_encoder.soft_targets(\n                azimuth * torch.pi * 2.0, zenith * torch.pi\n            )\n            target_data = {\n                \"direction\": torch.tensor(\n                    [dx, dy, dz], dtype=torch.float32\n                ),  # [3]\n                \"azimuth\": azimuth,  # ()\n                \"zenith\": zenith,  # ()\n                \"azimuth_zenith\": torch.stack([azimuth, zenith]),  # [2]\n                \"azimuth_target\": azimuth_target,\n                \"zenith_target\": zenith_target,\n            }\n\n        return {\n            \"event_id\": event.event_id,\n            \"x\": self._pad(sensors_positions[:, 0]),  # [L]\n            \"y\": self._pad(sensors_positions[:, 1]),  # [L]\n            \"z\": self._pad(sensors_positions[:, 2]),  # [L]\n            \"auxiliary\": self._pad(event.auxiliary[order].long()),  # [L]\n            \"time\": self._pad(time.float()),  # [L]\n            \"charge\": self._pad(charge.float()),  # [L]\n            \"core\": self._pad(is_core.long()),  # [L]\n            \"mask\": mask,  # [L]\n            \"ans_mask\": ans_mask,  # [L]\n            \"time_diff\": self._pad(time_diff.float()),  # [L]\n            \"rank\": self._pad(rank[order].long()),  # [N, L]\n            **target_data,\n        }\n\n\ndef jax_collate_repeat(num_devices, batch):\n    \"\"\"\n    Collate function that repeats given batch on each device available.\n    \"\"\"\n    # batch = data.default_collate(batch)  # works in pytorch 1.13\n    batch = data._utils.collate.default_collate(batch)\n    return jax.tree_map(\n        lambda x: repeat(x, \"n ... -> d n ...\", d=num_devices), batch\n    )","metadata":{"execution":{"iopub.status.busy":"2023-04-27T16:03:11.243759Z","iopub.execute_input":"2023-04-27T16:03:11.244095Z","iopub.status.idle":"2023-04-27T16:03:11.271173Z","shell.execute_reply.started":"2023-04-27T16:03:11.244068Z","shell.execute_reply":"2023-04-27T16:03:11.270054Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Transformer model definition\n\nThe provided code defines a custom Transformer-based architecture tailored for processing data from an array of sensors with various features. The model predicts azimuth and zenith angles based on the input features, which include coordinates `(x, y, z)`, `times`, `charge`, `auxiliary`, `core`, `mask`, and `rank`.\n\nIn order to process the input features, the model first creates embeddings for each feature using separate embedding modules. These embeddings are then concatenated and passed through a linear layer. An \"answer token\" is set at the first position of the input tensor, which will be used to predict the angles.","metadata":{}},{"cell_type":"code","source":"def layer_norm(x: jnp.ndarray) -> jnp.ndarray:\n    \"\"\"Applies a unique LayerNorm to x with default settings.\"\"\"\n    ln = hk.LayerNorm(axis=-1, create_scale=True, create_offset=True)\n    return ln(x)\n\n\n@dataclass\nclass TransformerEncoder(hk.Module):\n    num_heads: int\n    key_size: int\n    num_layers: int\n    feedforward_size: int = 2048\n    dropout_rate: float = 0.1\n\n    def __call__(\n        self,\n        inputs: jnp.ndarray,\n        mask: jnp.ndarray,\n        *,\n        is_training: bool = True,\n    ) -> jnp.ndarray:\n        initializer = hk.initializers.VarianceScaling(1.0, mode=\"fan_in\")\n        dropout_rate = self.dropout_rate if is_training else 0.0\n        _, _, model_size = inputs.shape\n\n        h = inputs\n        for _ in range(self.num_layers):\n            attn_block = hk.MultiHeadAttention(\n                num_heads=self.num_heads,\n                key_size=self.key_size,\n                model_size=model_size,\n                w_init=initializer,\n            )\n            h_norm = layer_norm(h)\n            h_attn = attn_block(\n                h_norm, h_norm, h_norm, mask=mask\n            )  # self-attention\n            h_attn = hk.dropout(hk.next_rng_key(), dropout_rate, h_attn)\n            h = h + h_attn\n\n            # then the dense block.\n            dense_block = hk.Sequential(\n                [\n                    hk.Linear(self.feedforward_size, w_init=initializer),\n                    jax.nn.gelu,\n                    hk.Linear(model_size, w_init=initializer),\n                ]\n            )\n            h_norm = layer_norm(h)\n            h_dense = dense_block(h_norm)\n            h_dense = hk.dropout(hk.next_rng_key(), dropout_rate, h_dense)\n            h = h + h_dense\n\n        return layer_norm(h)\n\n\nclass Rearrange(hk.Module):\n    def __init__(self, pattern, **axes_length):\n        super().__init__()\n        self._pattern = pattern\n        self._axes_length = axes_length\n\n    def __call__(self, x):\n        return rearrange(x, self._pattern, **self._axes_length)\n\n\ndef _one_to_embedding(embed_size):\n    \"\"\"Create module to change single continues values into `embeddings`\"\"\"\n    return hk.Sequential(\n        [\n            Rearrange(\"... -> ... ()\"),\n            hk.Linear(embed_size),\n            jax.nn.relu,\n            hk.Linear(embed_size),\n        ]\n    )\n\n\n@dataclass\nclass NaiveTransformer(hk.Module):\n    num_heads: int = 8\n    num_layers: int = 8\n    model_size: int = 512\n    key_size: int = 1024\n    feedforward_size: int = 2048\n    num_bins: int = 128\n    dropout_rate: float = 0.1\n    name: Optional[str] = None\n\n    def _feedforward(self, inputs):\n        h = layer_norm(inputs)\n        h = hk.Sequential(\n            [\n                hk.Linear(self.feedforward_size),\n                jax.nn.gelu,\n                hk.Linear(self.model_size),\n            ]\n        )(h)\n        h = h + inputs\n        return h\n\n    def __call__(\n        self,\n        sample: \"dict[str, jnp.ndarray]\",\n        *,\n        is_training: bool = True,\n    ) -> jnp.ndarray:  # [B, T, D]\n        \"\"\"\n        sample = {\n            \"x\" (float) -- [N, L]\n            \"y\" (float) -- [N, L]\n            \"z\" (float) -- [N, L]\n            \"auxiliary\" (long) -- [N, L]\n            \"times\" (float) -- [N, L]\n            \"charge\" (float) -- [N, L]\n            \"core\" (long) -- [N, L]\n            \"mask\" (bool) -- [N, L]\n            \"ans_mask\" (float) -- [N, L]\n            \"rank\" (long) -- [N, L]\n        }\n        \"\"\"\n\n        # for discrete value create Embedding layer\n        # for continuous values use stack of few linear layers\n        embedding_layers = {\n            \"auxiliary\": hk.Embed(2, self.model_size),\n            \"core\": hk.Embed(2, self.model_size),\n            \"rank\": hk.Embed(4, self.model_size),\n            **{\n                feat: _one_to_embedding(self.model_size)\n                for feat in [\"x\", \"y\", \"z\", \"time\", \"charge\"]\n            },\n        }\n\n        # project contatenated embeddings to model_size\n        embeds = hk.Linear(self.model_size)(\n            jnp.concatenate(\n                [\n                    module(sample[key])\n                    for key, module in embedding_layers.items()\n                ],\n                axis=-1,\n            )\n        )  # [N, L, M]\n\n        # replace first token with ans_token (by zeroing it first, then overwriting)\n        embeds = embeds * sample[\"ans_mask\"][..., None]  # [N, L, 1]\n        ans_token = hk.get_parameter(\n            \"ans_token\",\n            [self.model_size],\n            init=hk.initializers.UniformScaling(scale=0.1),\n        )\n        embeds = embeds.at[:, 0, :].set(ans_token)\n\n        # mask represents which values to hide when processing attention\n        # shape [batch, heads, out_seq, in_seq]\n        full_mask = sample[\"mask\"][:, None, None, :]  # [N, 1, 1, L]\n\n        # transformer with full attention\n        output = TransformerEncoder(\n            self.num_heads,\n            self.key_size,\n            self.num_layers,\n            self.feedforward_size,\n            self.dropout_rate,\n        )(embeds, full_mask, is_training=is_training)\n\n        # pick only ans token features, those are our embeddings for LSTM model\n        output = output[:, 0, :]\n\n        # apply few layers for azimuth prediction\n        azimuth = hk.Linear(self.num_bins)(self._feedforward(output))\n\n        # apply few layers for zenith prediction\n        zenith = hk.Linear(self.num_bins)(self._feedforward(output))\n        return azimuth, zenith, output\n","metadata":{"execution":{"iopub.status.busy":"2023-04-27T16:03:11.272762Z","iopub.execute_input":"2023-04-27T16:03:11.273342Z","iopub.status.idle":"2023-04-27T16:03:11.297150Z","shell.execute_reply.started":"2023-04-27T16:03:11.273305Z","shell.execute_reply":"2023-04-27T16:03:11.296120Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def transform_module(cls, with_state=True):\n    \"\"\"\n    Takes a class as an input and for each method applies\n    hk.transform, replacing it with values returned by hk.transform.\n    Created class have 'Haiku' prefix added to it's name.\n\n    Example:\n        class MyModule(hk.Module):\n            def __init__(self, out_features)\n                self.linear = hk.Linear(out_features)\n            def __call__(self, inputs):\n                return self.linear(inputs)\n            def predict(self, inputs):\n                return jnp.argmax(self.linear(inputs))\n\n        # in plain haiku\n        def _module_forward(inputs):\n            return MyModule(128)(inputs)\n        def _module_predict(inputs):\n            return MyModule(128).predict(inputs)\n        my_module_forward = hk.transform(_module_forward)\n        my_module_predict = hk.transform(_module_predict)\n        params = my_module_forward.init(jnp.zeros((32, 10)))\n        predictions = my_module_predict.apply(\n            params, None, jnp.zeros((32, 10)))\n        ...\n\n        # using this method\n        my_module = transform_module(MyModule)(128)\n        params = my_module.init(jnp.zeros(32, 10))\n        prediction = my_module.predict.apply(params, None, jnp.zeros((32, 10)))\n        ...\n\n    Assumptions:\n        - cls is subclass of hk.Module\n        - cls have __call__ method\n        - cls does NOT have init method\n        - cls does NOT have apply method\n    \"\"\"\n    assert issubclass(cls, hk.Module)\n    assert getattr(cls, \"__call__\", None) is not None\n    assert getattr(cls, \"init\", None) is None\n    assert getattr(cls, \"apply\", None) is None\n\n    def _wrap_fn(meth, *init_args, **init_kwargs):\n        @wraps(meth)\n        def _call_fn(*args, **kwargs):\n            obj = cls(*init_args, **init_kwargs)\n            return meth(obj, *args, **kwargs)\n\n        return _call_fn\n\n    hk_transform_fn = (\n        hk.transform if not with_state else hk.transform_with_state\n    )\n\n    @wraps(cls.__init__)\n    def _init_fn(self, *args, **kwargs):\n        self.init, self.apply = hk_transform_fn(\n            _wrap_fn(cls.__call__, *args, **kwargs)\n        )\n\n        excluded_methods = {\"__init__\", \"__call__\"}\n        extra_methods = set(dir(cls)) - set(dir(hk.Module)) - excluded_methods\n        methods = [\n            (method, getattr(cls, method))\n            for method in extra_methods\n            if callable(getattr(cls, method))\n        ]\n\n        for name, method in methods:\n            transformed = hk_transform_fn(_wrap_fn(method, *args, **kwargs))\n\n            jit_args = method.__dict__.get(\"jit_args\")\n            jit_kwargs = method.__dict__.get(\"jit_kwargs\")\n            if jit_args is not None and jit_kwargs is not None:\n                transformed = transformed._replace(\n                    apply=jax.jit(transformed.apply, *jit_args, **jit_kwargs)\n                )\n            setattr(self, name, transformed)\n\n    return type(f\"Haiku{cls.__name__}\", (), {\"__init__\": _init_fn})","metadata":{"_kg_hide-input":false,"execution":{"iopub.status.busy":"2023-04-27T16:03:11.298651Z","iopub.execute_input":"2023-04-27T16:03:11.299222Z","iopub.status.idle":"2023-04-27T16:03:11.312538Z","shell.execute_reply.started":"2023-04-27T16:03:11.299176Z","shell.execute_reply":"2023-04-27T16:03:11.311579Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def create_transformer_dataloader(config, num_devices):\n    \"\"\"Dataloader for Transformer and LSTM\"\"\"\n    transform = PositionalFeatures(\n        config.dataset / \"sensor_geometry.csv\",\n        length=256,\n        num_bins=config.num_bins,\n        training=(config.subset == \"train\"),\n    )\n\n    dataset = IceCubeNeutrinos(\n        root_path=config.dataset,\n        subset=config.subset,\n        transform=transform,\n        shuffle_files=False,\n    )\n\n    dataloader = data.DataLoader(\n        dataset,\n        config.batch_size,\n        num_workers=config.num_workers,\n        drop_last=(config.subset == \"train\"),\n        collate_fn=partial(jax_collate_repeat, num_devices),\n    )\n\n    return dataloader\n\n\ndef _prefetch_jax(devices, xs):\n    return jax.tree_map(\n        lambda x: jax.device_put_sharded(list(x.numpy()), devices), xs\n    )\n\n\ndef load_transformer_model(config, devices):\n    if config.checkpoint is None:\n        assert False, \"Missing Transformer checkpoint\"\n\n    print(f'> restoring model from checkpoint - \"{config.checkpoint}\"...')\n    with open(config.checkpoint, \"rb\") as file:\n        checkpoint = pickle.load(file)\n        assert isinstance(\n            checkpoint, dict\n        ), \"checkpoint should contain a dictionary\"\n\n    model_params = dict(\n        num_heads=8,\n        num_layers=8,\n        model_size=512,\n        key_size=128,\n        feedforward_size=2048,\n        dropout_rate=0.0,\n        num_bins=config.num_bins,\n    )\n    ckpt_model_params = checkpoint[\"model_params\"]\n    diff_keys = ckpt_model_params.keys() ^ model_params.keys()\n    if len(diff_keys) > 0:\n        print(f\"!> Extra parameter in model_params! = {diff_keys}\")\n    model_params.update(ckpt_model_params)\n    model = transform_module(NaiveTransformer)(**model_params)\n\n    if \"params\" in checkpoint and \"states\" in checkpoint:\n        params = replicate(checkpoint[\"params\"], devices)\n        states = replicate(checkpoint[\"states\"], devices)\n    else:\n        assert False, \"Missing Transformer parameters\"\n\n    num_params = jax.tree_util.tree_reduce(\n        lambda a, x: a + x[0].size, params, initializer=0\n    )\n    print(f\"> model parameters = {num_params}\")\n\n    return model, params, states\n\n\ndef create_evaluate_function(config, devices, model, params, states):\n    \"\"\"\n    Create a function that runs a Transformer model on a given data\n    and returns predictions as well as embeddings.\n    \"\"\"\n    _az_enc = AzimuthZenithEncoder(config.num_bins, config.num_bins)\n    az_mapping = replicate(_az_enc._azimuth.numpy(), devices)\n    zen_mapping = replicate(_az_enc._zenith.numpy(), devices)\n\n    @partial(jax.pmap, axis_name=\"device\", devices=devices)\n    def evaluate_step(az_map, zen_map, params, states, sample):\n        (az_logits, zen_logits, embeddings), _ = model.apply(\n            params, states, jax.random.PRNGKey(0), sample, is_training=False\n        )\n\n        if config.infer_method == \"argmax\":\n            azimuth = az_map[jnp.argmax(az_logits, axis=-1)]\n            zenith = zen_map[jnp.argmax(zen_logits, axis=-1)]\n        elif config.infer_method == \"sum\":\n            azimuth = jnp.sum(az_map * jax.nn.softmax(az_logits), axis=-1)\n            zenith = jnp.sum(zen_map * jax.nn.softmax(zen_logits), axis=-1)\n        else:\n            raise RuntimeError(\"unknown method\")\n\n        return azimuth, zenith, embeddings\n\n    return partial(evaluate_step, az_mapping, zen_mapping, params, states)\n","metadata":{"execution":{"iopub.status.busy":"2023-04-27T16:03:11.317514Z","iopub.execute_input":"2023-04-27T16:03:11.318052Z","iopub.status.idle":"2023-04-27T16:03:11.336465Z","shell.execute_reply.started":"2023-04-27T16:03:11.318024Z","shell.execute_reply":"2023-04-27T16:03:11.335380Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# GraphNet\nChanges made:\n- added ice trainsparency\n- added features to preprocessing in `graphnet_transform`: `wrong_charge_dom`, `quantum_effeciency`(`qe`)\n- added `features_subset_rest` argument to DynEdge model, changing `features_subset` for non-first layer\n- added two models IceCubeModelv4(bigger one, predicting azimuth and zenith) and IceCubeModelDir(larger post processing layer and predicting direction)\n- added option for models to train with other optimizers, used: AdaBelief\n- added CyclicLR but not used in training \n- added forward function `fwd_emb` returning graph embeddings from last layer ","metadata":{}},{"cell_type":"code","source":"def ice_transparency(data_path=\"data/nvme_1t/ice_transparency.csv\", datum=1950):\n    # Data from page 31 of https://arxiv.org/pdf/1301.5361.pdf\n    # Datum is from footnote 8 of page 29\n    df = pd.read_csv(data_path, delim_whitespace=True)\n    df[\"z\"] = df[\"depth\"] - datum\n    df[\"z_norm\"] = df[\"z\"] / 500\n    df[\n        [\"scattering_len_norm\", \"absorption_len_norm\"]\n    ] = RobustScaler().fit_transform(df[[\"scattering_len\", \"absorption_len\"]])\n\n    # These are both roughly equivalent after scaling\n    f_scattering = interp1d(df[\"z_norm\"], df[\"scattering_len_norm\"])\n    f_absorption = interp1d(df[\"z_norm\"], df[\"absorption_len_norm\"])\n    return f_scattering, f_absorption\n\n\ndef graphnet_transform(\n    event, geo_t, f_scattering, f_absorption, pulse_limit=500\n):\n    event_id = torch.tensor(event.event_id, dtype=torch.int64)\n    time = (event.time - 1.0e04) / 3.0e4\n    charge = torch.log10(event.charge) / 3.0\n    auxiliary = event.auxiliary.to(int) - 0.5\n    x_t = torch.take(geo_t.T[0, :], event.sensor_id.to(torch.int64)) / 500\n    y_t = torch.take(geo_t.T[1, :], event.sensor_id.to(torch.int64)) / 500\n    z_t = torch.take(geo_t.T[2, :], event.sensor_id.to(torch.int64)) / 500\n\n    # quantum effeciency in deep core\n    qe = torch.take(geo_t.T[3, :], event.sensor_id.to(torch.int64))\n\n    # due to dataset_stats.ipynb: \"Average charge per sensors\"\n    wrong_charge_dom = torch.zeros((geo_t.shape[0]))\n    wrong_charge_dom[1229] = 1.0\n    wrong_charge_dom_t = torch.take(\n        wrong_charge_dom, event.sensor_id.to(torch.int64)\n    )\n    x = torch.stack(\n        [x_t, y_t, z_t, time, charge, auxiliary, qe, wrong_charge_dom_t], dim=1\n    )\n    y = None\n    if event.azimuth is not None or event.zenith is not None:\n        y = torch.tensor([event.azimuth, event.zenith], dtype=torch.float32)\n    data = Data(x=x, y=y, n_pulses=torch.tensor(x.shape[0], dtype=torch.int32))\n    data.x = data.x.to(torch.float32)\n\n    z = data.x[:, 2].numpy()\n    scattering = torch.tensor(f_scattering(z), dtype=torch.float32).view(-1, 1)\n    absorption = torch.tensor(f_absorption(z), dtype=torch.float32).view(-1, 1)\n\n    data.x = torch.cat([data.x, scattering, absorption], dim=1)\n\n    # Downsample the large events\n    if data.n_pulses > pulse_limit:\n        data.x = data.x[np.random.choice(data.n_pulses, pulse_limit)]\n        data.n_pulses = torch.tensor(pulse_limit, dtype=torch.int32)\n\n    return {\"data\": data, \"event_ids\": event_id}\n\n\ndef angular_dist_score(y_pred, y_true, eps=1e-7):\n    \"\"\"\n    calculate the MAE of the angular distance between two directions.\n    The two vectors are first converted to cartesian unit vectors,\n    and then their scalar product is computed, which is equal to\n    the cosine of the angle between the two vectors. The inverse\n    cosine (arccos) thereof is then the angle between the two itorchut vectors\n\n    # https://www.kaggle.com/code/sohier/mean-angular-error\n\n    Parameters:\n    -----------\n\n    y_pred : float (torch.Tensor)\n        Prediction array of [N, 2], where the second dim is azimuth & zenith\n    y_true : float (torch.Tensor)\n        Ground truth array of [N, 2], where the second dim is azimuth & zenith\n\n    Returns:\n    --------\n\n    dist : float (torch.Tensor)\n        mean over the angular distance(s) in radian\n    \"\"\"\n\n    az_true = y_true[:, 0]\n    zen_true = y_true[:, 1]\n\n    az_pred = y_pred[:, 0]\n    zen_pred = y_pred[:, 1]\n\n    # pre-compute all sine and cosine values\n    sa1 = torch.sin(az_true)\n    ca1 = torch.cos(az_true)\n    sz1 = torch.sin(zen_true)\n    cz1 = torch.cos(zen_true)\n\n    sa2 = torch.sin(az_pred)\n    ca2 = torch.cos(az_pred)\n    sz2 = torch.sin(zen_pred)\n    cz2 = torch.cos(zen_pred)\n\n    # scalar product of the two cartesian vectors (x = sz*ca, y = sz*sa, z = cz)\n    scalar_prod = sz1 * sz2 * (ca1 * ca2 + sa1 * sa2) + (cz1 * cz2)\n\n    # scalar product of two unit vectors is always between -1 and 1, this is against nummerical instability\n    # that might otherwise occure from the finite precision of the sine and cosine functions\n    scalar_prod = torch.clamp(scalar_prod, -1 + eps, 1 - eps)\n\n    # convert back to an angle (in radian)\n    return torch.mean(torch.abs(torch.arccos(scalar_prod)))\n\n\ndef add_weight_decay(\n    model,\n    weight_decay=1e-5,\n    skip_list=(\"bias\", \"bn\", \"LayerNorm.bias\", \"LayerNorm.weight\"),\n):\n    decay = []\n    no_decay = []\n    for name, param in model.named_parameters():\n        if not param.requires_grad:\n            continue  # frozen weights\n        if len(param.shape) == 1 or name.endswith(\".bias\") or name in skip_list:\n            no_decay.append(param)\n        else:\n            decay.append(param)\n    return [\n        {\"params\": no_decay, \"weight_decay\": 0.0},\n        {\"params\": decay, \"weight_decay\": weight_decay},\n    ]\n\n\ndef get_sensor_geo_df(sensor_geo_path, with_qe=True):\n    df = pd.read_csv(sensor_geo_path, index_col=\"sensor_id\")\n    if with_qe:\n        df[\"qe\"] = -1\n        for i in range(len(df) // 60):\n            # High Quantum Efficiency in the lower 50 DOMs - https://arxiv.org/pdf/2209.03042.pdf (Figure 1)\n            if i in range(78, 86):\n                start_veto, end_veto = i * 60, (i * 60) + 10\n                start_core, end_core = end_veto + 1, (i * 60) + 60\n                df.loc[start_core:end_core, \"qe\"] = 0.4\n    return df\n\n\nGLOBAL_POOLINGS = {\n    \"min\": scatter_min,\n    \"max\": scatter_max,\n    \"sum\": scatter_sum,\n    \"mean\": scatter_mean,\n}\n\n\nclass DynEdgeConv(EdgeConv, pl.LightningModule):\n    \"\"\"Dynamical edge convolution layer.\"\"\"\n\n    def __init__(\n        self,\n        nn: Callable,\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        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        # Base class constructor\n        super().__init__(nn=nn, aggr=aggr, **kwargs)\n\n        # Additional member variables\n        self.nb_neighbors = nb_neighbors\n        self.features_subset = features_subset\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        x = super().forward(x, edge_index)\n\n        # Recompute adjacency\n        edge_index = knn_graph(\n            x=x[:, self.features_subset],\n            k=self.nb_neighbors,\n            batch=batch,\n        ).to(self.device)\n\n        return x, edge_index\n\n\nclass DynEdgeConv(EdgeConv, pl.LightningModule):\n    \"\"\"Dynamical edge convolution layer.\"\"\"\n\n    def __init__(\n        self,\n        nn: Callable,\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        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        # Base class constructor\n        super().__init__(nn=nn, aggr=aggr, **kwargs)\n\n        # Additional member variables\n        self.nb_neighbors = nb_neighbors\n        self.features_subset = features_subset\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        x = super().forward(x, edge_index)\n\n        # Recompute adjacency\n        edge_index = knn_graph(\n            x=x[:, self.features_subset],\n            k=self.nb_neighbors,\n            batch=batch,\n        ).to(self.device)\n\n        return x, edge_index\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        features_subset_rest: int = None,\n    ):\n        \"\"\"Construct `DynEdge`.\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, 3)\n\n        if features_subset_rest is None:\n            features_subset_rest = features_subset\n        features_subset_rest = slice(0, features_subset_rest)\n\n        # DynEdge layer sizes\n        if dynedge_layer_sizes is None:\n            dynedge_layer_sizes = [\n                (128, 256),\n                (336, 256),\n                (336, 256),\n                (336, 256),\n            ]\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                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        # 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.GELU()\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        self._features_subset_rest = features_subset_rest\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 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 i, sizes in enumerate(self._dynedge_layer_sizes):\n            layers = []\n            layer_sizes = [nb_latent_features] + list(sizes)\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 *= 2\n                layers.append(torch.nn.Linear(nb_in, nb_out))\n                layers.append(nn.LayerNorm(nb_out))\n                layers.append(self._activation)\n\n            feature_subset = (\n                self._features_subset if i == 0 else self._features_subset_rest\n            )\n            conv_layer = DynEdgeConv(\n                torch.nn.Sequential(*layers),\n                aggr=\"add\",\n                nb_neighbors=self._nb_neighbours,\n                features_subset=feature_subset,\n            )\n            self._conv_layers.append(conv_layer)\n\n            nb_latent_features = nb_out\n\n        # Post-processing operations\n        nb_latent_features = (\n            sum(sizes[-1] for sizes in self._dynedge_layer_sizes)\n            + nb_input_features\n        )\n\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(nn.LayerNorm(nb_out))\n            post_processing_layers.append(self._activation)\n\n        self._post_processing = torch.nn.Sequential(*post_processing_layers)\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 = nb_out * nb_poolings\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(nn.LayerNorm(nb_out))\n            readout_layers.append(self._activation)\n\n        self._readout = torch.nn.Sequential(*readout_layers)\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        global_variables = self._calculate_global_variables(\n            x,\n            edge_index,\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) * global_variables.unsqueeze(dim=0),\n                dim=1,\n            )\n\n            x = torch.cat((x, global_variables_distributed), dim=1)\n\n        # DynEdge-convolutions\n        skip_connections = [x]\n        for conv_layer in self._conv_layers:\n            x, edge_index = conv_layer(x, edge_index, batch)\n            skip_connections.append(x)\n\n        # Skip-cat\n        x = torch.cat(skip_connections, dim=1)\n\n        # Post-processing\n        x = self._post_processing(x)\n\n        # (Optional) Global pooling\n        if self._global_pooling_schemes:\n            x = self._global_pooling(x, batch=batch)\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        # Read-out\n        x = self._readout(x)\n\n        return x\n\n\nclass IceCubeModelDir(pl.LightningModule):\n    def __init__(\n        self,\n        model_name: str = \"DynEdge-Dir\",\n        learning_rate: float = 0.001,\n        weight_decay: float = 0.01,\n        warmup: float = 0.1,\n        T_max: int = 1000,\n        nb_inputs: int = 10,\n        nearest_neighbours: int = 8,\n        enable_cosine_loss: bool = False,\n        cyclic_scheduler: bool = False,\n        optimizer: str = \"adamw\",\n        features_subset_rest: int = 4,\n        **kwargs,\n    ):\n        super().__init__()\n        self.save_hyperparameters()\n\n        self.loss_dir = VonMisesFisher3DLoss()\n        self.loss_fn_cos = angular_dist_score\n\n        self.model = DynEdge(\n            nb_inputs=nb_inputs,\n            nb_neighbours=nearest_neighbours,\n            global_pooling_schemes=[\"min\", \"max\", \"mean\", \"sum\"],\n            features_subset=slice(0, 4),  # NN search using xyzt\n            dynedge_layer_sizes=[\n                (128, 256),\n                (336, 256),\n                (336, 256),\n                (336, 256),\n            ],\n            readout_layer_sizes=[\n                READOUT_FIRST_LAYER_SIZE,\n            ],\n            post_processing_layer_sizes=[512, 336],\n            features_subset_rest=features_subset_rest,\n        )\n        self.direction_task = DirectionReconstructionWithKappa(\n            hidden_size=self.model.nb_outputs,\n            loss_function=self.loss_dir,\n            target_labels=[\"dir_x\", \"dir_y\", \"dir_z\", \"kappa\"],\n        )\n        self.norm = nn.LayerNorm(self.model.nb_outputs)\n\n    def _to_direction(self, azimuth, zenith):\n        dx = torch.sin(zenith) * torch.cos(azimuth)\n        dy = torch.sin(zenith) * torch.sin(azimuth)\n        dz = torch.cos(zenith)\n        return torch.stack((dx, dy, dz), dim=-1)\n\n    def _to_angles(self, dir_out, ret_tuple=False):\n        x, y, z = dir_out[:, :3].unbind(dim=-1)\n        zenith = torch.arccos(z.clamp(-1.0, 1.0))\n        azimuth = torch.arctan2(y, x) + 2.0 * torch.pi * (y < 0.0)\n        if ret_tuple:\n            return azimuth, zenith\n        return torch.stack((azimuth, zenith), dim=-1)\n\n    def forward(self, x):\n        emb = self.model(x)\n        emb = self.norm(emb)\n        dir_out = self.direction_task(emb)\n        return dir_out\n\n    def fwd_embed(self, x):\n        emb = self.model(x)\n        emb = self.norm(emb)\n        return self.direction_task._affine(emb)\n\n    def training_step(self, batch, batch_idx):\n        dir_out = self.forward(batch)\n\n        target = batch.y.reshape(-1, 2)\n        target_dir = self._to_direction(target[:, 0], target[:, 1])\n        loss = self.loss_dir(dir_out, target_dir)\n\n        metric = angular_dist_score(self._to_angles(dir_out), target)\n\n        self.log_dict(\n            {\n                \"train_loss\": loss,\n                \"train_mae\": metric,\n            }\n        )\n        return {\n            \"loss\": loss,\n            \"mae\": metric,\n        }\n\n    def training_epoch_end(self, training_step_outputs):\n        avg_loss = torch.stack(\n            [x[\"loss\"] for x in training_step_outputs]\n        ).mean()\n        self.log(\"train_total_loss\", avg_loss)\n\n    def validation_step(self, batch, batch_idx):\n        dir_out = self.forward(batch)\n\n        target = batch.y.reshape(-1, 2)\n        target_dir = self._to_direction(target[:, 0], target[:, 1])\n\n        loss = self.loss_dir(dir_out, target_dir)\n\n        metric = angular_dist_score(self._to_angles(dir_out), target)\n\n        output = {\n            \"val_loss\": loss,\n            \"val_mae\": metric,\n        }\n\n        return output\n\n    def validation_epoch_end(self, outputs):\n        loss_val = torch.stack([x[\"val_loss\"] for x in outputs]).mean()\n        metric = torch.stack([x[\"val_mae\"] for x in outputs]).mean()\n\n        self.log_dict(\n            {\"val_loss\": loss_val, \"val_mae\": metric},\n            prog_bar=True,\n        )\n\n    def configure_optimizers(self):\n        parameters = add_weight_decay(\n            self,\n            self.hparams.weight_decay,\n            skip_list=[\"bias\", \"LayerNorm.bias\"],\n        )\n        if self.hparams.optimizer == \"adamw\":\n            opt = torch.optim.AdamW(parameters, lr=self.hparams.learning_rate)\n        elif self.hparams.optimizer == \"adam\":\n            opt = torch.optim.Adam(parameters, lr=self.hparams.learning_rate)\n        elif self.hparams.optimizer == \"adabelief\":\n            opt = AdaBelief(parameters, lr=self.hparams.learning_rate)\n        elif self.hparams.optimizer == \"sgd\":\n            opt = torch.optim.SGD(parameters, lr=self.hparams.learning_rate)\n        else:\n            raise NotImplementedError(\n                f\"Optimizer: {self.hparams.optimizer} not implemented\"\n            )\n        ret = {\"optimizer\": opt}\n\n        if self.hparams.cyclic_scheduler:\n            sch = torch.optim.lr_scheduler.CyclicLR(\n                optimizer=opt,\n                base_lr=self.hparams.learning_rate,\n                max_lr=0.1,\n                cycle_momentum=False,\n            )\n            ret[\"lr_scheduler\"] = {\n                \"scheduler\": sch,\n                \"interval\": \"step\",\n                \"monitor\": \"train_mae\",\n            }\n\n        return ret\n\n\ndef get_full_transform(nb_nearest_neighbours=8, ret_event_ids=False):\n    geo_df = get_sensor_geo_df(GEOMETRY_PATH)\n    geo_t = torch.tensor(geo_df.to_numpy(), dtype=torch.float32)\n    f_scattering, f_absorption = ice_transparency(TRANSPARENCY_PATH)\n    post_transform = KNNGraphBuilder(\n        nb_nearest_neighbours=nb_nearest_neighbours\n    )\n\n    def transform(event):\n        d = graphnet_transform(event, geo_t, f_scattering, f_absorption)\n        data = post_transform(d[\"data\"])\n        if ret_event_ids:\n            return d[\"event_ids\"], data\n        else:\n            return data\n\n    return transform\n\n\ndef create_dataloader(\n    dataset_path, nb_nearest_neighbours=8, batch_size=512, num_workers=0\n):\n    return DataLoader(\n        IceCubeNeutrinos(\n            root_path=dataset_path,\n            subset=\"test\",\n            transform=get_full_transform(\n                nb_nearest_neighbours=nb_nearest_neighbours, ret_event_ids=True\n            ),\n        ),\n        batch_size=batch_size,\n        num_workers=num_workers,\n    )","metadata":{"execution":{"iopub.status.busy":"2023-04-27T16:03:11.338417Z","iopub.execute_input":"2023-04-27T16:03:11.338754Z","iopub.status.idle":"2023-04-27T16:03:11.423198Z","shell.execute_reply.started":"2023-04-27T16:03:11.338716Z","shell.execute_reply":"2023-04-27T16:03:11.421996Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# LSTM Ensembler\n\nThis code defines an LSTMEnsembler class and consists of an 3 Bi-directional LSTM layers and a decision layer. The code also includes a function `prepare_embeddings` to concatenate azimuth, zenith, and embeddings data.\n\nThis model ensembles information from multiple sources, specifically Transformer embeddings and GraphNet embeddings. By concatenating these embeddings, the LSTMEnsembler aims to leverage the strengths of both the transformer and the graph network, potentially leading to more accurate and robust predictions.\n\nFor each sample in the test dataset, we processes the data using both transformer and graph network models, obtaining their respective azimuth, zenith, and embeddings. The LSTMEnsembler then takes the concatenated embeddings and makes a decision (using sigmoid activation) whether to choose the Transformer output (True) or the GraphNet output (False) for each sample. The selected azimuth and zenith values are saved as the final output.","metadata":{}},{"cell_type":"code","source":"class LSTMEnsembler(nn.Module):\n    def __init__(self, hidden_size=512 + 256, num_layers=2):\n        super().__init__()\n        self.num_layers = num_layers\n        self.lstm = nn.LSTM(\n            7,\n            hidden_size,\n            num_layers=num_layers,\n            bidirectional=True,\n            batch_first=True,\n        )\n        self.decision = nn.Sequential(\n            nn.LayerNorm(hidden_size * 2),\n            nn.GELU(),\n            nn.Linear(hidden_size * 2, 1),\n        )\n\n    def forward(self, sample, transformer_embeds, graphnet_embeds):\n        padded_inputs = torch.stack(\n            [\n                # remove \"ans\" token from input data (because we are shring\n                # dataloading with Transformer).\n                sample[\"x\"][..., 1:],\n                sample[\"y\"][..., 1:],\n                sample[\"z\"][..., 1:],\n                sample[\"charge\"][..., 1:],\n                sample[\"time\"][..., 1:],\n                sample[\"time_diff\"][..., 1:],\n                sample[\"auxiliary\"][..., 1:].float() * 2.0 - 1.0,\n            ],\n            dim=-1,\n        )  # [N, L, *]\n        lengths = sample[\"mask\"][..., 1:].cpu().sum(dim=-1)\n        packed_inputs = nn.utils.rnn.pack_padded_sequence(\n            padded_inputs,\n            lengths=lengths,\n            batch_first=True,\n            enforce_sorted=False,\n        )\n        cell_state = repeat(\n            torch.cat(\n                [transformer_embeds, graphnet_embeds], dim=-1\n            ),  # [N, 512 + 256 + 4]\n            \"n h -> d n h\",\n            d=2 * self.num_layers,\n        )\n        hidden_states = (cell_state, cell_state)\n        packed_outputs, _ = self.lstm(packed_inputs, hidden_states)\n        padded_outputs, output_lengths = nn.utils.rnn.pad_packed_sequence(\n            packed_outputs,\n            batch_first=True,\n        )\n        mask = (\n            torch.arange(\n                output_lengths.max(), device=output_lengths.device\n            )[None, :] < output_lengths[:, None]\n        ).to(padded_outputs.device)\n\n        decision = self.decision(padded_outputs)[..., 0]  # [N, L]\n        return torch.sum(decision * mask.float(), dim=-1) / mask.sum(dim=-1)\n        # return torch.masked.mean(decision, dim=-1, mask=mask)  # [N]\n\n\ndef prepare_embeddings(azimuth, zenith, embeddings):\n    return torch.from_numpy(\n        np.concatenate(\n            (\n                azimuth[..., None],\n                zenith[..., None],\n                embeddings,\n            ),\n            axis=-1,\n        )\n    )\n\n\ndef load_lstm_ensembler(path):\n    checkpoint = torch.load(path, map_location=\"cpu\")\n    model_params = checkpoint[\"model_params\"]\n    lstm_ensembler = LSTMEnsembler(**model_params)\n    lstm_ensembler.load_state_dict(checkpoint[\"state_dict\"])\n    return lstm_ensembler","metadata":{"execution":{"iopub.status.busy":"2023-04-27T16:03:11.426182Z","iopub.execute_input":"2023-04-27T16:03:11.427034Z","iopub.status.idle":"2023-04-27T16:03:11.441935Z","shell.execute_reply.started":"2023-04-27T16:03:11.426992Z","shell.execute_reply":"2023-04-27T16:03:11.440912Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transformer_device = jax.local_devices()[0]\ngraphnet_device = torch.device(\"cuda:0\")\nensembler_device = torch.device(\"cuda:0\")\ntransformer_device, graphnet_device, ensembler_device","metadata":{"execution":{"iopub.status.busy":"2023-04-27T16:03:11.443613Z","iopub.execute_input":"2023-04-27T16:03:11.445103Z","iopub.status.idle":"2023-04-27T16:03:11.671936Z","shell.execute_reply.started":"2023-04-27T16:03:11.445063Z","shell.execute_reply":"2023-04-27T16:03:11.670999Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create models\nwith jax.default_device(transformer_device):\n    transformer_model, params, states = load_transformer_model(\n        config, [transformer_device]\n    )\n    transformer_model = create_evaluate_function(\n        config, [transformer_device], transformer_model, params, states\n    )\n\ngraphnet_model = IceCubeModelDir(\n    features_subset_rest=FEATURES_SUBSET_REST\n).load_from_checkpoint(CHPT_PATH)\ngraphnet_model.to(graphnet_device)\n\nlstm_ensembler = (\n    load_lstm_ensembler(\n        \"/kaggle/input/icecube-5th-place-models/lstm_model.ckpt\"\n    )\n    .eval()\n    .to(ensembler_device)\n)","metadata":{"execution":{"iopub.status.busy":"2023-04-27T16:03:11.673621Z","iopub.execute_input":"2023-04-27T16:03:11.673977Z","iopub.status.idle":"2023-04-27T16:03:27.829166Z","shell.execute_reply.started":"2023-04-27T16:03:11.673921Z","shell.execute_reply":"2023-04-27T16:03:27.828135Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create dataloaders\ntransformer_dataloader = create_transformer_dataloader(config, num_devices=1)\ngraphnet_dataloader = create_dataloader(\n    DATASET_PATH, batch_size=config.batch_size, num_workers=config.num_workers\n)","metadata":{"execution":{"iopub.status.busy":"2023-04-27T16:03:27.830692Z","iopub.execute_input":"2023-04-27T16:03:27.831081Z","iopub.status.idle":"2023-04-27T16:03:27.952594Z","shell.execute_reply.started":"2023-04-27T16:03:27.831040Z","shell.execute_reply":"2023-04-27T16:03:27.949172Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"progress_bar = tqdm(\n    zip(transformer_dataloader, graphnet_dataloader),\n    mininterval=20.0,\n    maxinterval=100.0,\n)\n\n\nresults = []\n\nfor transformer_sample, graphnet_sample in progress_bar:\n    event_id, graphnet_sample = graphnet_sample\n    ensembler_sample = {\n        k: v[0].to(ensembler_device) for k, v in transformer_sample.items()\n    }\n    transformer_sample = _prefetch_jax([transformer_device], transformer_sample)\n    graphnet_sample = graphnet_sample.to(graphnet_device)\n\n    tr_azimuth, tr_zenith, tr_embed = [\n        np.array(jax.device_get(res)[0])\n        for res in transformer_model(transformer_sample)\n    ]\n\n    with torch.no_grad():\n        gn_embed = graphnet_model.norm(graphnet_model.model(graphnet_sample))\n        gn_azm, gn_zen = graphnet_model._to_angles(\n            graphnet_model.direction_task(gn_embed), ret_tuple=True\n        )\n        gn_embed, gn_azimuth, gn_zenith = [\n            res.cpu().numpy() for res in [gn_embed, gn_azm, gn_zen]\n        ]\n\n        # True -> transformer; False -> graphnet\n        model_picked = (\n            torch.sigmoid(\n                lstm_ensembler(\n                    ensembler_sample,\n                    prepare_embeddings(\n                        tr_azimuth, tr_zenith, tr_embed\n                    ).to(ensembler_device),\n                    prepare_embeddings(\n                        gn_azimuth, gn_zenith, gn_embed\n                    ).to(ensembler_device),\n                )\n            ).cpu().numpy() > 0.5\n        )\n\n    azimuth = np.where(model_picked, tr_azimuth, gn_azimuth)\n    zenith = np.where(model_picked, tr_zenith, gn_zenith)\n\n    results.extend(\n        list(\n            zip(\n                event_id.numpy().tolist(),\n                azimuth.tolist(),\n                zenith.tolist(),\n            )\n        )\n    )\n\nsubmission = pd.DataFrame.from_records(\n    results, columns=[\"event_id\", \"azimuth\", \"zenith\"]\n)\nsubmission.to_csv(\"submission.csv\", index=False)\nsubmission.head()","metadata":{"execution":{"iopub.status.busy":"2023-04-27T16:03:27.958447Z","iopub.execute_input":"2023-04-27T16:03:27.959138Z","iopub.status.idle":"2023-04-27T16:03:33.625579Z","shell.execute_reply.started":"2023-04-27T16:03:27.959082Z","shell.execute_reply":"2023-04-27T16:03:33.624413Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!head -10 submission.csv","metadata":{"execution":{"iopub.status.busy":"2023-04-27T16:03:33.627656Z","iopub.execute_input":"2023-04-27T16:03:33.628076Z","iopub.status.idle":"2023-04-27T16:03:34.640351Z","shell.execute_reply.started":"2023-04-27T16:03:33.628036Z","shell.execute_reply":"2023-04-27T16:03:34.639087Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from IPython.display import FileLink\nFileLink('submission.csv')","metadata":{"execution":{"iopub.status.busy":"2023-04-27T16:03:34.642629Z","iopub.execute_input":"2023-04-27T16:03:34.643104Z","iopub.status.idle":"2023-04-27T16:03:34.651475Z","shell.execute_reply.started":"2023-04-27T16:03:34.643052Z","shell.execute_reply":"2023-04-27T16:03:34.650205Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}