{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"try:\n    import polars as pls\nexcept:\n    print('Installing polars, please wait about 35 seconds...')\n    !pip install /kaggle/input/polars01516/polars-0.15.16-cp37-abi3-manylinux_2_17_x86_64.manylinux2014_x86_64.whl\n    import polars as pls","metadata":{"execution":{"iopub.status.busy":"2023-03-27T08:03:10.026943Z","iopub.execute_input":"2023-03-27T08:03:10.027287Z","iopub.status.idle":"2023-03-27T08:03:42.672560Z","shell.execute_reply.started":"2023-03-27T08:03:10.027209Z","shell.execute_reply":"2023-03-27T08:03:42.671521Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport random\nimport os\nimport gc\nimport bisect\nimport threading\nfrom concurrent.futures import ProcessPoolExecutor\nfrom tqdm import tqdm\nimport time\nfrom joblib import Parallel, delayed\nimport multiprocessing as mp\nfrom multiprocessing import get_context, set_start_method\nfrom sklearn.preprocessing import RobustScaler\nfrom scipy.interpolate import interp1d\nimport torch\n\nimport getpass\nfrom pathlib import Path\nfrom typing import Any, Callable, List, Optional, Sequence, Tuple, Union\n\nimport numpy as np\nimport pyarrow as pa\nimport pytorch_lightning as pl\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom pytorch_lightning import Trainer\nfrom pytorch_lightning.loggers import TensorBoardLogger, CSVLogger\nfrom pytorch_lightning.callbacks import ModelCheckpoint\n\n\nfrom scipy.interpolate import interp1d\nfrom sklearn.preprocessing import RobustScaler\nfrom torch import LongTensor, Tensor\nfrom torch.utils.data import DataLoader, Dataset\nfrom tqdm import tqdm\n\nimport sys\nimport subprocess\nsys.path.append(\"/kaggle/input/graphnet/graphnet-main/src\")\n\nwhls = [\n    \"/kaggle/input/pytorchgeometric/torch_cluster-1.6.0-cp37-cp37m-linux_x86_64.whl\",\n    \"/kaggle/input/pytorchgeometric/torch_scatter-2.1.0-cp37-cp37m-linux_x86_64.whl\",\n    \"/kaggle/input/pytorchgeometric/torch_sparse-0.6.16-cp37-cp37m-linux_x86_64.whl\",\n    \"/kaggle/input/pytorchgeometric/torch_spline_conv-1.2.1-cp37-cp37m-linux_x86_64.whl\",\n    \"/kaggle/input/pytorchgeometric/torch_geometric-2.2.0-py3-none-any.whl\",\n    \"/kaggle/input/pytorchgeometric/ruamel.yaml-0.17.21-py3-none-any.whl\",\n]\n\nfor w in whls:\n    print(\"Installing\", w)\n    subprocess.call([\"pip\", \"install\", w, \"--no-deps\", \"--upgrade\"])\n        \nfrom graphnet.models.graph_builders import KNNGraphBuilder\nfrom graphnet.models.task.reconstruction import DirectionReconstructionWithKappa\nfrom graphnet.training.loss_functions import VonMisesFisher2DLoss, VonMisesFisher3DLoss\nfrom graphnet.training.callbacks import PiecewiseLinearLR\nfrom torch_geometric.data import Data, Dataset\nfrom torch_geometric.loader import DataLoader\nfrom graphnet.models.gnn.gnn import GNN\nfrom graphnet.models.utils import calculate_xyzt_homophily\nfrom graphnet.utilities.config import save_model_config\nfrom torch_geometric.data import Data\nfrom torch_geometric.nn import EdgeConv\nfrom torch_geometric.nn.pool import knn_graph\nfrom torch_geometric.typing import Adj\nfrom torch_scatter import scatter_max, scatter_mean, scatter_min, scatter_sum","metadata":{"execution":{"iopub.status.busy":"2023-03-27T08:03:42.675471Z","iopub.execute_input":"2023-03-27T08:03:42.676541Z","iopub.status.idle":"2023-03-27T08:05:57.151253Z","shell.execute_reply.started":"2023-03-27T08:03:42.676499Z","shell.execute_reply":"2023-03-27T08:05:57.150204Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    VER = 2\n    OUTPUT_PATH = f\"/content/drive/MyDrive/kaggle/IceCube/{VER}/\"\n    EPOCHS = 10\n    BATCH_SIZE = 512\n    lr = 1e-3\n    T_max = 1000\n    weight_decay = 0.1\n    nb_inputs = 11\n    nearest_neighbours = 8\n\ninput_dir = '/kaggle/input/icecube-neutrinos-in-deep-ice'\ntransparency_dir = '/kaggle/input/icecubetransparency/ice_transparency.txt'","metadata":{"execution":{"iopub.status.busy":"2023-03-27T08:05:57.155009Z","iopub.execute_input":"2023-03-27T08:05:57.157712Z","iopub.status.idle":"2023-03-27T08:05:57.164142Z","shell.execute_reply.started":"2023-03-27T08:05:57.157679Z","shell.execute_reply":"2023-03-27T08:05:57.162300Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_dir = 'data'\n!mkdir data","metadata":{"execution":{"iopub.status.busy":"2023-03-27T08:05:57.167317Z","iopub.execute_input":"2023-03-27T08:05:57.167711Z","iopub.status.idle":"2023-03-27T08:05:58.166088Z","shell.execute_reply.started":"2023-03-27T08:05:57.167685Z","shell.execute_reply":"2023-03-27T08:05:58.164804Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def ice_transparency(data_path, datum=1950):\n    # Data from page 31 of https://arxiv.org/pdf/1301.5361.pdf\n    # Datum is from footnote 8 of page 29\n    df = pls.read_csv(data_path, sep = ' ')\n    \n    z_norm =  (df.get_column('depth').to_numpy()-datum)/500\n    \n    # \"scattering_len_norm\", \"absorption_len_norm\"\n    df = RobustScaler().fit_transform(\n        df.select([\n            pls.col(\"scattering_len\"),\n            pls.col(\"absorption_len\")\n        ]).to_numpy()\n    )\n\n    # These are both roughly equivalent after scaling\n    f_scattering = interp1d(z_norm, df[:, 0])\n    f_absorption = interp1d(z_norm, df[:, 1])\n    return f_scattering, f_absorption","metadata":{"execution":{"iopub.status.busy":"2023-03-27T08:05:58.168272Z","iopub.execute_input":"2023-03-27T08:05:58.168701Z","iopub.status.idle":"2023-03-27T08:05:58.176168Z","shell.execute_reply.started":"2023-03-27T08:05:58.168660Z","shell.execute_reply":"2023-03-27T08:05:58.175152Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"c_const = 0.299792458\nc_ice = 0.229\nt_valid_length = 6199.700247193777","metadata":{"execution":{"iopub.status.busy":"2023-03-27T08:05:58.177616Z","iopub.execute_input":"2023-03-27T08:05:58.178611Z","iopub.status.idle":"2023-03-27T08:05:58.184982Z","shell.execute_reply.started":"2023-03-27T08:05:58.178573Z","shell.execute_reply":"2023-03-27T08:05:58.183820Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# X_dataの特徴量は、x, y, z, time, charge, auxiliary, valid_time_window, qe, string, scatter, absorption\n# y_dataの値は、target_x, target_y, target_z\ndef read_data(meta, batch_idx, mode):\n    \"\"\"\n    mode : \"train\" or \"test\"\n    batch_idx: 使うbatch_idx\n    \"\"\"\n    \n    # [\"event_id\", \"time\", \"sensor_id\", \"charge\", \"auxiliary\"]を読み込み\n    X_data = (pls.scan_parquet(f\"{input_dir}/{mode}/batch_{batch_idx}.parquet\")\n              .with_columns([\n                  pls.col('event_id').cast(pls.Int64, strict=False),\n                  pls.col('sensor_id').cast(pls.Int32, strict=False),\n                  pls.when(pls.col(\"auxiliary\")==False).then(1).otherwise(0).cast(pls.Float32, strict=False).alias(\"auxiliary\")\n              ])).with_row_count('row_count')\n    \n    # パルスの数は1イベント当たり200個まで\n    # valid_time_window特徴量\n    peak = X_data.select((pls.col('row_count').min()+pls.col('charge').arg_max()).over('event_id').alias('row_count'))\n    peak = peak.join(X_data.select(['row_count', 'time']), on='row_count', how='left').rename({'time':'peak'}).collect()\n    X_data = (X_data.collect().with_columns([\n        pls.all(),\n        *peak.select(['peak'])\n    ]).lazy().with_columns([\n         pls.when((pls.col('peak')-pls.lit(t_valid_length)<=pls.col('time'))\n                 & (pls.col('time')<=pls.col('peak')+pls.lit(t_valid_length)))\n         .then(1.0).otherwise(0.0).alias('valid_time_window')\n        ]).drop(['row_count', 'peak'])\n              .sort([pls.col('auxiliary'), pls.col('valid_time_window'), pls.col('charge')], reverse=True)\n              .filter(pls.arange(0, pls.count()).over(\"event_id\") < 200)\n              .sort(['event_id', 'time']))\n    \n    del peak\n    gc.collect()\n    \n    # time, chargeを変換\n    X_data = X_data.with_columns([\n        (pls.col(\"time\")-1.0e04)/3.0e4,\n        pls.col(\"charge\").log10()/3.0\n    ]).collect().lazy()\n    \n    # [\"sensor_id\", \"x\", \"y\", \"z\", \"string\", \"qe\"]を読み込み\n    geometry_table = pls.scan_csv('/kaggle/input/sensor-features/sensors.csv').select([\"sensor_id\", \"x\", \"y\", \"z\", \"string\", \"qe\"])\n    geometry_table = geometry_table.with_columns([\n        pls.col('sensor_id').cast(pls.Int32, strict=False),\n        pls.col(\"x\")/500,\n        pls.col(\"y\")/500,\n        pls.col(\"z\")/500,\n        pls.col(\"string\")/10,\n        (pls.col(\"qe\")-1.25)/0.25\n    ])\n    \n    # sensor_idに応じて、それに対応するx,y,zをgeometry_tableから抽出\n    X_data = X_data.join(geometry_table, on=\"sensor_id\", how=\"left\")\n    \n    del geometry_table\n    gc.collect()\n    \n    # ice transparency\n    f_scattering, f_absorption = ice_transparency(transparency_dir)\n    X_data = X_data.collect()\n    X_data = X_data.with_columns([\n        pls.Series('scatter', f_scattering(X_data.get_column('z').to_numpy())),\n        pls.Series('absorption', f_absorption(X_data.get_column('z').to_numpy())),\n    ]).lazy()\n    \n    # 行番号を追加\n    X_data = X_data.with_row_count('row_number')\n\n    # eventごとに最初のパルスと最後のパルスのindexを取り出す\n    index = X_data.groupby('event_id').agg([\n        pls.col('row_number').min().alias('row_number_min'),\n        pls.col('row_number').max().alias('row_number_max')\n    ]).select(['row_number_min', 'row_number_max']).collect()\n\n    # X_dataからrow_number,event_id, sensor_idを削除する\n    X_data = X_data.drop(['row_number', 'event_id', 'sensor_id'])\n    \n    # データの順番を整える\n    X_data = X_data.select([pls.col('x'), pls.col('y'), pls.col('z'), pls.col('time'), pls.exclude(['x', 'y', 'z', 'time'])]).collect()\n    \n    # データを保存\n#     np.save(f'{data_dir}/{mode}_X_{batch_idx}.npy', X_data.to_numpy())\n    X_data.write_parquet(f'{data_dir}/{mode}_X_{batch_idx}.parquet')\n    index.write_parquet(f'{data_dir}/{mode}_index_{batch_idx}.parquet')\n    \n    \n    if mode=='train':\n        # [\"azimuth\", \"zenith\"]を読み込み\n        y_data = (meta.lazy().filter(pls.col(\"batch_id\")==batch_idx)\n                  .select([\"azimuth\", \"zenith\"])\n                  .with_columns([\n                      pls.col('azimuth').cast(pls.Float32, strict=False),\n                      pls.col('zenith').cast(pls.Float32, strict=False),\n        ]))\n\n        # 極座標から直交座標に変換\n        y_data = y_data.with_columns([\n            (pls.col('zenith').sin() * pls.col('azimuth').cos()).alias('target_x'),\n            (pls.col('zenith').sin() * pls.col('azimuth').sin()).alias('target_y'),\n            pls.col('zenith').cos().alias('target_z')\n        ]).drop(['azimuth', 'zenith']).collect()\n        \n        # データを保存\n        y_data.write_parquet(f'{data_dir}/{mode}_y_{batch_idx}.parquet')\n        \n        del y_data\n\n    del X_data, index\n    gc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-03-27T08:05:58.186856Z","iopub.execute_input":"2023-03-27T08:05:58.187792Z","iopub.status.idle":"2023-03-27T08:05:58.211468Z","shell.execute_reply.started":"2023-03-27T08:05:58.187692Z","shell.execute_reply":"2023-03-27T08:05:58.210575Z"},"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\n_dtype = {\n    \"batch_id\": \"int16\",\n    \"event_id\": \"int64\",\n    \"azimuth\" : \"float32\",\n    \"zenith\" : \"float32\"\n}\n\ndef to_cartesian(target):\n    return torch.sin(target[:, 1])*torch.cos(target[:, 0]), torch.sin(target[:, 1])*torch.sin(target[:, 0]), torch.cos(target[:, 1])\n\n# def to_polar(target):\n#     cos_phi = target[:, 0]/torch.sqrt(target[:, 0]**2+target[:, 1]**2+1e-10)\n#     cos_theta = target[:, 2]/torch.sqrt(target[:, 0]**2+target[:, 1]**2+target[:, 2]**2+1e-10)\n    \n#     cos_phi = torch.clamp(cos_phi, -1, 1)\n#     cos_theta = torch.clamp(cos_theta, -1, 1)\n    \n#     return torch.acos(cos_phi), torch.sgn(target[:, 1])*torch.acos(cos_theta)\n\nclass IceCubeModel(pl.LightningModule):\n    def __init__(\n        self,\n        data_len: int = 200000,\n        max_epochs: int = 10, # for the learning scheduler\n        nb_inputs: int = CFG.nb_inputs,\n        nearest_neighbours: int = CFG.nearest_neighbours,\n        T_max: int = CFG.T_max,\n        lr: float = CFG.lr,\n        weight_decay = CFG.weight_decay,\n        **kwargs,\n    ):\n        super().__init__()\n        self.save_hyperparameters()\n\n        self.loss_fn = VonMisesFisher3DLoss()\n        \n        self.model = DynEdge(\n            nb_inputs=nb_inputs,\n            nb_neighbours=nearest_neighbours,\n            global_pooling_schemes=[\"min\", \"max\", \"mean\", \"sum\"],\n            features_subset=slice(0, 4),  # NN search using xyzt\n        )\n        self.fc1 = nn.Linear(self.model.nb_outputs, self.model.nb_outputs)\n        self.fc2 = nn.Linear(self.model.nb_outputs, self.model.nb_outputs)\n        \n        self.direction_task = DirectionReconstructionWithKappa(\n            hidden_size=self.model.nb_outputs,\n            target_labels=\"direction\",\n            loss_function=VonMisesFisher3DLoss(),\n        )\n        self.direction_task1 = DirectionReconstructionWithKappa(\n            hidden_size=self.model.nb_outputs,\n            target_labels=\"direction\",\n            loss_function=VonMisesFisher3DLoss(),\n        )\n        self.direction_task2 = DirectionReconstructionWithKappa(\n            hidden_size=self.model.nb_outputs,\n            target_labels=\"direction\",\n            loss_function=VonMisesFisher3DLoss(),\n        )\n        self.training_step_outputs = []\n        self.validation_step_outputs = []\n\n    def forward(self, x):\n        emb = self.model(x)\n        emb1 = self.fc1(emb)\n        emb2 = self.fc2(emb)\n\n        # emb = self.norm(emb)\n        out1 = self.direction_task1(emb1)\n        out2 = self.direction_task2(emb2)\n\n        return out1, out2\n\n    def training_step(self, batch, batch_idx):\n        pred1, pred2 = self.forward(batch)\n        target1 = batch.y\n        target2 = target1*-1\n        loss = (self.loss_fn(pred1, target1) + self.loss_fn(pred2, target2))/2\n        self.training_step_outputs.append(loss)\n        \n        return {\"loss\": loss}\n\n    def on_train_epoch_end(self):\n        loss = torch.stack(self.training_step_outputs).mean()\n        self.log('loss', loss)\n        self.training_step_outputs.clear()\n\n    def validation_step(self, batch, batch_idx):\n        pred1, pred2 = self.forward(batch)\n        target = batch.y\n        pred2[:, :3] = pred2[:, :3]*-1\n        pred = (pred1 + pred2)/2\n        loss = self.loss_fn(pred, target)\n\n        metric = angular_dist_score(pred, target)\n\n        output = {\n            \"val_loss\": loss,\n            \"metric\": metric,\n        }\n        self.validation_step_outputs.append(output)\n\n        return output\n\n    def on_validation_epoch_end(self):\n        val_loss = torch.stack([x[\"val_loss\"] for x in self.validation_step_outputs]).mean()\n        metric = torch.stack([x[\"metric\"] for x in self.validation_step_outputs]).mean()\n        self.log_dict({'val_loss':val_loss, 'metric':metric}, prog_bar=True)\n        self.validation_step_outputs.clear()\n    \n    def test_step(self, batch, batch_idx):\n        pred = self.forward(batch)\n        return pred\n\n    def configure_optimizers(self):\n        parameters = add_weight_decay(\n            self,\n            self.hparams.weight_decay,\n            skip_list=[\"bias\", \"LayerNorm.bias\"],  # , \"LayerNorm.weight\"],\n        )\n\n        opt = torch.optim.Adam(self.parameters(), lr=self.hparams.lr)\n        sch = PiecewiseLinearLR(\n            optimizer = opt,\n            milestones = [\n                0,\n                self.hparams.data_len / 2,\n                self.hparams.data_len * self.hparams.max_epochs,\n            ],\n            factors=[1e-02, 1, 1e-02]\n        )\n        # sch = get_cosine_schedule_with_warmup(\n        #     opt,\n        #     num_warmup_steps=int(0 * self.hparams.T_max),\n        #     num_training_steps=self.hparams.T_max,\n        #     num_cycles=0.5,  # 1,\n        #     last_epoch=-1,\n        # )\n        \n        return {\n            \"optimizer\": opt,\n            \"lr_scheduler\": {\"scheduler\": sch, \"interval\": \"step\"},\n        }\n\ndef add_weight_decay(\n    model,\n    weight_decay=1e-5,\n    skip_list=(\"bias\", \"bn\", \"LayerNorm.bias\", \"LayerNorm.weight\"),\n):\n    decay = []\n    no_decay = []\n    for name, param in model.named_parameters():\n        if not param.requires_grad:\n            continue  # frozen weights\n        if len(param.shape) == 1 or name.endswith(\".bias\") or name in skip_list:\n            no_decay.append(param)\n        else:\n            decay.append(param)\n    return [\n        {\"params\": no_decay, \"weight_decay\": 0.0},\n        {\"params\": decay, \"weight_decay\": weight_decay},\n    ]\n\n\nclass DynEdgeConv(EdgeConv, pl.LightningModule):\n    \"\"\"Dynamical edge convolution layer.\"\"\"\n\n    def __init__(\n        self,\n        nn: Callable,\n        aggr: str = \"max\",\n        nb_neighbors: int = 8,\n        features_subset: Optional[Union[Sequence[int], slice]] = None,\n        **kwargs: Any,\n    ):\n        \"\"\"Construct `DynEdgeConv`.\n        Args:\n            nn: The MLP/torch.Module to be used within the `EdgeConv`.\n            aggr: Aggregation method to be used with `EdgeConv`.\n            nb_neighbors: Number of neighbours to be clustered after the\n                `EdgeConv` operation.\n            features_subset: Subset of features in `Data.x` that should be used\n                when dynamically performing the new graph clustering after the\n                `EdgeConv` operation. Defaults to all features.\n            **kwargs: Additional features to be passed to `EdgeConv`.\n        \"\"\"\n        # Check(s)\n        if features_subset is None:\n            features_subset = slice(None)  # Use all features\n        assert isinstance(features_subset, (list, slice))\n\n        # Base class constructor\n        super().__init__(nn=nn, aggr=aggr, **kwargs)\n\n        # Additional member variables\n        self.nb_neighbors = nb_neighbors\n        self.features_subset = features_subset\n\n    def forward(\n        self, x: Tensor, edge_index: Adj, batch: Optional[Tensor] = None\n    ) -> Tensor:\n\n        \"\"\"Forward pass.\"\"\"\n        # Standard EdgeConv forward pass\n        x = super().forward(x, edge_index)\n\n        # Recompute adjacency\n        edge_index = knn_graph(\n            x=x[:, self.features_subset],\n            k=self.nb_neighbors,\n            batch=batch,\n        ).to(self.device)\n\n        return x, edge_index\n\n\nclass DynEdge(GNN):\n    \"\"\"DynEdge (dynamical edge convolutional) model.\"\"\"\n\n    @save_model_config\n    def __init__(\n        self,\n        nb_inputs: int,\n        *,\n        nb_neighbours: int = 8,\n        features_subset: Optional[Union[List[int], slice]] = None,\n        dynedge_layer_sizes: Optional[List[Tuple[int, ...]]] = None,\n        post_processing_layer_sizes: Optional[List[int]] = None,\n        readout_layer_sizes: Optional[List[int]] = None,\n        global_pooling_schemes: Optional[Union[str, List[str]]] = None,\n        add_global_variables_after_pooling: bool = False,\n    ):\n        \"\"\"Construct `DynEdge`.\n        Args:\n            nb_inputs: Number of input features on each node.\n            nb_neighbours: Number of neighbours to used in the k-nearest\n                neighbour clustering which is performed after each (dynamical)\n                edge convolution.\n            features_subset: The subset of latent features on each node that\n                are used as metric dimensions when performing the k-nearest\n                neighbours clustering. Defaults to [0,1,2].\n            dynedge_layer_sizes: The layer sizes, or latent feature dimenions,\n                used in the `DynEdgeConv` layer. Each entry in\n                `dynedge_layer_sizes` corresponds to a single `DynEdgeConv`\n                layer; the integers in the corresponding tuple corresponds to\n                the layer sizes in the multi-layer perceptron (MLP) that is\n                applied within each `DynEdgeConv` layer. That is, a list of\n                size-two tuples means that all `DynEdgeConv` layers contain a\n                two-layer MLP.\n                Defaults to [(128, 256), (336, 256), (336, 256), (336, 256)].\n            post_processing_layer_sizes: Hidden layer sizes in the MLP\n                following the skip-concatenation of the outputs of each\n                `DynEdgeConv` layer. Defaults to [336, 256].\n            readout_layer_sizes: Hidden layer sizes in the MLP following the\n                post-processing _and_ optional global pooling. As this is the\n                last layer(s) in the model, the last layer in the read-out\n                yields the output of the `DynEdge` model. Defaults to [128,].\n            global_pooling_schemes: The list global pooling schemes to use.\n                Options are: \"min\", \"max\", \"mean\", and \"sum\".\n            add_global_variables_after_pooling: Whether to add global variables\n                after global pooling. The alternative is to  added (distribute)\n                them to the individual nodes before any convolutional\n                operations.\n        \"\"\"\n        # Latent feature subset for computing nearest neighbours in DynEdge.\n        if features_subset is None:\n            features_subset = slice(0, 3)\n\n        # DynEdge layer sizes\n        if dynedge_layer_sizes is None:\n            dynedge_layer_sizes = [\n                (\n                    128,\n                    256,\n                ),\n                (\n                    336,\n                    256,\n                ),\n                (\n                    336,\n                    256,\n                ),\n                (\n                    336,\n                    256,\n                ),\n            ]\n\n        assert isinstance(dynedge_layer_sizes, list)\n        assert len(dynedge_layer_sizes)\n        assert all(isinstance(sizes, tuple) for sizes in dynedge_layer_sizes)\n        assert all(len(sizes) > 0 for sizes in dynedge_layer_sizes)\n        assert all(all(size > 0 for size in sizes) for sizes in dynedge_layer_sizes)\n\n        self._dynedge_layer_sizes = dynedge_layer_sizes\n\n        # Post-processing layer sizes\n        if post_processing_layer_sizes is None:\n            post_processing_layer_sizes = [\n                336,\n                256,\n            ]\n\n        assert isinstance(post_processing_layer_sizes, list)\n        assert len(post_processing_layer_sizes)\n        assert all(size > 0 for size in post_processing_layer_sizes)\n\n        self._post_processing_layer_sizes = post_processing_layer_sizes\n\n        # Read-out layer sizes\n        if readout_layer_sizes is None:\n            readout_layer_sizes = [\n                128,\n            ]\n\n        assert isinstance(readout_layer_sizes, list)\n        assert len(readout_layer_sizes)\n        assert all(size > 0 for size in readout_layer_sizes)\n\n        self._readout_layer_sizes = readout_layer_sizes\n\n        # Global pooling scheme(s)\n        if isinstance(global_pooling_schemes, str):\n            global_pooling_schemes = [global_pooling_schemes]\n\n        if isinstance(global_pooling_schemes, list):\n            for pooling_scheme in global_pooling_schemes:\n                assert (\n                    pooling_scheme in GLOBAL_POOLINGS\n                ), f\"Global pooling scheme {pooling_scheme} not supported.\"\n        else:\n            assert global_pooling_schemes is None\n\n        self._global_pooling_schemes = global_pooling_schemes\n\n        if add_global_variables_after_pooling:\n            assert self._global_pooling_schemes, (\n                \"No global pooling schemes were request, so cannot add global\"\n                \" variables after pooling.\"\n            )\n        self._add_global_variables_after_pooling = add_global_variables_after_pooling\n\n        # Base class constructor\n        super().__init__(nb_inputs, self._readout_layer_sizes[-1])\n\n        # Remaining member variables()\n        self._activation = torch.nn.GELU()\n        self._nb_inputs = nb_inputs\n        self._nb_global_variables = 5 + nb_inputs\n        self._nb_neighbours = nb_neighbours\n        self._features_subset = features_subset\n\n        self._construct_layers()\n\n    def _construct_layers(self) -> None:\n        \"\"\"Construct layers (torch.nn.Modules).\"\"\"\n        # Convolutional operations\n        nb_input_features = self._nb_inputs\n        if not self._add_global_variables_after_pooling:\n            nb_input_features += self._nb_global_variables\n\n        self._conv_layers = torch.nn.ModuleList()\n        nb_latent_features = nb_input_features\n        for 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(nn.BatchNorm1d(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(nn.BatchNorm1d(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_latent_features = nb_out * nb_poolings\n        if self._add_global_variables_after_pooling:\n            nb_latent_features += self._nb_global_variables\n\n        readout_layers = []\n        layer_sizes = [nb_latent_features] + list(self._readout_layer_sizes)\n        for nb_in, nb_out in zip(layer_sizes[:-1], layer_sizes[1:]):\n            readout_layers.append(torch.nn.Linear(nb_in, nb_out))\n            readout_layers.append(nn.BatchNorm1d(nb_out))\n            readout_layers.append(self._activation)\n\n        self._readout = torch.nn.Sequential(*readout_layers)\n\n    def _global_pooling(self, x: Tensor, batch: LongTensor) -> Tensor:\n        \"\"\"Perform global pooling.\"\"\"\n        assert self._global_pooling_schemes\n        pooled = []\n        for pooling_scheme in self._global_pooling_schemes:\n            pooling_fn = GLOBAL_POOLINGS[pooling_scheme]\n            pooled_x = pooling_fn(x, index=batch, dim=0)\n            if isinstance(pooled_x, tuple) and len(pooled_x) == 2:\n                # `scatter_{min,max}`, which return also an argument, vs.\n                # `scatter_{mean,sum}`\n                pooled_x, _ = pooled_x\n            pooled.append(pooled_x)\n\n        return torch.cat(pooled, dim=1)\n\n    def _calculate_global_variables(\n        self,\n        x: Tensor,\n        edge_index: LongTensor,\n        batch: LongTensor,\n        *additional_attributes: Tensor,\n    ) -> Tensor:\n        \"\"\"Calculate global variables.\"\"\"\n        # Calculate homophily (scalar variables)\n        h_x, h_y, h_z, h_t = calculate_xyzt_homophily(x, edge_index, batch)\n\n        # Calculate mean features\n        global_means = scatter_mean(x, batch, dim=0)\n\n        # Add global variables\n        global_variables = torch.cat(\n            [\n                global_means,\n                h_x,\n                h_y,\n                h_z,\n                h_t,\n            ]\n            + [attr.unsqueeze(dim=1) for attr in additional_attributes],\n            dim=1,\n        )\n\n        return global_variables\n\n    def forward(self, data: Data) -> Tensor:\n        \"\"\"Apply learnable forward pass.\"\"\"\n        # Convenience variables\n        x, edge_index, batch = data.x, data.edge_index, data.batch\n\n        global_variables = self._calculate_global_variables(\n            x,\n            edge_index,\n            batch,\n            torch.log10(data.n_pulses),\n        )\n\n        # Distribute global variables out to each node\n        if not self._add_global_variables_after_pooling:\n            distribute = (\n                batch.unsqueeze(dim=1) == torch.unique(batch).unsqueeze(dim=0)\n            ).type(torch.float)\n\n            global_variables_distributed = torch.sum(\n                distribute.unsqueeze(dim=2) * global_variables.unsqueeze(dim=0),\n                dim=1,\n            )\n\n            x = torch.cat((x, global_variables_distributed), dim=1)\n\n        # DynEdge-convolutions\n        skip_connections = [x]\n        for conv_layer in self._conv_layers:\n            x, edge_index = conv_layer(x, edge_index, batch)\n            skip_connections.append(x)\n\n        # Skip-cat\n        x = torch.cat(skip_connections, dim=1)\n\n        # Post-processing\n        x = self._post_processing(x)\n\n        # (Optional) Global pooling\n        if self._global_pooling_schemes:\n            x = self._global_pooling(x, batch=batch)\n            if self._add_global_variables_after_pooling:\n                x = torch.cat(\n                    [\n                        x,\n                        global_variables,\n                    ],\n                    dim=1,\n                )\n\n        # Read-out\n        x = self._readout(x)\n\n        return x\n\n\n\nclass DataManager():\n    X_data = {}\n    y_data = {}\n    index = {}\n    count = {}\n    mode = 'train'\n    data_size = {}\n    \n    def __init__(self, mode, use_flip=False):\n        self.mode = mode\n        self.use_flip = use_flip\n        \n    def get(self, i, j):        \n        \n        if i not in self.count:\n            print(f'created batch_{i} data.')\n            self.X_data = {}\n            self.y_data = {}\n            self.index = {}\n            self.count = {}\n            self.data_size = {}\n#             gc.collect()  \n            self.X_data[i] = pa.parquet.read_table(f'{data_dir}/test_X_{i}.parquet').to_pandas().values\n#             self.X_data[i] = np.load(f'{data_dir}/train_X_{i}.npy')\n            \n            if self.mode=='train':\n                self.y_data[i] = pa.parquet.read_table(f'{data_dir}/train_y_{i}.parquet').to_pandas().values\n                \n            self.index[i] = pa.parquet.read_table(f'{data_dir}/test_index_{i}.parquet').to_pandas().values\n            self.data_size[i] = len(self.index[i])\n            self.count[i] = 0\n        \n        self.count[i] += 1\n        \n\n\n        # event毎のパルスのデータ\n        x = self.X_data[i][self.index[i][j, 0]:self.index[i][j, 1]+1]\n        \n        # pca = PCA(n_components=1)\n        # transformed = pca.fit_transform(x[:, :3])\n        # x = np.hstack([x, transformed])\n        # transformed = pca.fit_transform(x[:, :4])\n        # x = np.hstack([x, pca.inverse_transform(transformed)])\n        if np.random.uniform(0, 1) < 0.5 and self.use_flip:\n            flip = True\n        else:\n            flip = False\n        if flip:\n            x[:, :2] = x[:, :2]*-1\n        x = torch.tensor(x, dtype=torch.float32)\n        \n        if self.mode=='train':\n            # event毎のtargetのデータ\n            y = np.array([self.y_data[i][j, :]])\n            if flip:\n                y[:, :2] = y[:, :2]*-1\n            y = torch.tensor(y, dtype=torch.float32)\n        \n            data = Data(x=x, y=y, n_pulses=torch.tensor(x.shape[0], dtype=torch.int32))\n        else:\n            data = Data(x=x, n_pulses=torch.tensor(x.shape[0], dtype=torch.int32))\n        \n#         if self.count[i]==self.data_size[i]:\n#             print(f'deleted batch_{i} data.')\n#             self.X_data.pop(i)\n            \n#             if self.mode=='train':\n#                 self.y_data.pop(i)\n                \n#             self.index.pop(i)\n#             self.count.pop(i)\n#             self.data_size.pop(i)\n#             gc.collect()  \n        \n        return data\n       \nclass IceCubeDataset(Dataset):\n    def __init__(\n        self,\n        manager,\n        idx,\n        count,\n        pulse_limit=300,\n        transform=None,\n        pre_transform=None,\n        pre_filter=None,\n    ):\n        super().__init__(transform, pre_transform, pre_filter)\n        \n        self.manager = manager\n        self.count = count\n        self.length = np.sum(count)\n        self.idx = idx\n        batch = np.arange(len(self.idx))\n\n        p = 0\n        j = 0\n        self.batch_idx = np.zeros((self.length, 2)).astype(np.int32)\n        for i in range(self.length):\n            if j >= self.count[batch[p]]:\n                p += 1\n                j = 0\n            self.batch_idx[i, 0] = self.idx[batch[p]]\n            self.batch_idx[i, 1] = j\n            j += 1\n\n    def len(self):\n        \"\"\"\n        eventの数を返す\n        \"\"\"\n        return self.length\n\n    def get(self, idx):\n        i, j = self.batch_idx[idx, 0], self.batch_idx[idx, 1]\n        return self.manager.get(i, j)\n\nclass ValidDataset(Dataset):\n    def __init__(\n        self,\n        manager,\n        idx,\n        count,\n        pulse_limit=300,\n        transform=None,\n        pre_transform=None,\n        pre_filter=None,\n    ):\n        super().__init__(transform, pre_transform, pre_filter)\n        \n        self.manager = manager\n        self.count = count\n        self.length = np.sum(count)\n        self.idx = idx\n        batch = np.arange(len(self.idx))\n\n        p = 0\n        j = 0\n        self.batch_idx = np.zeros((self.length, 2)).astype(np.int32)\n        for i in range(self.length):\n            if j >= self.count[batch[p]]:\n                p += 1\n                j = 0\n            self.batch_idx[i, 0] = self.idx[batch[p]]\n            self.batch_idx[i, 1] = j\n            j += 1\n\n    def len(self):\n        \"\"\"\n        eventの数を返す\n        \"\"\"\n        return self.length\n\n    def get(self, idx):\n        i, j = self.batch_idx[idx, 0], self.batch_idx[idx, 1]\n        return self.manager.get(i, j)","metadata":{"execution":{"iopub.status.busy":"2023-03-27T08:09:54.297135Z","iopub.execute_input":"2023-03-27T08:09:54.297604Z","iopub.status.idle":"2023-03-27T08:09:54.379734Z","shell.execute_reply.started":"2023-03-27T08:09:54.297569Z","shell.execute_reply":"2023-03-27T08:09:54.378652Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# meta = (pls.scan_parquet(f'{input_dir}/test_meta.parquet')\n#         .select([\"batch_id\", \"azimuth\", \"zenith\"])\n#         .select(pls.all().shrink_dtype()).collect())\n# partitions = meta.partition_by('batch_id')","metadata":{"execution":{"iopub.status.busy":"2023-03-27T08:09:54.797245Z","iopub.execute_input":"2023-03-27T08:09:54.797592Z","iopub.status.idle":"2023-03-27T08:09:54.803781Z","shell.execute_reply.started":"2023-03-27T08:09:54.797564Z","shell.execute_reply":"2023-03-27T08:09:54.802511Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# partition_len = 660//6\n# chunk_size = partition_len//mp.cpu_count()\n# partition_chunks = [partitions[i:i+chunk_size] for i in range(0, partition_len, chunk_size)]","metadata":{"execution":{"iopub.status.busy":"2023-03-27T08:09:55.148558Z","iopub.execute_input":"2023-03-27T08:09:55.148941Z","iopub.status.idle":"2023-03-27T08:09:55.153830Z","shell.execute_reply.started":"2023-03-27T08:09:55.148906Z","shell.execute_reply":"2023-03-27T08:09:55.152747Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def read_data_wrapper(args):\n    meta, idx = args\n    for i, j in tqdm(zip(meta, idx), total=len(idx)):\n        read_data(i, j, 'test')","metadata":{"execution":{"iopub.status.busy":"2023-03-27T08:09:55.696954Z","iopub.execute_input":"2023-03-27T08:09:55.697654Z","iopub.status.idle":"2023-03-27T08:09:55.702874Z","shell.execute_reply.started":"2023-03-27T08:09:55.697619Z","shell.execute_reply":"2023-03-27T08:09:55.701865Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"meta = pls.read_parquet(f'{input_dir}/test_meta.parquet')\nmeta = (meta.lazy().groupby('batch_id')\n        .agg(pls.count().alias('count')).sort('batch_id').collect())\n\nbatch_ids = meta.select(pls.col(\"batch_id\").unique())\nmeta","metadata":{"execution":{"iopub.status.busy":"2023-03-27T08:09:56.264740Z","iopub.execute_input":"2023-03-27T08:09:56.266195Z","iopub.status.idle":"2023-03-27T08:09:56.283663Z","shell.execute_reply.started":"2023-03-27T08:09:56.266147Z","shell.execute_reply":"2023-03-27T08:09:56.282586Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import pandas as pd\n# pd.read_parquet('/kaggle/input/icecube-neutrinos-in-deep-ice/sample_submission.parquet')","metadata":{"execution":{"iopub.status.busy":"2023-03-27T08:09:57.428874Z","iopub.execute_input":"2023-03-27T08:09:57.429581Z","iopub.status.idle":"2023-03-27T08:09:57.433871Z","shell.execute_reply.started":"2023-03-27T08:09:57.429545Z","shell.execute_reply":"2023-03-27T08:09:57.432868Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# idx = np.arange(1,partition_len+1)\n# value = [(partition_chunks[i], idx[j:j+chunk_size]) for i, j in enumerate(range(0, partition_len, chunk_size))]\n# idx = np.arange(1, 110+1)\n\n# for i in tqdm(idx):\nfor i in batch_ids.to_numpy().reshape(-1):\n    read_data(None, i, mode = 'test')\n# Parallel(n_jobs=mp.cpu_count())(\n#     delayed(read_data_wrapper)(chunk) for chunk in value\n# )","metadata":{"execution":{"iopub.status.busy":"2023-03-27T08:09:58.052955Z","iopub.execute_input":"2023-03-27T08:09:58.053313Z","iopub.status.idle":"2023-03-27T08:09:58.642626Z","shell.execute_reply.started":"2023-03-27T08:09:58.053285Z","shell.execute_reply":"2023-03-27T08:09:58.641646Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def to_polar(target):\n    cos_phi = target[:, 0]/np.sqrt(target[:, 0]**2+target[:, 1]**2+1e-10)\n    cos_theta = target[:, 2]/np.sqrt(target[:, 0]**2+target[:, 1]**2+target[:, 2]**2+1e-10)\n    \n    cos_phi = np.clip(cos_phi, -1, 1)\n    cos_theta = np.clip(cos_theta, -1, 1)\n    \n    phi =  np.where(target[:, 1]>=0, 1, -1)*np.arccos(cos_phi)\n    theta = np.arccos(cos_theta)\n    \n    phi = np.where(phi>=0, phi, phi+2*np.pi)\n    \n    return phi, theta\n\ndef inference(paths, val_idx, val_count):\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    pre_transform = KNNGraphBuilder(nb_nearest_neighbours = CFG.nearest_neighbours)\n    \n    manager = DataManager(mode = 'test', use_flip=False)\n\n#     dfs = []\n#     for idx in val_idx:\n#         df = pa.parquet.read_table(f'{data_dir}/train_y_{idx}.parquet').to_pandas()\n#         dfs.append(df)\n#     df = pd.concat(dfs).reset_index(drop=True)\n    \n    val_dataset = ValidDataset(\n         manager=manager, idx=val_idx, count=val_count, pre_transform=pre_transform\n    )\n    val_loader = DataLoader(val_dataset, batch_size=CFG.BATCH_SIZE, num_workers=1)\n    model = IceCubeModel(max_epochs = CFG.EPOCHS+10)\n\n    xss = []\n    yss = []\n    zss = []\n    for path in paths:\n        model.load_state_dict(torch.load(path, map_location=torch.device('cpu'))['state_dict'])\n        model.eval()\n        model.to(device)\n\n        xs = []\n        ys = []\n        zs = []\n        for inputs in val_loader:\n            inputs.to(device)\n            with torch.no_grad():\n                # y_preds = model.forward(inputs).to('cpu').numpy()\n                pred1, pred2 = model.forward(inputs)\n                pred1, pred2 = pred1.to('cpu').numpy(), pred2.to('cpu').numpy()\n                pred2[:, :3] = pred2[:, :3]*-1\n                y_preds = (pred1 + pred2)/2\n    #             y_preds = pred1 + pred2\n                # y_preds = to_polar(y_preds)\n                xs.append(y_preds[:, 0])\n                ys.append(y_preds[:, 1])\n                zs.append(y_preds[:, 2])\n        xs = np.concatenate(xs).reshape(-1, 1)\n        ys = np.concatenate(ys).reshape(-1, 1)\n        zs = np.concatenate(zs).reshape(-1, 1)\n        xss.append(xs)\n        yss.append(ys)\n        zss.append(zs)\n    xs = np.mean(xss, axis=0)\n    ys = np.mean(yss, axis=0)\n    zs = np.mean(zss, axis=0)\n    \n    azimuth, zenith = to_polar(np.concatenate([xs, ys, zs], axis=1))\n    \n#     df['pred_x'] = xs\n#     df['pred_y'] = ys\n#     df['pred_z'] = zs\n#     df.to_csv(CFG.OUTPUT_PATH + f'oof_{fold}.csv', index=False)\n\n#     del val_count, df\n    gc.collect()\n    return azimuth, zenith\n    ","metadata":{"execution":{"iopub.status.busy":"2023-03-27T08:09:59.896053Z","iopub.execute_input":"2023-03-27T08:09:59.896473Z","iopub.status.idle":"2023-03-27T08:09:59.912340Z","shell.execute_reply.started":"2023-03-27T08:09:59.896441Z","shell.execute_reply":"2023-03-27T08:09:59.911331Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\nsubmission = pd.read_parquet('/kaggle/input/icecube-neutrinos-in-deep-ice/sample_submission.parquet')\nval_count = (meta.lazy().filter(pls.col('batch_id').is_in(pls.Series(batch_ids.to_numpy().reshape(-1))))\n       .with_column(pls.col('count')).collect()\n       .get_column('count').to_numpy().astype(np.int32))\n\nazimuth, zenith = inference(['/kaggle/input/dynedge-model-weight/model_f0-val_loss1.4073.ckpt', '/kaggle/input/dynedge-model-weight/model_f2-val_loss1.4120.ckpt', '/kaggle/input/dynedge-model-weight/model_f3-val_loss1.4007.ckpt', '/kaggle/input/dynedge-model-weight/model_f1-val_loss1.4198.ckpt'], batch_ids.to_numpy().reshape(-1), val_count)","metadata":{"execution":{"iopub.status.busy":"2023-03-27T08:10:01.217254Z","iopub.execute_input":"2023-03-27T08:10:01.217613Z","iopub.status.idle":"2023-03-27T08:10:07.444910Z","shell.execute_reply.started":"2023-03-27T08:10:01.217583Z","shell.execute_reply":"2023-03-27T08:10:07.442855Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission['azimuth'] = azimuth\nsubmission['zenith'] = zenith\nsubmission = submission.sort_values(by = ['event_id'])\nsubmission","metadata":{"execution":{"iopub.status.busy":"2023-03-27T08:10:13.822914Z","iopub.execute_input":"2023-03-27T08:10:13.823294Z","iopub.status.idle":"2023-03-27T08:10:13.845751Z","shell.execute_reply.started":"2023-03-27T08:10:13.823264Z","shell.execute_reply":"2023-03-27T08:10:13.844859Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission.to_csv(\"submission.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2023-03-27T08:10:22.772377Z","iopub.execute_input":"2023-03-27T08:10:22.772730Z","iopub.status.idle":"2023-03-27T08:10:22.782740Z","shell.execute_reply.started":"2023-03-27T08:10:22.772701Z","shell.execute_reply":"2023-03-27T08:10:22.781691Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}