{"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":"# ICECUBE","metadata":{"papermill":{"duration":0.008138,"end_time":"2023-01-23T10:33:19.704397","exception":false,"start_time":"2023-01-23T10:33:19.696259","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# import torch\n# print(torch.__version__)  # -> 1.13.0\n# print(torch.version.cuda) # -> 11.3","metadata":{"execution":{"iopub.status.busy":"2023-04-18T20:39:46.780358Z","iopub.execute_input":"2023-04-18T20:39:46.780921Z","iopub.status.idle":"2023-04-18T20:39:46.808152Z","shell.execute_reply.started":"2023-04-18T20:39:46.780822Z","shell.execute_reply":"2023-04-18T20:39:46.807073Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# !pip install /kaggle/input/icm-setup/torch-1.12.0+cu113-cp37-cp37m-linux_x86_64.whl\n# !pip install /kaggle/input/icm-setup/torchvision-0.13.0+cu113-cp37-cp37m-linux_x86_64.whl\n!pip install /kaggle/input/icm-setup/torch-1.12.1+cu113-cp37-cp37m-linux_x86_64.whl\n!pip install /kaggle/input/icm-setup/torchvision-0.13.1+cu113-cp37-cp37m-linux_x86_64.whl\n!pip install /kaggle/input/icm-setup/pyg_lib-0.1.0+pt112cu113-cp37-cp37m-linux_x86_64.whl\n!pip install /kaggle/input/icm-setup/torch_scatter-2.1.0+pt112cu113-cp37-cp37m-linux_x86_64.whl\n!pip install /kaggle/input/icm-setup/torch_sparse-0.6.16+pt112cu113-cp37-cp37m-linux_x86_64.whl\n!pip install /kaggle/input/icm-setup/torch_cluster-1.6.0+pt112cu113-cp37-cp37m-linux_x86_64.whl\n!pip install /kaggle/input/icm-setup/torch_spline_conv-1.2.1+pt112cu113-cp37-cp37m-linux_x86_64.whl\n!pip install /kaggle/input/icm-setup/torch_geometric-2.2.0.tar.gz","metadata":{"execution":{"iopub.status.busy":"2023-04-18T20:39:46.810315Z","iopub.execute_input":"2023-04-18T20:39:46.811065Z","iopub.status.idle":"2023-04-18T20:45:15.057052Z","shell.execute_reply.started":"2023-04-18T20:39:46.811024Z","shell.execute_reply":"2023-04-18T20:45:15.055834Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport gc\nimport sys\nimport scipy\nimport numpy as np\nimport pandas as pd\nimport pyarrow\nimport random","metadata":{"execution":{"iopub.status.busy":"2023-04-18T20:45:15.060955Z","iopub.execute_input":"2023-04-18T20:45:15.061284Z","iopub.status.idle":"2023-04-18T20:45:15.086986Z","shell.execute_reply.started":"2023-04-18T20:45:15.061242Z","shell.execute_reply":"2023-04-18T20:45:15.086080Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch_geometric\nfrom torch.cuda.amp import autocast\nfrom torch.cuda.amp.grad_scaler import GradScaler\nfrom torch_geometric.data import Data, Batch, Dataset\nfrom torch_geometric.loader import DataLoader\nfrom torch_geometric.nn import EdgeConv, knn_graph\nfrom torch_scatter import scatter_max, scatter_mean, scatter_min, scatter_sum\nfrom torch_geometric.typing import Adj\nfrom timm.scheduler import CosineLRScheduler","metadata":{"execution":{"iopub.status.busy":"2023-04-18T20:45:15.089964Z","iopub.execute_input":"2023-04-18T20:45:15.090341Z","iopub.status.idle":"2023-04-18T20:45:21.207581Z","shell.execute_reply.started":"2023-04-18T20:45:15.090290Z","shell.execute_reply":"2023-04-18T20:45:21.206523Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from typing import Optional\nfrom torch import Tensor\nfrom torch.nn import LayerNorm, Linear, MultiheadAttention, Parameter\nfrom torch_geometric.nn.aggr import Aggregation\n\nclass MultiheadAttentionBlock(torch.nn.Module):\n    r\"\"\"The Multihead Attention Block (MAB) from the `\"Set Transformer: A\n    Framework for Attention-based Permutation-Invariant Neural Networks\"\n    <https://arxiv.org/abs/1810.00825>`_ paper\n    .. math::\n        \\mathrm{MAB}(\\mathbf{x}, \\mathbf{y}) &= \\mathrm{LayerNorm}(\\mathbf{h} +\n        \\mathbf{W} \\mathbf{h})\n        \\mathbf{h} &= \\mathrm{LayerNorm}(\\mathbf{x} +\n        \\mathrm{Multihead}(\\mathbf{x}, \\mathbf{y}, \\mathbf{y}))\n    Args:\n        channels (int): Size of each input sample.\n        heads (int, optional): Number of multi-head-attentions.\n            (default: :obj:`1`)\n        norm (str, optional): If set to :obj:`False`, will not apply layer\n            normalization. (default: :obj:`True`)\n        dropout (float, optional): Dropout probability of attention weights.\n            (default: :obj:`0`)\n    \"\"\"\n    def __init__(self, channels: int, heads: int = 1, layer_norm: bool = True,\n                 dropout: float = 0.0):\n        super().__init__()\n\n        self.channels = channels\n        self.heads = heads\n        self.dropout = dropout\n\n        self.attn = MultiheadAttention(\n            channels,\n            heads,\n            batch_first=True,\n            dropout=dropout,\n        )\n        self.lin = Linear(channels, channels)\n        self.layer_norm1 = LayerNorm(channels) if layer_norm else None\n        self.layer_norm2 = LayerNorm(channels) if layer_norm else None\n\n    def reset_parameters(self):\n        self.attn._reset_parameters()\n        self.lin.reset_parameters()\n        if self.layer_norm1 is not None:\n            self.layer_norm1.reset_parameters()\n        if self.layer_norm2 is not None:\n            self.layer_norm2.reset_parameters()\n\n    def forward(self, x: Tensor, y: Tensor, x_mask: Optional[Tensor] = None, y_mask: Optional[Tensor] = None) -> Tensor:\n        if y_mask is not None:\n            y_mask = ~y_mask\n        out, _ = self.attn(x, y, y, y_mask, need_weights=False)\n        if x_mask is not None:\n            out[~x_mask] = 0.\n        out = out + x\n        if self.layer_norm1 is not None:\n            out = self.layer_norm1(out)\n        out = out + self.lin(out).relu()\n        if self.layer_norm2 is not None:\n            out = self.layer_norm2(out)\n\n        return out\n\n    def __repr__(self) -> str:\n        return (f'{self.__class__.__name__}({self.channels}, '\n                f'heads={self.heads}, '\n                f'layer_norm={self.layer_norm1 is not None}, '\n                f'dropout={self.dropout})')\n\nclass SetAttentionBlock(torch.nn.Module):\n    r\"\"\"The Set Attention Block (SAB) from the `\"Set Transformer: A\n    Framework for Attention-based Permutation-Invariant Neural Networks\"\n    <https://arxiv.org/abs/1810.00825>`_ paper\n    .. math::\n        \\mathrm{SAB}(\\mathbf{X}) = \\mathrm{MAB}(\\mathbf{x}, \\mathbf{y})\n    Args:\n        channels (int): Size of each input sample.\n        heads (int, optional): Number of multi-head-attentions.\n            (default: :obj:`1`)\n        norm (str, optional): If set to :obj:`False`, will not apply layer\n            normalization. (default: :obj:`True`)\n        dropout (float, optional): Dropout probability of attention weights.\n            (default: :obj:`0`)\n    \"\"\"\n    def __init__(self, channels: int, heads: int = 1, layer_norm: bool = True,\n                 dropout: float = 0.0):\n        super().__init__()\n        self.mab = MultiheadAttentionBlock(channels, heads, layer_norm,\n                                           dropout)\n\n    def reset_parameters(self):\n        self.mab.reset_parameters()\n\n    def forward(self, x: Tensor, mask: Optional[Tensor] = None) -> Tensor:\n        return self.mab(x, x, mask, mask)\n\n    def __repr__(self) -> str:\n        return (f'{self.__class__.__name__}({self.mab.channels}, '\n                f'heads={self.mab.heads}, '\n                f'layer_norm={self.mab.layer_norm1 is not None}, '\n                f'dropout={self.mab.dropout})')\n\nclass PoolingByMultiheadAttention(torch.nn.Module):\n    r\"\"\"The Pooling by Multihead Attention (PMA) layer from the `\"Set\n    Transformer: A Framework for Attention-based Permutation-Invariant Neural\n    Networks\" <https://arxiv.org/abs/1810.00825>`_ paper\n    .. math::\n        \\mathrm{PMA}(\\mathbf{X}) = \\mathrm{MAB}(\\mathbf{S}, \\mathbf{x})\n    where :math:`\\mathbf{S}` denotes :obj:`num_seed_points` learnable vectors.\n    Args:\n        channels (int): Size of each input sample.\n        num_seed_points (int, optional): Number of seed points.\n            (default: :obj:`1`)\n        heads (int, optional): Number of multi-head-attentions.\n            (default: :obj:`1`)\n        norm (str, optional): If set to :obj:`False`, will not apply layer\n            normalization. (default: :obj:`True`)\n        dropout (float, optional): Dropout probability of attention weights.\n            (default: :obj:`0`)\n    \"\"\"\n    def __init__(self, channels: int, num_seed_points: int = 1, heads: int = 1,\n                 layer_norm: bool = True, dropout: float = 0.0):\n        super().__init__()\n        self.lin = Linear(channels, channels)\n        self.seed = Parameter(torch.Tensor(1, num_seed_points, channels))\n        self.mab = MultiheadAttentionBlock(channels, heads, layer_norm,\n                                           dropout)\n        self.reset_parameters()\n\n    def reset_parameters(self):\n        self.lin.reset_parameters()\n        torch.nn.init.xavier_uniform_(self.seed)\n        self.mab.reset_parameters()\n\n    def forward(self, x: Tensor, mask: Optional[Tensor] = None) -> Tensor:\n        x = self.lin(x).relu()\n        return self.mab(self.seed.expand(x.size(0), -1, -1), x, y_mask=mask)\n\n    def __repr__(self) -> str:\n        return (f'{self.__class__.__name__}({self.seed.size(2)}, '\n                f'num_seed_points={self.seed.size(1)}, '\n                f'heads={self.mab.heads}, '\n                f'layer_norm={self.mab.layer_norm1 is not None}, '\n                f'dropout={self.mab.dropout})')\n\nclass SetTransformerAggregation(Aggregation):\n    r\"\"\"Performs \"Set Transformer\" aggregation in which the elements to\n    aggregate are processed by multi-head attention blocks, as described in\n    the `\"Graph Neural Networks with Adaptive Readouts\"\n    <https://arxiv.org/abs/2211.04952>`_ paper.\n\n    Args:\n        channels (int): Size of each input sample.\n        num_seed_points (int, optional): Number of seed points.\n            (default: :obj:`1`)\n        num_encoder_blocks (int, optional): Number of Set Attention Blocks\n            (SABs) in the encoder. (default: :obj:`1`).\n        num_decoder_blocks (int, optional): Number of Set Attention Blocks\n            (SABs) in the decoder. (default: :obj:`1`).\n        heads (int, optional): Number of multi-head-attentions.\n            (default: :obj:`1`)\n        concat (bool, optional): If set to :obj:`False`, the seed embeddings\n            are averaged instead of concatenated. (default: :obj:`True`)\n        norm (str, optional): If set to :obj:`True`, will apply layer\n            normalization. (default: :obj:`False`)\n        dropout (float, optional): Dropout probability of attention weights.\n            (default: :obj:`0`)\n    \"\"\"\n    def __init__(\n        self,\n        channels: int,\n        num_seed_points: int = 1,\n        num_encoder_blocks: int = 1,\n        num_decoder_blocks: int = 1,\n        heads: int = 1,\n        concat: bool = True,\n        layer_norm: bool = False,\n        dropout: float = 0.0,\n    ):\n        super().__init__()\n\n        self.channels = channels\n        self.num_seed_points = num_seed_points\n        self.heads = heads\n        self.concat = concat\n        self.layer_norm = layer_norm\n        self.dropout = dropout\n\n        self.encoders = torch.nn.ModuleList([\n            SetAttentionBlock(channels, heads, layer_norm, dropout)\n            for _ in range(num_encoder_blocks)\n        ])\n\n        self.pma = PoolingByMultiheadAttention(channels, num_seed_points,\n                                               heads, layer_norm, dropout)\n\n        self.decoders = torch.nn.ModuleList([\n            SetAttentionBlock(channels, heads, layer_norm, dropout)\n            for _ in range(num_decoder_blocks)\n        ])\n\n    def reset_parameters(self):\n        for encoder in self.encoders:\n            encoder.reset_parameters()\n        self.pma.reset_parameters()\n        for decoder in self.decoders:\n            decoder.reset_parameters()\n\n    def forward(self, x: Tensor, index: Optional[Tensor] = None, ptr: Optional[Tensor] = None, dim_size: Optional[int] = None, dim: int = -2) -> Tensor:\n        x, mask = self.to_dense_batch(x, index, ptr, dim_size, dim)\n        for encoder in self.encoders:\n            x = encoder(x, mask)\n        x = self.pma(x, mask)\n        for decoder in self.decoders:\n            x = decoder(x)\n        return x.flatten(1, 2) if self.concat else x.mean(dim=1)\n\n    def __repr__(self) -> str:\n        return (f'{self.__class__.__name__}({self.channels}, '\n                f'num_seed_points={self.num_seed_points}, '\n                f'heads={self.heads}, '\n                f'layer_norm={self.layer_norm}, '\n                f'dropout={self.dropout})')","metadata":{"execution":{"iopub.status.busy":"2023-04-18T20:45:21.209347Z","iopub.execute_input":"2023-04-18T20:45:21.209703Z","iopub.status.idle":"2023-04-18T20:45:21.240696Z","shell.execute_reply.started":"2023-04-18T20:45:21.209663Z","shell.execute_reply":"2023-04-18T20:45:21.239553Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from typing import Any, Callable, List, Dict, Optional, Sequence, Tuple, Union","metadata":{"execution":{"iopub.status.busy":"2023-04-18T20:45:21.242051Z","iopub.execute_input":"2023-04-18T20:45:21.242699Z","iopub.status.idle":"2023-04-18T20:45:21.252232Z","shell.execute_reply.started":"2023-04-18T20:45:21.242660Z","shell.execute_reply":"2023-04-18T20:45:21.251374Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DEBUG = False","metadata":{"execution":{"iopub.status.busy":"2023-04-18T20:45:21.253928Z","iopub.execute_input":"2023-04-18T20:45:21.254495Z","iopub.status.idle":"2023-04-18T20:45:21.265186Z","shell.execute_reply.started":"2023-04-18T20:45:21.254458Z","shell.execute_reply":"2023-04-18T20:45:21.264172Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    # device\n    device = 'cuda'\n    # rootdir\n    rootdir = '/kaggle/input/icecube-neutrinos-in-deep-ice'\n    # batch_data_dir\n    datadir = f'{rootdir}/train' if DEBUG else f'{rootdir}/test'\n    # sensor\n    sensor_file = f'{rootdir}/sensor_geometry.csv'\n    # batch_size\n    batch_size = 12\n    # num_workers\n    num_workers = 2\n    # TTA\n    angles = [ 0, ]\n    # angles = [ 0, 180 ]","metadata":{"execution":{"iopub.status.busy":"2023-04-18T20:45:21.266612Z","iopub.execute_input":"2023-04-18T20:45:21.267147Z","iopub.status.idle":"2023-04-18T20:45:21.275647Z","shell.execute_reply.started":"2023-04-18T20:45:21.267072Z","shell.execute_reply":"2023-04-18T20:45:21.274690Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_submit_or_scoring():\n    df = pd.read_parquet(os.path.join(CFG.rootdir, 'test_meta.parquet'))\n    return len(df['batch_id'].unique()) == 1 and len(df['event_id'].unique()) == 3\n\nSUBMIT = get_submit_or_scoring()\nprint(f'SUBMIT : {SUBMIT}')","metadata":{"execution":{"iopub.status.busy":"2023-04-18T20:45:21.277097Z","iopub.execute_input":"2023-04-18T20:45:21.277481Z","iopub.status.idle":"2023-04-18T20:45:21.403207Z","shell.execute_reply.started":"2023-04-18T20:45:21.277445Z","shell.execute_reply":"2023-04-18T20:45:21.402022Z"},"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 ICM_01011(torch.nn.Module):\n    def __init__(self, nb_inputs=6, nb_nearest_neighbours=32):\n        super().__init__()\n        self.detector = KNNGraphBuilder(nb_nearest_neighbours=nb_nearest_neighbours)\n        self.prefetch = PreFetch()\n        self.backbone = DynEdge(\n            nb_inputs=nb_inputs,\n            global_pooling_schemes=[ 'min', 'max', 'mean' ],\n        )\n        self.head = BASEHead(self.backbone._nb_outputs, 3)\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        h = self.detector(x)\n        h = self.prefetch(h)\n        h = self.backbone(h)\n        h = self.head(h)\n        return h\n\nclass ICM_01029(torch.nn.Module):\n    def __init__(self, nb_inputs=6, nb_nearest_neighbours=16):\n        super().__init__()\n        self.detector = KNNGraphBuilder(nb_nearest_neighbours=nb_nearest_neighbours, columns=[0, 1, 2, 3])\n        self.prefetch = PreFetch()\n        self.backbone = DynEdgeTrns(\n            nb_inputs=nb_inputs,\n        )\n        self.head = BASEHead(256, 3)\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        h = self.detector(x)\n        h = self.prefetch(h)\n        h = self.backbone(h)\n        h = self.head(h)\n        return h\n\nclass ICM_01030(torch.nn.Module):\n    def __init__(self, nb_inputs=6, nb_nearest_neighbours=32):\n        super().__init__()\n        self.detector = KNNGraphBuilder(nb_nearest_neighbours=nb_nearest_neighbours, columns=[0, 1, 2, 3])\n        self.prefetch = PreFetch()\n        self.backbone = DynEdgeTrns(\n            nb_inputs=nb_inputs,\n        )\n        self.head = BASEHead(256, 3)\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        h = self.detector(x)\n        h = self.prefetch(h)\n        h = self.backbone(h)\n        h = self.head(h)\n        return h\n\nclass ICM_01031(torch.nn.Module):\n    def __init__(self, nb_inputs=6, nb_nearest_neighbours=32):\n        super().__init__()\n        self.detector = KNNGraphBuilder(nb_nearest_neighbours=nb_nearest_neighbours, columns=[0, 1, 2, 3])\n        self.prefetch = PreFetch()\n        self.backbone = DynEdgeAttn2(\n            nb_inputs=nb_inputs,\n        )\n        self.head = BASEHead(256, 3)\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        h = self.detector(x)\n        h = self.prefetch(h)\n        h = self.backbone(h)\n        h = self.head(h)\n        return h\n\nclass ICM_12000(torch.nn.Module):\n    def __init__(self, nb_inputs=6, nb_nearest_neighbours=32):\n        super().__init__()\n        self.detector = KNNGraphBuilder(nb_nearest_neighbours=nb_nearest_neighbours, columns=[0, 1, 2, 3])\n        self.prefetch = PreFetch()\n        self.backbone = DynEdgeTrnsA(\n            nb_inputs=nb_inputs,\n        )\n        self.head = BASEHead(256, 3)\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        h = self.detector(x)\n        h = self.prefetch(h)\n        h = self.backbone(h)\n        h = self.head(h)\n        return h\n\nclass ICM_12001(torch.nn.Module):\n    def __init__(self, nb_inputs=6, nb_nearest_neighbours=32):\n        super().__init__()\n        self.detector = KNNGraphBuilder(nb_nearest_neighbours=nb_nearest_neighbours, columns=[0, 1, 2, 3])\n        self.prefetch = PreFetch()\n        self.backbone = DynEdgeTrnsB(\n            nb_inputs=nb_inputs,\n        )\n        self.head = BASEHead(256, 3)\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        h = self.detector(x)\n        h = self.prefetch(h)\n        h = self.backbone(h)\n        h = self.head(h)\n        return h\n\nclass ICM_12010(torch.nn.Module):\n    def __init__(self, nb_inputs=6, nb_nearest_neighbours=32):\n        super().__init__()\n        self.detector = KNNGraphBuilder(nb_nearest_neighbours=nb_nearest_neighbours)\n        self.preproc1 = PreProc1()\n        self.preproc2 = PreProc2()\n        self.backbone = DynEdgeTrnsA(\n            nb_inputs=nb_inputs,\n        )\n        self.head = BASEHead(256, 3)\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        h = self.preproc1(x)\n        h = self.detector(h)\n        h = self.preproc2(h)\n        h = self.backbone(h)\n        h = self.head(h)\n        return h\n\nclass ICM_12031(torch.nn.Module):\n    def __init__(self, nb_inputs=6, nb_nearest_neighbours=32):\n        super().__init__()\n        self.detector = KNNGraphBuilder(nb_nearest_neighbours=nb_nearest_neighbours, columns=[0, 1, 2, 3])\n        self.prefetch = PreFetch()\n        self.backbone = DynEdgeAttn2A(\n            nb_inputs=nb_inputs,\n        )\n        self.head = BASEHead(256, 3)\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        h = self.detector(x)\n        h = self.prefetch(h)\n        h, _, _ = self.backbone(h)\n        h = self.head(h)\n        return h\n    \nclass ICM_12034(torch.nn.Module):\n    def __init__(self, nb_inputs=6, nb_nearest_neighbours=16):\n        super().__init__()\n        self.detector = KNNGraphBuilder(nb_nearest_neighbours=nb_nearest_neighbours)\n        self.prefetch = PreFetch()\n        self.backbone = DynEdgeAttn2C(\n            nb_inputs=nb_inputs,\n        )\n        self.head = BASEHead(self.backbone._nb_outputs, 3)\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        h = self.detector(x)\n        h = self.prefetch(h)\n        h = self.backbone(h)\n        h = self.head(h)\n        return h\n\nclass PreFetch(torch.nn.Module):\n    def __init__(self):\n        super().__init__()\n\n    def forward(self, data):\n        # xyz\n        data.x[:, 0] /= 500.0\n        data.x[:, 1] /= 500.0\n        data.x[:, 2] /= 500.0\n        # time\n        data.x[:, 3] = (data.x[:, 3] - 1.0e04) / 3.0e4\n        # charge\n        data.x[:, 4] = torch.log10(data.x[:, 4]) / 3.0\n        return data\n\nclass PreProc1(torch.nn.Module):\n    def __init__(self):\n        super().__init__()\n\n    def forward(self, data):\n        # xyz\n        data.x[:, 2] *= 7.35\n        return data\n\nclass PreProc2(torch.nn.Module):\n    def __init__(self):\n        super().__init__()\n\n    def forward(self, data):\n        # xyz\n        data.x[:, 0] /= 500.0\n        data.x[:, 1] /= 500.0\n        data.x[:, 2] /= 500.0 * 7.35\n        # time\n        data.x[:, 3] = (data.x[:, 3] - 1.0e04) / 3.0e4\n        # charge\n        data.x[:, 4] = torch.log10(data.x[:, 4]) / 3.0\n        return data\n\nclass BASEHead(torch.nn.Module):\n    def __init__(self, feature_size, output_size):\n        super().__init__()\n        self.head = torch.nn.Linear(feature_size, output_size)\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        h = self.head(x)\n        kappa = torch.linalg.vector_norm(h, dim=1) + torch.finfo(h.dtype).eps\n        vec_x = h[:, 0] / kappa\n        vec_y = h[:, 1] / kappa\n        vec_z = h[:, 2] / kappa\n        return torch.stack((vec_x, vec_y, vec_z, kappa), dim=1)\n\nclass DynEdge(torch.nn.Module):\n    \"\"\" DynEdge : dynamical edge convolutional model \"\"\"\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        # Latent feature subset for computing nearest neighbours in DynEdge.\n        if features_subset is None:\n            features_subset = slice(0, 3)\n        # DynEdge layer sizes\n        if dynedge_layer_sizes is None:\n            dynedge_layer_sizes = [ (128, 256), (336, 256), (336, 256), (336, 256) ]\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(all(size > 0 for size in sizes) for sizes in dynedge_layer_sizes)\n        self._dynedge_layer_sizes = dynedge_layer_sizes\n        # Post-processing layer sizes\n        if post_processing_layer_sizes is None:\n            post_processing_layer_sizes = [ 336, 256 ]\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        self._post_processing_layer_sizes = post_processing_layer_sizes\n        # Read-out layer sizes\n        if readout_layer_sizes is None:\n            readout_layer_sizes = [\n                128,\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        self._readout_layer_sizes = readout_layer_sizes\n        # Global pooling scheme(s)\n        if isinstance(global_pooling_schemes, str):\n            global_pooling_schemes = [global_pooling_schemes]\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        self._global_pooling_schemes = global_pooling_schemes\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 = add_global_variables_after_pooling\n        # Base class constructor\n        self._nb_inputs = nb_inputs\n        self._nb_outputs = self._readout_layer_sizes[-1]\n        super().__init__()\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        self._construct_layers()\n        # Add Device Parameter:\n        self._device = torch.nn.Parameter(torch.empty(0))\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        self._conv_layers = torch.nn.ModuleList()\n        nb_latent_features = nb_input_features\n        for sizes in 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(self._activation)\n\n            conv_layer = DynEdgeConv(\n                torch.nn.Sequential(*layers),\n                aggr=\"add\",\n                nb_neighbors=self._nb_neighbours,\n                features_subset=self._features_subset,\n            )\n            self._conv_layers.append(conv_layer)\n            nb_latent_features = nb_out\n        # Post-processing operations\n        nb_latent_features = (\n            sum(sizes[-1] for sizes in self._dynedge_layer_sizes) + nb_input_features\n        )\n        post_processing_layers = []\n        layer_sizes = [nb_latent_features] + list(self._post_processing_layer_sizes)\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        self._post_processing = torch.nn.Sequential(*post_processing_layers)\n        # Read-out operations\n        nb_poolings = (\n            len(self._global_pooling_schemes) if self._global_pooling_schemes 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        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    def _global_pooling(self, x: torch.Tensor, batch: torch.LongTensor) -> torch.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                pooled_x, _ = pooled_x\n            pooled.append(pooled_x)\n        return torch.cat(pooled, dim=1)\n\n    def _calculate_global_variables(self, x: torch.Tensor, edge_index: torch.LongTensor, batch: torch.LongTensor, *additional_attributes: torch.Tensor) -> torch.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        # Calculate mean features\n        global_means = scatter_mean(x, batch, dim=0)\n        # Add global variables\n        global_variables = torch.cat([global_means, h_x, h_y, h_z, h_t ] + [attr.unsqueeze(dim=1) for attr in additional_attributes], dim=1)\n        return global_variables\n\n    def forward(self, data: Data) -> torch.Tensor:\n        # Convenience variables\n        x, edge_index, batch = data.x, data.edge_index, data.batch\n        global_variables = self._calculate_global_variables(x, edge_index, batch, torch.log10(data.n_pulses))\n        # Distribute global variables out to each node\n        if not self._add_global_variables_after_pooling:\n            distribute = (batch.unsqueeze(dim=1) == torch.unique(batch).unsqueeze(dim=0)).type(torch.float)\n            global_variables_distributed = torch.sum(distribute.unsqueeze(dim=2) * global_variables.unsqueeze(dim=0), dim=1)\n            x = torch.cat((x, global_variables_distributed), dim=1)\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        # Skip-cat\n        x = torch.cat(skip_connections, dim=1)\n        # Post-processing\n        x = self._post_processing(x)\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([ x, global_variables ], dim=1)\n        x = self._readout(x)\n        return x\n\n    @property\n    def device(self):\n        return self._device.device\n\nclass DynEdgeAttn(torch.nn.Module):\n    \"\"\" DynEdge : dynamical edge convolutional model \"\"\"\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    ):\n        # Latent feature subset for computing nearest neighbours in DynEdge.\n        if features_subset is None:\n            features_subset = slice(0, 3)\n        # DynEdge layer sizes\n        if dynedge_layer_sizes is None:\n            dynedge_layer_sizes = [ (128, 256), (336, 256), (336, 256), (336, 256) ]\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(all(size > 0 for size in sizes) for sizes in dynedge_layer_sizes)\n        self._dynedge_layer_sizes = dynedge_layer_sizes\n        # Post-processing layer sizes\n        if post_processing_layer_sizes is None:\n            post_processing_layer_sizes = [ 336, 256 ]\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        self._post_processing_layer_sizes = post_processing_layer_sizes\n        # Read-out layer sizes\n        if readout_layer_sizes is None:\n            readout_layer_sizes = [\n                # 128,\n                256,\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        self._readout_layer_sizes = readout_layer_sizes\n        self._add_global_variables_after_pooling = False\n        # Base class constructor\n        self._nb_inputs = nb_inputs\n        self._nb_outputs = self._readout_layer_sizes[-1]\n        super().__init__()\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        self._construct_layers()\n        # Attn\n        self.attn = torch_geometric.nn.AttentionalAggregation(torch.nn.Linear(256, 1))\n        # Add Device Parameter:\n        self._device = torch.nn.Parameter(torch.empty(0))\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 sizes in 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(self._activation)\n\n            conv_layer = DynEdgeConv(\n                torch.nn.Sequential(*layers),\n                aggr=\"add\",\n                nb_neighbors=self._nb_neighbours,\n                features_subset=self._features_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) + nb_input_features\n        )\n\n        post_processing_layers = []\n        layer_sizes = [nb_latent_features] + list(self._post_processing_layer_sizes)\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\n        self._post_processing = torch.nn.Sequential(*post_processing_layers)\n\n        # Read-out operations\n        # nb_poolings = (\n        #     len(self._global_pooling_schemes) if self._global_pooling_schemes else 1\n        # )\n        nb_poolings = (1)\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(self._activation)\n\n        self._readout = torch.nn.Sequential(*readout_layers)\n\n    def _calculate_global_variables(self, x: torch.Tensor, edge_index: torch.LongTensor, batch: torch.LongTensor, *additional_attributes: torch.Tensor) -> torch.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        # Calculate mean features\n        global_means = scatter_mean(x, batch, dim=0)\n        # Add global variables\n        global_variables = torch.cat([global_means, h_x, h_y, h_z, h_t ] + [attr.unsqueeze(dim=1) for attr in additional_attributes], dim=1)\n        return global_variables\n\n    def forward(self, data: Data) -> torch.Tensor:\n        # Convenience variables\n        x, edge_index, batch = data.x, data.edge_index, data.batch\n        global_variables = self._calculate_global_variables(x, edge_index, batch, torch.log10(data.n_pulses))\n        # Distribute global variables out to each node\n        if not self._add_global_variables_after_pooling:\n            distribute = (batch.unsqueeze(dim=1) == torch.unique(batch).unsqueeze(dim=0)).type(torch.float)\n            global_variables_distributed = torch.sum(distribute.unsqueeze(dim=2) * global_variables.unsqueeze(dim=0), dim=1)\n            x = torch.cat((x, global_variables_distributed), dim=1)\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        # Skip-cat\n        x = torch.cat(skip_connections, dim=1)\n\n        # Post-processing\n        x = self._post_processing(x)\n\n        # Attantion\n        x = self.attn(x, batch)\n\n        x = self._readout(x)\n        return x\n        # return x, q, m\n\n    @property\n    def device(self):\n        return self._device.device\n\nclass DynEdgeAttn2(torch.nn.Module):\n    \"\"\" DynEdge : dynamical edge convolutional model \"\"\"\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        layer_norm : bool = False,\n    ):\n        # Latent feature subset for computing nearest neighbours in DynEdge.\n        if features_subset is None:\n            features_subset = [ slice(0, 4), slice(0, 4), slice(0, 3), slice(0, 3) ]\n        # DynEdge layer sizes\n        if dynedge_layer_sizes is None:\n            dynedge_layer_sizes = [ (128, 256), (336, 256), (336, 256), (336, 256) ]\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(all(size > 0 for size in sizes) for sizes in dynedge_layer_sizes)\n        self._dynedge_layer_sizes = dynedge_layer_sizes\n        # Post-processing layer sizes\n        if post_processing_layer_sizes is None:\n            post_processing_layer_sizes = [ 336, 256 ]\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        self._post_processing_layer_sizes = post_processing_layer_sizes\n        # Read-out layer sizes\n        if readout_layer_sizes is None:\n            readout_layer_sizes = [\n                # 128,\n                256,\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        self._readout_layer_sizes = readout_layer_sizes\n        self._add_global_variables_after_pooling = False\n        # Base class constructor\n        self._nb_inputs = nb_inputs\n        self._nb_outputs = self._readout_layer_sizes[-1]\n        super().__init__()\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        self._construct_layers()\n        # Attn\n        self.attn = torch_geometric.nn.AttentionalAggregation(torch.nn.Linear(256, 1))\n        self.norm = torch.nn.LayerNorm(256) if layer_norm else None\n        # Add Device Parameter:\n        self._device = torch.nn.Parameter(torch.empty(0))\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 sizes in 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(self._activation)\n\n            conv_layer = DynEdgeConv(\n                torch.nn.Sequential(*layers),\n                aggr=\"add\",\n                nb_neighbors=self._nb_neighbours,\n                features_subset=self._features_subset[ix],\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) + nb_input_features\n        )\n\n        post_processing_layers = []\n        layer_sizes = [nb_latent_features] + list(self._post_processing_layer_sizes)\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\n        self._post_processing = torch.nn.Sequential(*post_processing_layers)\n\n        # Read-out operations\n        # nb_poolings = (\n        #     len(self._global_pooling_schemes) if self._global_pooling_schemes else 1\n        # )\n        nb_poolings = (1)\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(self._activation)\n\n        self._readout = torch.nn.Sequential(*readout_layers)\n\n    def _calculate_global_variables(self, x: torch.Tensor, edge_index: torch.LongTensor, batch: torch.LongTensor, *additional_attributes: torch.Tensor) -> torch.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        # Calculate mean features\n        global_means = scatter_mean(x, batch, dim=0)\n        # Add global variables\n        global_variables = torch.cat([global_means, h_x, h_y, h_z, h_t ] + [attr.unsqueeze(dim=1) for attr in additional_attributes], dim=1)\n        return global_variables\n\n    def forward(self, data: Data) -> torch.Tensor:\n        # Convenience variables\n        x, edge_index, batch = data.x, data.edge_index, data.batch\n        global_variables = self._calculate_global_variables(x, edge_index, batch, torch.log10(data.n_pulses))\n        # Distribute global variables out to each node\n        if not self._add_global_variables_after_pooling:\n            distribute = (batch.unsqueeze(dim=1) == torch.unique(batch).unsqueeze(dim=0)).type(torch.float)\n            global_variables_distributed = torch.sum(distribute.unsqueeze(dim=2) * global_variables.unsqueeze(dim=0), dim=1)\n            x = torch.cat((x, global_variables_distributed), dim=1)\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        # Skip-cat\n        x = torch.cat(skip_connections, dim=1)\n        \n        # Post-processing\n        x = self._post_processing(x)\n\n        # Attantion\n        x = self.attn(x, batch)\n        \n        # LayerNorm\n        if self.norm:\n            x = self.norm(x)\n\n        x = self._readout(x)\n        return x\n        # return x, q, m\n\n    @property\n    def device(self):\n        return self._device.device\n\nclass DynEdgeAttn2A(torch.nn.Module):\n    \"\"\" DynEdge : dynamical edge convolutional model \"\"\"\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    ):\n        # Latent feature subset for computing nearest neighbours in DynEdge.\n        if features_subset is None:\n            features_subset = [ slice(0, 4), slice(0, 4), slice(0, 3), slice(0, 3) ]\n        # DynEdge layer sizes\n        if dynedge_layer_sizes is None:\n            dynedge_layer_sizes = [ (128, 256), (336, 256), (336, 256), (336, 256) ]\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(all(size > 0 for size in sizes) for sizes in dynedge_layer_sizes)\n        self._dynedge_layer_sizes = dynedge_layer_sizes\n        # Post-processing layer sizes\n        if post_processing_layer_sizes is None:\n            post_processing_layer_sizes = [ 336, 256 ]\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        self._post_processing_layer_sizes = post_processing_layer_sizes\n        # Read-out layer sizes\n        if readout_layer_sizes is None:\n            readout_layer_sizes = [\n                # 128,\n                256,\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        self._readout_layer_sizes = readout_layer_sizes\n        self._add_global_variables_after_pooling = False\n        # Base class constructor\n        self._nb_inputs = nb_inputs\n        self._nb_outputs = self._readout_layer_sizes[-1]\n        super().__init__()\n        # Remaining member variables()\n        self._activation = torch.nn.PReLU(init=0.01)\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._construct_layers()\n        # Attn\n        self.attn = torch_geometric.nn.AttentionalAggregation(torch.nn.Linear(256, 1))\n        # Add Device Parameter:\n        self._device = torch.nn.Parameter(torch.empty(0))\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 sizes in 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(self._activation)\n\n            conv_layer = DynEdgeConv(\n                torch.nn.Sequential(*layers),\n                aggr=\"add\",\n                nb_neighbors=16,\n                features_subset=self._features_subset[ix],\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) + nb_input_features\n        )\n\n        post_processing_layers = []\n        layer_sizes = [nb_latent_features] + list(self._post_processing_layer_sizes)\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\n        self._post_processing = torch.nn.Sequential(*post_processing_layers)\n\n        # Read-out operations\n        # nb_poolings = (\n        #     len(self._global_pooling_schemes) if self._global_pooling_schemes else 1\n        # )\n        nb_poolings = (1)\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(self._activation)\n\n        self._readout = torch.nn.Sequential(*readout_layers)\n\n    def _calculate_global_variables(self, x: torch.Tensor, edge_index: torch.LongTensor, batch: torch.LongTensor, *additional_attributes: torch.Tensor) -> torch.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        # Calculate mean features\n        global_means = scatter_mean(x, batch, dim=0)\n        # Add global variables\n        global_variables = torch.cat([global_means, h_x, h_y, h_z, h_t ] + [attr.unsqueeze(dim=1) for attr in additional_attributes], dim=1)\n        return global_variables\n\n    def forward(self, data: Data) -> torch.Tensor:\n        # Convenience variables\n        x, edge_index, batch = data.x, data.edge_index, data.batch\n        global_variables = self._calculate_global_variables(x, edge_index, batch, torch.log10(data.n_pulses))\n        # Distribute global variables out to each node\n        if not self._add_global_variables_after_pooling:\n            distribute = (batch.unsqueeze(dim=1) == torch.unique(batch).unsqueeze(dim=0)).type(torch.float)\n            global_variables_distributed = torch.sum(distribute.unsqueeze(dim=2) * global_variables.unsqueeze(dim=0), dim=1)\n            x = torch.cat((x, global_variables_distributed), dim=1)\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        # Skip-cat\n        x = torch.cat(skip_connections, dim=1)\n\n        # Post-processing\n        x = self._post_processing(x)\n\n        # Attantion\n        x = self.attn(x, batch)\n\n        x = self._readout(x)\n        # return x\n        return x, None, None\n\n    @property\n    def device(self):\n        return self._device.device\n\nclass DynEdgeAttn2C(torch.nn.Module):\n    \"\"\" DynEdge : dynamical edge convolutional model \"\"\"\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    ):\n        # Latent feature subset for computing nearest neighbours in DynEdge.\n        if features_subset is None:\n            features_subset = [ slice(0, 4), slice(0, 4), slice(0, 3), slice(0, 3) ]\n        # DynEdge layer sizes\n        if dynedge_layer_sizes is None:\n            dynedge_layer_sizes = [ (128, 256), (336, 256), (336, 256), (336, 256) ]\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(all(size > 0 for size in sizes) for sizes in dynedge_layer_sizes)\n        self._dynedge_layer_sizes = dynedge_layer_sizes\n        # Post-processing layer sizes\n        if post_processing_layer_sizes is None:\n            post_processing_layer_sizes = [ 336, 256 ]\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        self._post_processing_layer_sizes = post_processing_layer_sizes\n        # Read-out layer sizes\n        if readout_layer_sizes is None:\n            readout_layer_sizes = [\n                # 128,\n                256,\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        self._readout_layer_sizes = readout_layer_sizes\n        self._add_global_variables_after_pooling = False\n        # Base class constructor\n        self._nb_inputs = nb_inputs\n        self._nb_outputs = self._readout_layer_sizes[-1]\n        super().__init__()\n        # Remaining member variables()\n        self._activation = torch.nn.PReLU(init=0.01) # 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        self._construct_layers()\n        # Attn\n        self.attn = torch_geometric.nn.AttentionalAggregation(torch.nn.Linear(256, 1))\n        # Add Device Parameter:\n        self._device = torch.nn.Parameter(torch.empty(0))\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 sizes in 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(self._activation)\n\n            conv_layer = DynEdgeConv(\n                torch.nn.Sequential(*layers),\n                aggr=\"add\",\n                nb_neighbors=16, # self._nb_neighbours,\n                features_subset=self._features_subset[ix],\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) + nb_input_features\n        )\n\n        post_processing_layers = []\n        layer_sizes = [nb_latent_features] + list(self._post_processing_layer_sizes)\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\n        self._post_processing = torch.nn.Sequential(*post_processing_layers)\n\n        # Read-out operations\n        # nb_poolings = (\n        #     len(self._global_pooling_schemes) if self._global_pooling_schemes else 1\n        # )\n        nb_poolings = (1)\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(self._activation)\n\n        self._readout = torch.nn.Sequential(*readout_layers)\n\n    def _calculate_global_variables(self, x: torch.Tensor, edge_index: torch.LongTensor, batch: torch.LongTensor, *additional_attributes: torch.Tensor) -> torch.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        # Calculate mean features\n        global_means = scatter_mean(x, batch, dim=0)\n        # Add global variables\n        global_variables = torch.cat([global_means, h_x, h_y, h_z, h_t ] + [attr.unsqueeze(dim=1) for attr in additional_attributes], dim=1)\n        return global_variables\n\n    def forward(self, data: Data) -> torch.Tensor:\n        # Convenience variables\n        x, edge_index, batch = data.x, data.edge_index, data.batch\n        global_variables = self._calculate_global_variables(x, edge_index, batch, torch.log10(data.n_pulses))\n        # Distribute global variables out to each node\n        if not self._add_global_variables_after_pooling:\n            distribute = (batch.unsqueeze(dim=1) == torch.unique(batch).unsqueeze(dim=0)).type(torch.float)\n            global_variables_distributed = torch.sum(distribute.unsqueeze(dim=2) * global_variables.unsqueeze(dim=0), dim=1)\n            x = torch.cat((x, global_variables_distributed), dim=1)\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        # Skip-cat\n        x = torch.cat(skip_connections, dim=1)\n\n        # Post-processing\n        x = self._post_processing(x)\n\n        # Attantion\n        x = self.attn(x, batch)\n\n        x = self._readout(x)\n\n        return x\n\n    @property\n    def device(self):\n        return self._device.device\n    \nclass DynEdgeTrns(torch.nn.Module):\n    \"\"\" DynEdge : dynamical edge convolutional model \"\"\"\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    ):\n        # Latent feature subset for computing nearest neighbours in DynEdge.\n        if features_subset is None:\n            features_subset = slice(0, 3)\n        # DynEdge layer sizes\n        if dynedge_layer_sizes is None:\n            dynedge_layer_sizes = [ (128, 256), (336, 256), (336, 256), (336, 256) ]\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(all(size > 0 for size in sizes) for sizes in dynedge_layer_sizes)\n        self._dynedge_layer_sizes = dynedge_layer_sizes\n        # Post-processing layer sizes\n        if post_processing_layer_sizes is None:\n            post_processing_layer_sizes = [ 336, 256 ]\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        self._post_processing_layer_sizes = post_processing_layer_sizes\n        # Read-out layer sizes\n        if readout_layer_sizes is None:\n            readout_layer_sizes = [\n                # 128,\n                256,\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        self._readout_layer_sizes = readout_layer_sizes\n        self._add_global_variables_after_pooling = False\n        # Base class constructor\n        self._nb_inputs = nb_inputs\n        self._nb_outputs = self._readout_layer_sizes[-1]\n        super().__init__()\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        self._construct_layers()\n        # Attn\n        self.trns = SetTransformerAggregation(256, heads=8)\n        # Add Device Parameter:\n        self._device = torch.nn.Parameter(torch.empty(0))\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 sizes in 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(self._activation)\n\n            conv_layer = DynEdgeConv(\n                torch.nn.Sequential(*layers),\n                aggr=\"add\",\n                nb_neighbors=self._nb_neighbours,\n                features_subset=self._features_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) + nb_input_features\n        )\n\n        post_processing_layers = []\n        layer_sizes = [nb_latent_features] + list(self._post_processing_layer_sizes)\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\n        self._post_processing = torch.nn.Sequential(*post_processing_layers)\n\n        # Read-out operations\n        # nb_poolings = (\n        #     len(self._global_pooling_schemes) if self._global_pooling_schemes else 1\n        # )\n        nb_poolings = (1)\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(self._activation)\n\n        self._readout = torch.nn.Sequential(*readout_layers)\n\n    def _calculate_global_variables(self, x: torch.Tensor, edge_index: torch.LongTensor, batch: torch.LongTensor, *additional_attributes: torch.Tensor) -> torch.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        # Calculate mean features\n        global_means = scatter_mean(x, batch, dim=0)\n        # Add global variables\n        global_variables = torch.cat([global_means, h_x, h_y, h_z, h_t ] + [attr.unsqueeze(dim=1) for attr in additional_attributes], dim=1)\n        return global_variables\n\n    def forward(self, data: Data) -> torch.Tensor:\n        # Convenience variables\n        x, edge_index, batch = data.x, data.edge_index, data.batch\n        global_variables = self._calculate_global_variables(x, edge_index, batch, torch.log10(data.n_pulses))\n        # Distribute global variables out to each node\n        if not self._add_global_variables_after_pooling:\n            distribute = (batch.unsqueeze(dim=1) == torch.unique(batch).unsqueeze(dim=0)).type(torch.float)\n            global_variables_distributed = torch.sum(distribute.unsqueeze(dim=2) * global_variables.unsqueeze(dim=0), dim=1)\n            x = torch.cat((x, global_variables_distributed), dim=1)\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        # Skip-cat\n        x = torch.cat(skip_connections, dim=1)\n\n        # Post-processing\n        x = self._post_processing(x)\n\n        # Attantion\n        # x = self.attn(x, batch)\n        x = self.trns(x, batch)\n\n        x = self._readout(x)\n        return x\n\n    @property\n    def device(self):\n        return self._device.device\n\nclass DynEdgeTrnsA(torch.nn.Module):\n    \"\"\" DynEdge : dynamical edge convolutional model \"\"\"\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    ):\n        # Latent feature subset for computing nearest neighbours in DynEdge.\n        if features_subset is None:\n            features_subset = slice(0, 3)\n        # DynEdge layer sizes\n        if dynedge_layer_sizes is None:\n            dynedge_layer_sizes = [ (128, 256), (336, 256), (336, 256), (336, 256) ]\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(all(size > 0 for size in sizes) for sizes in dynedge_layer_sizes)\n        self._dynedge_layer_sizes = dynedge_layer_sizes\n        # Post-processing layer sizes\n        if post_processing_layer_sizes is None:\n            post_processing_layer_sizes = [ 336, 256 ]\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        self._post_processing_layer_sizes = post_processing_layer_sizes\n        # Read-out layer sizes\n        if readout_layer_sizes is None:\n            readout_layer_sizes = [\n                # 128,\n                256,\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        self._readout_layer_sizes = readout_layer_sizes\n        self._add_global_variables_after_pooling = False\n        # Base class constructor\n        self._nb_inputs = nb_inputs\n        self._nb_outputs = self._readout_layer_sizes[-1]\n        super().__init__()\n        # Remaining member variables()\n        self._activation = torch.nn.PReLU(init=0.01)\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._construct_layers()\n        # Attn\n        self.trns = SetTransformerAggregation(256, heads=8)\n        # Add Device Parameter:\n        self._device = torch.nn.Parameter(torch.empty(0))\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 sizes in 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(self._activation)\n\n            conv_layer = DynEdgeConv(\n                torch.nn.Sequential(*layers),\n                aggr=\"add\",\n                nb_neighbors=16,\n                features_subset=self._features_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) + nb_input_features\n        )\n\n        post_processing_layers = []\n        layer_sizes = [nb_latent_features] + list(self._post_processing_layer_sizes)\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\n        self._post_processing = torch.nn.Sequential(*post_processing_layers)\n\n        # Read-out operations\n        # nb_poolings = (\n        #     len(self._global_pooling_schemes) if self._global_pooling_schemes else 1\n        # )\n        nb_poolings = (1)\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(self._activation)\n\n        self._readout = torch.nn.Sequential(*readout_layers)\n\n    def _calculate_global_variables(self, x: torch.Tensor, edge_index: torch.LongTensor, batch: torch.LongTensor, *additional_attributes: torch.Tensor) -> torch.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        # Calculate mean features\n        global_means = scatter_mean(x, batch, dim=0)\n        # Add global variables\n        global_variables = torch.cat([global_means, h_x, h_y, h_z, h_t ] + [attr.unsqueeze(dim=1) for attr in additional_attributes], dim=1)\n        return global_variables\n\n    def forward(self, data: Data) -> torch.Tensor:\n        # Convenience variables\n        x, edge_index, batch = data.x, data.edge_index, data.batch\n        global_variables = self._calculate_global_variables(x, edge_index, batch, torch.log10(data.n_pulses))\n        # Distribute global variables out to each node\n        if not self._add_global_variables_after_pooling:\n            distribute = (batch.unsqueeze(dim=1) == torch.unique(batch).unsqueeze(dim=0)).type(torch.float)\n            global_variables_distributed = torch.sum(distribute.unsqueeze(dim=2) * global_variables.unsqueeze(dim=0), dim=1)\n            x = torch.cat((x, global_variables_distributed), dim=1)\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        # Skip-cat\n        x = torch.cat(skip_connections, dim=1)\n\n        # Post-processing\n        x = self._post_processing(x)\n\n        x = self.trns(x, batch)\n\n        x = self._readout(x)\n        \n        # return x\n        return x\n\n    @property\n    def device(self):\n        return self._device.device\n\nclass DynEdgeTrnsB(torch.nn.Module):\n    \"\"\" DynEdge : dynamical edge convolutional model \"\"\"\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    ):\n        # Latent feature subset for computing nearest neighbours in DynEdge.\n        if features_subset is None:\n            features_subset = slice(0, 3)\n        # DynEdge layer sizes\n        if dynedge_layer_sizes is None:\n            dynedge_layer_sizes = [ (128, 256), (336, 256), (336, 256), (336, 256) ]\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(all(size > 0 for size in sizes) for sizes in dynedge_layer_sizes)\n        self._dynedge_layer_sizes = dynedge_layer_sizes\n        # Post-processing layer sizes\n        if post_processing_layer_sizes is None:\n            post_processing_layer_sizes = [ 336, 256 ]\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        self._post_processing_layer_sizes = post_processing_layer_sizes\n        # Read-out layer sizes\n        if readout_layer_sizes is None:\n            readout_layer_sizes = [\n                256,\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        self._readout_layer_sizes = readout_layer_sizes\n        self._add_global_variables_after_pooling = False\n        # Base class constructor\n        self._nb_inputs = nb_inputs\n        self._nb_outputs = self._readout_layer_sizes[-1]\n        super().__init__()\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        self._construct_layers()\n        # Attn\n        self.attn = torch_geometric.nn.AttentionalAggregation(\n            torch.nn.Linear(256, 1),\n            nn=torch.nn.Sequential(\n                torch.nn.Linear(256, 256),\n                torch.nn.ReLU(),\n            )\n        )\n        # Add Device Parameter:\n        self._device = torch.nn.Parameter(torch.empty(0))\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 sizes in 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(self._activation)\n\n            conv_layer = DynEdgeConv(\n                torch.nn.Sequential(*layers),\n                aggr=\"max\", # aggr=\"add\",\n                nb_neighbors=self._nb_neighbours,\n                features_subset=self._features_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) + nb_input_features\n        )\n\n        post_processing_layers = []\n        layer_sizes = [nb_latent_features] + list(self._post_processing_layer_sizes)\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\n        self._post_processing = torch.nn.Sequential(*post_processing_layers)\n\n        nb_poolings = (1)\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(self._activation)\n\n        self._readout = torch.nn.Sequential(*readout_layers)\n\n    def _calculate_global_variables(self, x: torch.Tensor, edge_index: torch.LongTensor, batch: torch.LongTensor, *additional_attributes: torch.Tensor) -> torch.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        # Calculate mean features\n        global_means = scatter_mean(x, batch, dim=0)\n        # Add global variables\n        global_variables = torch.cat([global_means, h_x, h_y, h_z, h_t ] + [attr.unsqueeze(dim=1) for attr in additional_attributes], dim=1)\n        return global_variables\n\n    def forward(self, data: Data) -> torch.Tensor:\n        # Convenience variables\n        x, edge_index, batch = data.x, data.edge_index, data.batch\n        global_variables = self._calculate_global_variables(x, edge_index, batch, torch.log10(data.n_pulses))\n        # Distribute global variables out to each node\n        if not self._add_global_variables_after_pooling:\n            distribute = (batch.unsqueeze(dim=1) == torch.unique(batch).unsqueeze(dim=0)).type(torch.float)\n            global_variables_distributed = torch.sum(distribute.unsqueeze(dim=2) * global_variables.unsqueeze(dim=0), dim=1)\n            x = torch.cat((x, global_variables_distributed), dim=1)\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        # Skip-cat\n        x = torch.cat(skip_connections, dim=1)\n\n        # LSTMデータ構築部分:\n        # (q, m)を出力する\n\n        # Post-processing\n        x = self._post_processing(x)\n\n        x = self.attn(x, batch)\n        \n        x = self._readout(x)\n        \n        # return x\n        return x\n\n    @property\n    def device(self):\n        return self._device.device\n    \ndef calculate_xyzt_homophily(x: torch.Tensor, edge_index: torch.LongTensor, batch: Batch) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:\n    hx = torch_geometric.utils.homophily(edge_index, x[:, 0], batch).reshape(-1, 1)\n    hy = torch_geometric.utils.homophily(edge_index, x[:, 1], batch).reshape(-1, 1)\n    hz = torch_geometric.utils.homophily(edge_index, x[:, 2], batch).reshape(-1, 1)\n    ht = torch_geometric.utils.homophily(edge_index, x[:, 3], batch).reshape(-1, 1)\n    return hx, hy, hz, ht\n\nclass DynEdgeConv(EdgeConv):\n    \"\"\" Dynamical edge convolution layer \"\"\"\n    def __init__(self, nn: Callable, aggr: str = \"max\", nb_neighbors: int = 8, features_subset: Optional[Union[Sequence[int], slice]] = None, **kwargs: Any):\n        if features_subset is None:\n            features_subset = slice(None)\n        assert isinstance(features_subset, (list, slice))\n        super().__init__(nn=nn, aggr=aggr, **kwargs)\n        self.nb_neighbors = nb_neighbors\n        self.features_subset = features_subset\n        self._device = torch.nn.Parameter(torch.empty(0))\n\n    def forward(self, x: torch.Tensor, edge_index: Adj, batch: Optional[torch.Tensor] = None) -> torch.Tensor:\n        x = super().forward(x, edge_index)\n        edge_index = knn_graph(x=x[:, self.features_subset], k=self.nb_neighbors, batch=batch).to(self.device)\n        return x, edge_index\n\n    @property\n    def device(self):\n        return self._device.device\n\nclass KNNGraphBuilder(torch.nn.Module):\n    \"\"\"Builds graph from the k-nearest neighbours.\"\"\"\n    def __init__(self, nb_nearest_neighbours: int, columns: List[int] = None):\n        super().__init__()\n        if columns is None:\n            columns = [0, 1, 2]\n        self._nb_nearest_neighbours = nb_nearest_neighbours\n        self._columns = columns\n        self._device = torch.nn.Parameter(torch.empty(0))\n\n    def forward(self, data: Data) -> Data:\n        # NOTE: CPUとCUDAで結果が変わるので留意すること\n        data.edge_index = knn_graph(data.x[:, self._columns], k=self._nb_nearest_neighbours, batch=data.batch).to(self.device)\n        return data\n\n    @property\n    def device(self):\n        return self._device.device","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","papermill":{"duration":0.027664,"end_time":"2023-01-23T10:33:19.73899","exception":false,"start_time":"2023-01-23T10:33:19.711326","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-18T20:45:21.409100Z","iopub.execute_input":"2023-04-18T20:45:21.409533Z","iopub.status.idle":"2023-04-18T20:45:21.647370Z","shell.execute_reply.started":"2023-04-18T20:45:21.409493Z","shell.execute_reply":"2023-04-18T20:45:21.646159Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def vector_to_angle(vector):\n    vs = np.linalg.norm(vector[:, :3], axis=1)\n    vx = np.where(vs > 0, vector[:, 0] / vs, vector[:, 0])\n    vy = np.where(vs > 0, vector[:, 1] / vs, vector[:, 1])\n    vz = np.where(vs > 0, vector[:, 2] / vs, vector[:, 2])\n    vz = np.clip(vz, -1, 1)\n    ze = np.arccos(vz)\n    az = np.arctan2(vy, vx)\n    az = np.where(az < 0, az + 2.0 * np.pi, az)\n    return az, ze\n\ndef vector_to_angle_tensor(vector):\n    vs = torch.norm(vector[:, :3], dim=1)\n    vx = torch.where(vs > 0, vector[:, 0] / vs, vector[:, 0])\n    vy = torch.where(vs > 0, vector[:, 1] / vs, vector[:, 1])\n    vz = torch.where(vs > 0, vector[:, 2] / vs, vector[:, 2])\n    vz = torch.clip(vz, -1, 1)\n    ze = torch.arccos(vz)\n    az = torch.arctan2(vy, vx)\n    az = torch.where(az < 0, az + 2.0 * torch.pi, az)\n    return az, ze","metadata":{"execution":{"iopub.status.busy":"2023-04-18T20:45:21.648746Z","iopub.execute_input":"2023-04-18T20:45:21.649567Z","iopub.status.idle":"2023-04-18T20:45:21.660818Z","shell.execute_reply.started":"2023-04-18T20:45:21.649527Z","shell.execute_reply":"2023-04-18T20:45:21.659443Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#\n# TTA\n#\nclass TTAOperator(torch.nn.Module):\n    def __init__(self, angle, device):\n        super().__init__()\n        self.device = device\n        self.fwdmat = self.build_matrix( (angle * np.pi / 180.0))\n        self.invmat = self.build_matrix(-(angle * np.pi / 180.0))\n\n    def build_matrix(self, theta):\n        return torch.tensor([\n            [ np.cos(theta), -np.sin(theta), 0 ],\n            [ np.sin(theta),  np.cos(theta), 0 ],\n            [             0,              0, 1 ],\n        ], dtype=torch.float32).unsqueeze(0).to(self.device)\n    \n    def forward(self, dat):\n        buf = dat.detach().clone()\n        buf.x[:, :3] = torch.matmul(dat.x[:, :3], self.fwdmat)\n        return buf\n    \n    def inverse(self, x):\n        y = torch.matmul(x[:, :3], self.invmat)\n        return y\n\nclass NOPOperator(torch.nn.Module):\n    def __init__(self, angle, device):\n        super().__init__()\n        self.device = device\n    \n    def forward(self, dat):\n        buf = dat.detach().clone()\n        return buf\n\n    def inverse(self, x):\n        return x[:, :3]\n\nclass MagOperator(torch.nn.Module):\n    def __init__(self, angle, device):\n        super().__init__()\n        self.device = device\n\n    def forward(self, dat):\n        buf = dat.detach().clone()\n        return buf\n    \n    def inverse(self, x):\n        y = x[:, 3].unsqueeze(1) * x[:, :3]\n        return y","metadata":{"execution":{"iopub.status.busy":"2023-04-18T20:45:21.662548Z","iopub.execute_input":"2023-04-18T20:45:21.662963Z","iopub.status.idle":"2023-04-18T20:45:21.676647Z","shell.execute_reply.started":"2023-04-18T20:45:21.662924Z","shell.execute_reply":"2023-04-18T20:45:21.675639Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_sensor_info():\n    df = pd.read_csv(CFG.sensor_file).astype({\n        'sensor_id' : np.int16,\n        'x' : np.float32,\n        'y' : np.float32,\n        'z' : np.float32,\n    })    \n    return df","metadata":{"execution":{"iopub.status.busy":"2023-04-18T20:45:21.678282Z","iopub.execute_input":"2023-04-18T20:45:21.678674Z","iopub.status.idle":"2023-04-18T20:45:21.688743Z","shell.execute_reply.started":"2023-04-18T20:45:21.678635Z","shell.execute_reply":"2023-04-18T20:45:21.687797Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ICDataset(Dataset):\n    def __init__(self, metadata, event_ids, nbatch, sensor):\n        super().__init__()\n        self.event_ids = event_ids\n        self.metadata = metadata\n        self.nbatch = nbatch\n        self.sensor = sensor\n\n    def len(self):\n        return len(self.event_ids)\n\n    def get(self, idx):\n        eid = self.event_ids[idx]\n        edt = self.nbatch.loc[eid]\n        edt = pd.merge(edt, self.sensor , on='sensor_id', how='left')\n        edt[\"auxiliary\"] = edt[\"auxiliary\"].astype(int)        \n        # Downsample the large events\n        dat1 = downsample(edt, eid, 300)\n        dat2 = downsample(edt, eid, 375)\n        dat3 = downsample(edt, eid, 450)\n        dat4 = downsample(edt, eid, 600)\n        return dat1, dat2, dat3, dat4\n\ndef downsample(df, eid, pulse_limit):\n    edt = df.copy()\n    # Downsample the large events\n    count = len(edt)\n    if len(edt) > pulse_limit:\n        edt.reset_index(inplace=True, drop=True)\n        ids = edt[edt['auxiliary'] == 0].index.tolist()\n        if len(ids) >= pulse_limit:\n            ids = list(np.random.choice(ids, pulse_limit))\n        else:\n            num = pulse_limit - len(ids)\n            lst = edt[edt['auxiliary'] != 0].index.tolist()\n            ids = ids + list(np.random.choice(lst, num))\n        edt = edt.iloc[ids, :]\n        edt = edt.sort_values('time')\n    x = edt[[ 'x', 'y', 'z', 'time', 'charge', 'auxiliary' ]].values\n    x = torch.tensor(x, dtype=torch.float32)\n    data = Data(x=x, event_id=eid, n_pulses=torch.tensor(x.shape[0], dtype=torch.int32))\n    return data","metadata":{"execution":{"iopub.status.busy":"2023-04-18T20:45:21.690507Z","iopub.execute_input":"2023-04-18T20:45:21.690894Z","iopub.status.idle":"2023-04-18T20:45:21.703956Z","shell.execute_reply.started":"2023-04-18T20:45:21.690850Z","shell.execute_reply":"2023-04-18T20:45:21.702844Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_parquet(filepath):\n    df = pd.read_parquet(filepath)\n    return df","metadata":{"execution":{"iopub.status.busy":"2023-04-18T20:45:21.705977Z","iopub.execute_input":"2023-04-18T20:45:21.706620Z","iopub.status.idle":"2023-04-18T20:45:21.716438Z","shell.execute_reply.started":"2023-04-18T20:45:21.706423Z","shell.execute_reply":"2023-04-18T20:45:21.715385Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_loader(batch_id, metadata, sensor):\n    infer_data = load_parquet(os.path.join(CFG.datadir, f'batch_{batch_id}.parquet'))\n    infer_event_ids = infer_data.index.unique()\n    # infer_event_ids = infer_data['event_id'].unique()\n    infer_loader = DataLoader(\n        ICDataset(metadata, infer_event_ids, infer_data, sensor),\n        shuffle=False,\n        drop_last=False,\n        batch_size=CFG.batch_size,\n        num_workers=CFG.num_workers,\n    )\n    return infer_loader","metadata":{"execution":{"iopub.status.busy":"2023-04-18T20:45:21.717988Z","iopub.execute_input":"2023-04-18T20:45:21.718393Z","iopub.status.idle":"2023-04-18T20:45:21.726607Z","shell.execute_reply.started":"2023-04-18T20:45:21.718317Z","shell.execute_reply":"2023-04-18T20:45:21.725763Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def submit_writer_header():\n    with open('submission.csv', mode='w') as fw:\n        fw.write('event_id,azimuth,zenith\\n')\n\ndef submit_writer(ev_list, az_list, ze_list):\n    assert len(ev_list) == len(az_list)\n    assert len(ev_list) == len(ze_list)\n    if not os.path.exists('submission.csv'):\n        with open('submission.csv', mode='w') as fw:\n            fw.write('event_id,azimuth,zenith\\n')\n    with open('submission.csv', mode='a') as fw:\n        for e, az, ze in zip(ev_list, az_list, ze_list):\n            fw.write('{},{},{}\\n'.format(e, az, ze))","metadata":{"execution":{"iopub.status.busy":"2023-04-18T20:45:21.727985Z","iopub.execute_input":"2023-04-18T20:45:21.728677Z","iopub.status.idle":"2023-04-18T20:45:21.740530Z","shell.execute_reply.started":"2023-04-18T20:45:21.728639Z","shell.execute_reply":"2023-04-18T20:45:21.739731Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def inference(models1, models2, models3, models4, tta_operator):\n    submit_writer_header()\n    sens = get_sensor_info()\n    file = 'train_meta.parquet' if DEBUG else 'test_meta.parquet'\n    meta = pd.read_parquet(os.path.join(CFG.rootdir, file), columns=[ 'batch_id', 'event_id' ]).astype({ 'batch_id' : 'int16', 'event_id' : 'int64' })\n    for nb in meta.iloc[:, 0].unique():\n        ev_list = meta.loc[meta['batch_id'] == nb, 'event_id'].unique()\n        az_list = np.array([], dtype=np.float32)\n        ze_list = np.array([], dtype=np.float32)\n        loader = get_loader(nb, meta, sens)\n        for itr in loader:\n            w, x, y, z = itr\n            # pulses = x.n_pulses.numpy()\n            w = w.to(CFG.device)\n            x = x.to(CFG.device)\n            y = y.to(CFG.device)\n            z = z.to(CFG.device)\n            v = None\n            for op in tta_operator:\n                # m1\n                for mdl, wts in models1:\n                    q = op.inverse(mdl(op.forward(w)))\n                    v = (v + wts * q) if v is not None else (wts * q)\n                # boost\n                for mdl, wts in models2:\n                    q = op.inverse(mdl(op.forward(x)))\n                    v = (v + wts * q) if v is not None else (wts * q)\n                for mdl, wts in models3:\n                    q = op.inverse(mdl(op.forward(y)))\n                    v = (v + wts * q) if v is not None else (wts * q)\n                for mdl, wts in models4:\n                    q = op.inverse(mdl(op.forward(z)))\n                    v = (v + wts * q) if v is not None else (wts * q)\n            az_y, ze_y = vector_to_angle(v.cpu().numpy().reshape(-1, 3))\n            az_list = np.append(az_list, az_y)\n            ze_list = np.append(ze_list, ze_y)\n        submit_writer(ev_list, az_list, ze_list)\n        del loader, ev_list, az_list, ze_list\n        gc.collect()\n        if DEBUG:\n            break","metadata":{"execution":{"iopub.status.busy":"2023-04-18T20:45:21.741879Z","iopub.execute_input":"2023-04-18T20:45:21.742541Z","iopub.status.idle":"2023-04-18T20:45:21.754516Z","shell.execute_reply.started":"2023-04-18T20:45:21.742506Z","shell.execute_reply":"2023-04-18T20:45:21.753521Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# pulse=300\nmodelgroups1 = [ ]\n# pulse=375\nmodelgroups2 = [ ]\n# pulse=450\nmodelgroups3 = [ ]\n# pulse=600\nmodelgroups4 = [\n    [ ICM_12034,\n        [\n            ( 1.0, '/kaggle/input/icm-model/train_12034_lp17_ep144.pth'),\n        ]\n    ],\n]","metadata":{"execution":{"iopub.status.busy":"2023-04-18T20:45:21.755977Z","iopub.execute_input":"2023-04-18T20:45:21.756439Z","iopub.status.idle":"2023-04-18T20:45:21.767657Z","shell.execute_reply.started":"2023-04-18T20:45:21.756403Z","shell.execute_reply":"2023-04-18T20:45:21.766655Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Model:\nmodels1 = [ ]\nfor modelfunc, modellist in modelgroups1:\n    for w, pth in modellist:\n        m = modelfunc()\n        m.load_state_dict(torch.load(pth, map_location='cpu'), strict=True)\n        m = m.to(CFG.device)\n        m.eval()\n        models1.append((m, w))\nmodels2 = [ ]\nfor modelfunc, modellist in modelgroups2:\n    for w, pth in modellist:\n        m = modelfunc()\n        m.load_state_dict(torch.load(pth, map_location='cpu'), strict=True)\n        m = m.to(CFG.device)\n        m.eval()\n        models2.append((m, w))\nmodels3 = [ ]\nfor modelfunc, modellist in modelgroups3:\n    for w, pth in modellist:\n        m = modelfunc()\n        m.load_state_dict(torch.load(pth, map_location='cpu'), strict=True)\n        m = m.to(CFG.device)\n        m.eval()\n        models3.append((m, w))\nmodels4 = [ ]\nfor modelfunc, modellist in modelgroups4:\n    for w, pth in modellist:\n        m = modelfunc()\n        m.load_state_dict(torch.load(pth, map_location='cpu'), strict=True)\n        m = m.to(CFG.device)\n        m.eval()\n        models4.append((m, w))\n# TTA:\ntta_operator = [ MagOperator(0, CFG.device) ]","metadata":{"execution":{"iopub.status.busy":"2023-04-18T20:45:21.769118Z","iopub.execute_input":"2023-04-18T20:45:21.769626Z","iopub.status.idle":"2023-04-18T20:45:24.844410Z","shell.execute_reply.started":"2023-04-18T20:45:21.769589Z","shell.execute_reply":"2023-04-18T20:45:24.843383Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with torch.no_grad():\n    inference(models1, models2, models3, models4, tta_operator)","metadata":{"execution":{"iopub.status.busy":"2023-04-18T20:45:24.845727Z","iopub.execute_input":"2023-04-18T20:45:24.846079Z","iopub.status.idle":"2023-04-18T20:46:19.371827Z","shell.execute_reply.started":"2023-04-18T20:45:24.846042Z","shell.execute_reply":"2023-04-18T20:46:19.370052Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if SUBMIT:\n    df = pd.read_csv('submission.csv')\n    print(df.head())","metadata":{"papermill":{"duration":1.364722,"end_time":"2023-01-23T11:51:22.460904","exception":false,"start_time":"2023-01-23T11:51:21.096182","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-18T20:46:19.373221Z","iopub.status.idle":"2023-04-18T20:46:19.374306Z","shell.execute_reply.started":"2023-04-18T20:46:19.374006Z","shell.execute_reply":"2023-04-18T20:46:19.374042Z"},"trusted":true},"execution_count":null,"outputs":[]}]}