{"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":"","metadata":{"_uuid":"e85c38d1-5d54-48e9-8ebd-b9146334024a","_cell_guid":"955bd05e-0ad0-4eb9-b652-3fb6be591c63","jupyter":{"outputs_hidden":false}}},{"cell_type":"markdown","source":"# Update Packages\n\nTo have the same package as my local environment, we need to run the following commands in Kaggle to:\n1. Update pytorch from cuda 11.3 to 11.6.\n2. Install pytorch geometric and its depencencies.\n3. Update pyarrow from 5.0 to 11.0.","metadata":{"_uuid":"c2b2203b-d56e-4eff-910a-ef1bb8c03e8b","_cell_guid":"9f153194-ead5-4510-a121-710a85f2ee98","trusted":true}},{"cell_type":"code","source":"# !pip install -U torch torchvision torchaudio --extra-index-url https://download.pytorch.org/whl/cu116\n# !pip install pyg_lib torch_scatter torch_sparse torch_cluster torch_spline_conv torch_geometric -f https://data.pyg.org/whl/torch-1.13.0+cu116.html\n# !pip install pyarrow==11","metadata":{"_uuid":"75a1f8bc-b1db-4a9d-a268-77cf7856f038","_cell_guid":"acc382e1-79db-4589-85e9-b7766cc0d333","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-04-18T13:23:52.764462Z","iopub.execute_input":"2023-04-18T13:23:52.764811Z","iopub.status.idle":"2023-04-18T13:23:52.788284Z","shell.execute_reply.started":"2023-04-18T13:23:52.764727Z","shell.execute_reply":"2023-04-18T13:23:52.787362Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nimport getpass\nfrom pathlib import Path\nfrom typing import Any, Callable, List, Optional, Sequence, Tuple, Union\n\nimport numpy as np\nimport pandas as pd\nimport pytorch_lightning as pl\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\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\nfrom transformers import get_cosine_schedule_with_warmup\n\nKERNEL = False if getpass.getuser() == \"anjum\" else True\nCOMP_NAME = \"icecube-neutrinos-in-deep-ice\"\n\nif not KERNEL:\n    INPUT_PATH = Path(f\"/mnt/storage_dimm2/kaggle_data/{COMP_NAME}\")\n    OUTPUT_PATH = Path(f\"/mnt/storage_dimm2/kaggle_output/{COMP_NAME}\")\n    MODEL_CACHE = Path(\"/mnt/storage/model_cache/torch\")\n    #TRANSPARENCY_PATH = INPUT_PATH / \"ice_transparency.txt\"\nelse:\n    INPUT_PATH = Path(f\"/kaggle/input/{COMP_NAME}\")\n    MODEL_CACHE = None\n    #TRANSPARENCY_PATH = \"/kaggle/input/icecubetransparency/ice_transparency.txt\"\n\n    # Install packages\n    import subprocess\n\n    whls = [\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\n    for w in whls:\n        print(\"Installing\", w)\n        subprocess.call([\"pip\", \"install\", w, \"--no-deps\", \"--upgrade\"])\n\n    import sys\n    sys.path.append(\"/kaggle/input/graphnet/graphnet-main/src\")\n\n    \n# from graphnet.models.graph_builders import KNNGraphBuilder\n# from graphnet.models.task.reconstruction import (\n#     AzimuthReconstructionWithKappa,\n#     ZenithReconstruction,\n# )\n# from graphnet.training.loss_functions import VonMisesFisher2DLoss\n# from torch_geometric.data import Data, Dataset\n# from torch_geometric.loader import DataLoader\n# from graphnet.models.gnn.gnn import GNN\n# from graphnet.models.utils import calculate_xyzt_homophily\n# from graphnet.utilities.config import save_model_config\n# from torch_geometric.data import Data\n# from torch_geometric.nn import EdgeConv\n# from torch_geometric.nn.pool import knn_graph\n# from torch_geometric.typing import Adj\n# from torch_scatter import scatter_max, scatter_mean, scatter_min, scatter_sum","metadata":{"_uuid":"5562ef73-fd74-4185-812a-e6437d0ac027","_cell_guid":"bba3b4dd-8718-4a43-b25c-230e0593115b","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-04-18T13:23:52.790752Z","iopub.execute_input":"2023-04-18T13:23:52.79144Z","iopub.status.idle":"2023-04-18T13:26:11.448389Z","shell.execute_reply.started":"2023-04-18T13:23:52.791405Z","shell.execute_reply":"2023-04-18T13:26:11.44687Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Prepare\n\nThe code in my compter is structured as:\n\n```\n<root>\n├── data\n│   ├── test_meta.parquet\n│   ├── test_meta.parquet\n│   └── train\n│       ├── batch_1.parquet\n│       ├── ...\n│       └── batch_51.parquet\n├── meta\n│   ├── meta_1.parquet\n│   ├── ...\n│   └── meta_51.parquet\n├── state_dict.pth\n├── sensor_geometry.csv\n├── split_meta.py\n├── exp-gnn.py\n└── utils.py\n```\n\n* `meta` keeps the split meta. It should not exist now.\n* `state_dict.pth` is the baseline's pretrained weight.\n* `split_meta.py` is the script to split meta.\n* `exp-gnn.py` contains the code for this experiment.\n* `utils.py` contains functions shared by experiments.\n\nFollowing commands make this notebook to have same file structure.","metadata":{"_uuid":"ae1db0db-8394-4d46-9f4c-db79cf69a990","_cell_guid":"3caf9a9f-e24f-42e5-9309-ced45de226f1","trusted":true}},{"cell_type":"code","source":"!ln -fs /kaggle/input/icecube-neutrinos-in-deep-ice/ data\n!cp /kaggle/input/dynedge-pretrained/dynedge_pretrained_batch_1_to_50/state_dict.pth .\n!cp /kaggle/input/icecube-neutrinos-in-deep-ice/sensor_geometry.csv .\n!md5sum state_dict.pth","metadata":{"_uuid":"e23cf84d-a73d-4192-bee7-4fc269cd65ed","_cell_guid":"f7450e7f-3e02-4813-ac9a-a628dd886b8f","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-04-18T13:26:11.450856Z","iopub.execute_input":"2023-04-18T13:26:11.452355Z","iopub.status.idle":"2023-04-18T13:26:15.501494Z","shell.execute_reply.started":"2023-04-18T13:26:11.452301Z","shell.execute_reply":"2023-04-18T13:26:15.500303Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Split Meta\n\n`train_meta.parquet` is so large that I cannot load it in `Dataset`, so I split it to chunks just like `batch_X.parquet`.","metadata":{"_uuid":"b21adc50-66a5-4039-9000-265ff197774d","_cell_guid":"5c1f236f-58e3-4a49-a64d-83daad4bd221","trusted":true}},{"cell_type":"code","source":"%%writefile split_meta.py\n\nimport pyarrow\nimport pyarrow.parquet\nimport pandas as pd\nfrom tqdm import tqdm\nfrom pathlib import Path\n\n\nout_dir = Path('meta')\nout_dir.mkdir(parents=True, exist_ok=True)\n\nmeta = pyarrow.parquet.read_table('data/train_meta.parquet')\n#meta = pyarrow.parquet.read_table('data/test_meta.parquet')\nfor batch_id in tqdm(range(1, 660)):  #??? how sholud i know how much meta ?\n    group = meta.filter(pyarrow.compute.field('batch_id') == batch_id)\n    group = group.select(['event_id', 'azimuth', 'zenith'])\n    pyarrow.parquet.write_table(group, out_dir / f'meta_{batch_id}.parquet')\n\ndf = pd.read_parquet('data/test_meta.parquet')\ndf['azimuth'] = 0.0\ndf['zenith'] = 0.0\ndf.to_parquet(out_dir / 'meta_661.parquet')","metadata":{"_uuid":"6b77b145-28f4-4fb0-97cc-c744f81990ef","_cell_guid":"43a3edfc-223d-422c-a4b4-4fdfa2f8d228","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-04-18T13:26:15.505366Z","iopub.execute_input":"2023-04-18T13:26:15.506464Z","iopub.status.idle":"2023-04-18T13:26:15.51432Z","shell.execute_reply.started":"2023-04-18T13:26:15.506418Z","shell.execute_reply":"2023-04-18T13:26:15.513197Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#!python split_meta.py","metadata":{"_uuid":"9e7f4e50-09e4-41a2-9785-10f319c8acd4","_cell_guid":"402ce381-6f4a-494e-9992-64b5e4e1d32d","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-04-18T13:26:15.515722Z","iopub.execute_input":"2023-04-18T13:26:15.516159Z","iopub.status.idle":"2023-04-18T13:26:15.525148Z","shell.execute_reply.started":"2023-04-18T13:26:15.516107Z","shell.execute_reply":"2023-04-18T13:26:15.524001Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Utils.py\n\nIt contains the common functions shared by different experiments.\n\nThe `VonMisesFisher3DLoss` and `LogCMK` are extracted from Graphnet.","metadata":{"_uuid":"12692ef4-45c0-4306-899e-4bc81f6caa64","_cell_guid":"e8c5e516-6551-474c-8f37-bd175835e7ee","trusted":true}},{"cell_type":"code","source":"%%writefile utils.py\n\nimport torch\nfrom torch import nn\nfrom torch import Tensor\nimport numpy as np\nimport scipy\n\n\ndef angle_to_xyz(angles_b):\n    az, zen = angles_b.t()\n    x = torch.cos(az) * torch.sin(zen)\n    y = torch.sin(az) * torch.sin(zen)\n    z = torch.cos(zen)\n    return torch.stack([x, y, z], dim=1)\n\n\ndef xyz_to_angle(xyz_b):\n    x, y, z = xyz_b.t()\n    az = torch.arccos(x / torch.sqrt(x**2 + y**2)) * torch.sign(y)\n    zen = torch.arccos(z / torch.sqrt(x**2 + y**2 + z**2))\n    return torch.stack([az, zen], dim=1)\n\n\ndef angular_error(xyz_pred_b, xyz_true_b):\n    return torch.arccos(torch.sum(xyz_pred_b * xyz_true_b, dim=1))\n\n\nclass LogCMK(torch.autograd.Function):\n    \"\"\"MIT License.\n\n    Copyright (c) 2019 Max Ryabinin\n\n    Permission is hereby granted, free of charge, to any person obtaining a copy\n    of this software and associated documentation files (the \"Software\"), to deal\n    in the Software without restriction, including without limitation the rights\n    to use, copy, modify, merge, publish, distribute, sublicense, and/or sell\n    copies of the Software, and to permit persons to whom the Software is\n    furnished to do so, subject to the following conditions:\n\n    The above copyright notice and this permission notice shall be included in all\n    copies or substantial portions of the Software.\n\n    THE SOFTWARE IS PROVIDED \"AS IS\", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR\n    IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,\n    FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE\n    AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER\n    LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,\n    OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE\n    SOFTWARE.\n    _____________________\n\n    From [https://github.com/mryab/vmf_loss/blob/master/losses.py]\n    Modified to use modified Bessel function instead of exponentially scaled ditto\n    (i.e. `.ive` -> `.iv`) as indiciated in [1812.04616] in spite of suggestion in\n    Sec. 8.2 of this paper. The change has been validated through comparison with\n    exact calculations for `m=2` and `m=3` and found to yield the correct results.\n    \"\"\"\n\n    @staticmethod\n    def forward(ctx, m, kappa):  # pylint: disable=invalid-name,arguments-differ\n        \"\"\"Forward pass.\"\"\"\n        dtype = kappa.dtype\n        ctx.save_for_backward(kappa)\n        ctx.m = m\n        ctx.dtype = dtype\n        kappa = kappa.double()\n        iv = torch.from_numpy(scipy.special.iv(m / 2.0 - 1, kappa.cpu().numpy())).to(\n            kappa.device\n        )\n        return (\n            (m / 2.0 - 1) * torch.log(kappa)\n            - torch.log(iv)\n            - (m / 2) * np.log(2 * np.pi)\n        ).type(dtype)\n\n    @staticmethod\n    def backward(ctx, grad_output):  # pylint: disable=invalid-name,arguments-differ\n        \"\"\"Backward pass.\"\"\"\n        kappa = ctx.saved_tensors[0]\n        m = ctx.m\n        dtype = ctx.dtype\n        kappa = kappa.double().cpu().numpy()\n        grads = -(\n            (scipy.special.iv(m / 2.0, kappa)) / (scipy.special.iv(m / 2.0 - 1, kappa))\n        )\n        return (\n            None,\n            grad_output * torch.from_numpy(grads).to(grad_output.device).type(dtype),\n        )\n\n\nclass VonMisesFisher3DLoss(nn.Module):\n    \"\"\"General class for calculating von Mises-Fisher loss.\n\n    Requires implementation for specific dimension `m` in which the target and\n    prediction vectors need to be prepared.\n    \"\"\"\n\n    @classmethod\n    def log_cmk_exact(\n        cls, m: int, kappa: Tensor\n    ) -> Tensor:  # pylint: disable=invalid-name\n        \"\"\"Calculate $log C_{m}(k)$ term in von Mises-Fisher loss exactly.\"\"\"\n        return LogCMK.apply(m, kappa)\n\n    @classmethod\n    def log_cmk_approx(\n        cls, m: int, kappa: Tensor\n    ) -> Tensor:  # pylint: disable=invalid-name\n        \"\"\"Calculate $log C_{m}(k)$ term in von Mises-Fisher loss approx.\n\n        [https://arxiv.org/abs/1812.04616] Sec. 8.2 with additional minus sign.\n        \"\"\"\n        v = m / 2.0 - 0.5\n        a = torch.sqrt((v + 1) ** 2 + kappa**2)\n        b = v - 1\n        return -a + b * torch.log(b + a)\n\n    @classmethod\n    def log_cmk(\n        cls, m: int, kappa: Tensor, kappa_switch: float = 100.0\n    ) -> Tensor:  # pylint: disable=invalid-name\n        \"\"\"Calculate $log C_{m}(k)$ term in von Mises-Fisher loss.\n\n        Since `log_cmk_exact` is diverges for `kappa` >~ 700 (using float64\n        precision), and since `log_cmk_approx` is unaccurate for small `kappa`,\n        this method automatically switches between the two at `kappa_switch`,\n        ensuring continuity at this point.\n        \"\"\"\n        kappa_switch = torch.tensor([kappa_switch]).to(kappa.device)\n        mask_exact = kappa < kappa_switch\n\n        # Ensure continuity at `kappa_switch`\n        offset = cls.log_cmk_approx(m, kappa_switch) - cls.log_cmk_exact(\n            m, kappa_switch\n        )\n        ret = cls.log_cmk_approx(m, kappa) - offset\n        ret[mask_exact] = cls.log_cmk_exact(m, kappa[mask_exact])\n        return ret\n\n    def _evaluate(self, prediction: Tensor, target: Tensor) -> Tensor:\n        \"\"\"Calculate von Mises-Fisher loss for a vector in D dimensons.\n\n        This loss utilises the von Mises-Fisher distribution, which is a\n        probability distribution on the (D - 1) sphere in D-dimensional space.\n\n        Args:\n            prediction: Predicted vector, of shape [batch_size, D].\n            target: Target unit vector, of shape [batch_size, D].\n\n        Returns:\n            Elementwise von Mises-Fisher loss terms.\n        \"\"\"\n        # Check(s)\n        assert prediction.dim() == 2\n        assert target.dim() == 2\n        assert prediction.size() == target.size()\n\n        # Computing loss\n        m = target.size()[1]\n        k = torch.norm(prediction, dim=1)\n        dotprod = torch.sum(prediction * target, dim=1)\n        elements = -self.log_cmk(m, k) - dotprod\n        return elements\n\n    def forward(self, prediction: Tensor, target: Tensor) -> Tensor:\n        \"\"\"Calculate von Mises-Fisher loss for a direction in the 3D.\n\n        Args:\n            prediction: Output of the model. Must have shape [N, 4] where\n                columns 0, 1, 2 are predictions of `direction` and last column\n                is an estimate of `kappa`.\n            target: Target tensor, extracted from graph object.\n\n        Returns:\n            Elementwise von Mises-Fisher loss terms. Shape [N,]\n        \"\"\"\n        target = target.reshape(-1, 3)\n        # Check(s)\n        assert prediction.dim() == 2 and prediction.size()[1] == 4\n        assert target.dim() == 2\n        assert prediction.size()[0] == target.size()[0]\n\n        kappa = prediction[:, 3]\n        p = kappa.unsqueeze(1) * prediction[:, [0, 1, 2]]\n        return self._evaluate(p, target)","metadata":{"_uuid":"b400b495-347e-46b3-93c5-3dfdb2d3e21f","_cell_guid":"5548d51a-4c4a-4046-94f5-267f903c58ba","collapsed":false,"_kg_hide-input":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-04-18T13:26:15.526523Z","iopub.execute_input":"2023-04-18T13:26:15.526796Z","iopub.status.idle":"2023-04-18T13:26:15.539894Z","shell.execute_reply.started":"2023-04-18T13:26:15.526772Z","shell.execute_reply":"2023-04-18T13:26:15.538717Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Experiment Code\n\nThe code is split in 3 parts.\n\n## IceCube as [IterableDataset](https://pytorch.org/docs/stable/data.html#torch.utils.data.IterableDataset)\n\nThe dataset reads each chunk (`batch_X.parquet` and its corresponding `meta_X.parquet`) into memory iteratively instead of random accessing them. This reduces the memory consumption so I can train the model on home computer.\n\nThe dataset collates the samples to mini-batch itself, so when instantiating the `DataLoader`, please use `batch_size=1`,`num_worker=1` and `collate_fn=lambda x: x[0]`. \n\nEach sample in dataset is a Data instance and a mini-batch is a Batch instance. Please refer to [torch_geomeric](https://pytorch-geometric.readthedocs.io/en/latest/modules/data.html#data-objects) for details.\n\n## Model as [pl.LightningModule](https://lightning.ai/pages/open-source/)\n\nI extract and rewrite the code from graphnet. I make the code as clear as possible. Beaware that the input of `forward` is a [`Batch`](https://pytorch-geometric.readthedocs.io/en/latest/get_started/introduction.html#mini-batches), not a single sample `Data`.\n\n## Validation\n\nFinally, I let the model load weights from baseline's pretrained weight and perform validation on `batch_51.parquet`. The angular error is 1.02 indicating the `forward` implementation is roughly the same as baseline.","metadata":{"_uuid":"5230fd65-86a5-425e-b3d0-14f1d5a51d95","_cell_guid":"0a7fbb54-7107-4343-a1d3-d352138b684f","trusted":true}},{"cell_type":"code","source":"%%writefile exp-gnn.py\n\nimport gc\nimport random\nimport pandas as pd\nfrom pathlib import Path\nimport pyarrow.parquet\n\nimport torch\nfrom torch import nn\nfrom torch.nn import functional as F\nfrom torch.utils.data import IterableDataset, DataLoader\n\nimport pytorch_lightning as pl\nfrom torch_geometric.nn import knn_graph, EdgeConv\nfrom torch_geometric.utils import homophily\nfrom torch_geometric.data import Data, Batch\nfrom torch_scatter import scatter_add, scatter_mean, scatter_max, scatter_min\n\nfrom utils import angle_to_xyz, xyz_to_angle\nfrom utils import VonMisesFisher3DLoss, angular_error\n\nfrom typing import Any, Callable, List, Optional, Sequence, Tuple, Union\n\nimport numpy as np\nimport pandas as pd\n\nCOMP_NAME = \"icecube-neutrinos-in-deep-ice\"\nINPUT_PATH = Path(f\"/kaggle/input/{COMP_NAME}\")\nMODEL_CACHE = None\nTRANSPARENCY_PATH = \"/kaggle/input/icecubetransparency/ice_transparency.txt\"\n\n_dtype = {\n    \"batch_id\": \"int16\",\n    \"event_id\": \"int64\",\n}\n\n\nclass IceCube(IterableDataset):\n    def __init__(\n        self,\n        parquet_dir,\n        meta_dir,\n        chunk_ids,\n        batch_size=200,\n        max_pulses=200,\n        shuffle=False,\n    ):\n        self.parquet_dir = parquet_dir\n        self.meta_dir = meta_dir\n        self.chunk_ids = chunk_ids\n        self.batch_size = batch_size\n        self.max_pulses = max_pulses\n        self.shuffle = shuffle\n\n        if self.shuffle:\n            random.shuffle(self.chunk_ids)\n\n    def __iter__(self):\n        # Handle num_workers > 1 and multi-gpu\n        is_dist = torch.distributed.is_initialized()\n        world_size = torch.distributed.get_world_size() if is_dist else 1\n        rank_id = torch.distributed.get_rank() if is_dist else 0\n\n        info = torch.utils.data.get_worker_info()\n        num_worker = info.num_workers if info else 1\n        worker_id = info.id if info else 0\n\n        num_replica = world_size * num_worker\n        offset = rank_id * num_worker + worker_id\n        chunk_ids = self.chunk_ids[offset::num_replica]\n\n        # Sensor data\n        #sensor_xyz = pd.read_csv('sensor_geometry.csv')[['x', 'y', 'z']]\n        sensor_xyz = pd.read_csv(INPUT_PATH / \"sensor_geometry.csv\")[['x', 'y', 'z']]\n\n        \n        sensor_xyz = torch.from_numpy(sensor_xyz.values).float()\n\n        # Read each chunk and meta iteratively into memory and build mini-batch\n        for c, chunk_id in enumerate(chunk_ids):\n            data = pd.read_parquet(self.parquet_dir / f'batch_{chunk_id}.parquet')\n            #data =pd.read_parquet(INPUT_PATH / f\"batch_{chunk_id}.parquet\")\n            \n            meta = pd.read_parquet(self.meta_dir / f'meta_{chunk_id}.parquet')\n            #meta = pd.read_parquet(INPUT_PATH / f'meta_{chunk_id}.parquet')\n            angles = meta[['azimuth', 'zenith']].values\n            angles = torch.from_numpy(angles).float()\n            xyzs = angle_to_xyz(angles)\n            meta = {eid: xyz for eid, xyz in zip(meta['event_id'].tolist(), xyzs)}\n\n            # Take all eventi_ids and split them into batches\n            eids = list(meta.keys())\n            if self.shuffle:\n                random.shuffle(eids)\n            eids_batches = [\n                eids[i : i + self.batch_size]\n                for i in range(0, len(eids), self.batch_size)\n            ]\n\n            for batch_eids in eids_batches:\n                batch = []\n\n                # For each sample, extract features\n                for eid in batch_eids:\n                    df = data.loc[eid]\n                    if len(df) > self.max_pulses:\n                        df = df.sample(n=self.max_pulses)\n                    df = df.sort_values(['time'])\n                    t = torch.from_numpy(df['time'].values).float()\n                    c = torch.from_numpy(df['charge'].values).float()\n                    s = torch.from_numpy(df['sensor_id'].values).long()\n                    p = sensor_xyz[s]\n                    a = torch.from_numpy(df['auxiliary'].values).float()\n                    feat = torch.stack([p[:, 0], p[:, 1], p[:, 2], t, c, a], dim=1)\n\n                    batch.append(\n                        Data(\n                            x=feat,\n                            gt=meta[eid],\n                            n_pulses=len(feat),\n                            eid=torch.tensor([eid]).long(),\n                        )\n                    )\n\n                yield Batch.from_data_list(batch)\n\n            del data\n            del meta\n            gc.collect()\n\n            \nclass IceCubeV2(IterableDataset):\n    def __init__(\n        self,\n        #parquet_dir,\n        #meta_dir,\n        g_batchid,\n        g_event_ids,\n        mode,\n        batch_size,\n        #max_pulses=200,\n        max_pulses=300,\n        shuffle=False,\n    ):\n        #self.parquet_dir = parquet_dir\n        #self.meta_dir = meta_dir\n        self.g_batchid=g_batchid\n        self.g_event_ids=g_event_ids\n        #self.sensers=sensers\n        self.mode = mode\n        #self.chunk_ids = chunk_ids\n        self.batch_size = batch_size\n        self.max_pulses = max_pulses\n        self.shuffle = shuffle\n\n        #if self.shuffle:\n        #    random.shuffle(self.chunk_ids)\n\n    def __iter__(self):\n        # Handle num_workers > 1 and multi-gpu\n#         is_dist = torch.distributed.is_initialized()\n#         world_size = torch.distributed.get_world_size() if is_dist else 1\n#         rank_id = torch.distributed.get_rank() if is_dist else 0\n\n#         info = torch.utils.data.get_worker_info()\n#         num_worker = info.num_workers if info else 1\n#         worker_id = info.id if info else 0\n\n#         num_replica = world_size * num_worker\n#         offset = rank_id * num_worker + worker_id\n        #chunk_ids = self.chunk_ids[offset::num_replica]\n        chunk_ids = self.g_batchid\n        # Sensor data\n        #sensor_xyz = pd.read_csv('sensor_geometry.csv')[['x', 'y', 'z']]\n        sensor_xyz = pd.read_csv(INPUT_PATH / \"sensor_geometry.csv\")[['x', 'y', 'z']]\n\n        \n        sensor_xyz = torch.from_numpy(sensor_xyz.values).float()\n        \n        # Read each chunk and meta iteratively into memory and build mini-batch\n        #or c, chunk_id in enumerate(chunk_ids):\n        for chunk_id in [chunk_ids]:#661 ->[661]  stupid haha~\n            print(chunk_id)\n            #data = pd.read_parquet(self.parquet_dir / f'batch_{chunk_id}.parquet')\n            data = pd.read_parquet(INPUT_PATH / self.mode / f\"batch_{chunk_id}.parquet\")\n            #data =pd.read_parquet(INPUT_PATH / f\"batch_{chunk_id}.parquet\")\n            #print(data.shape)#32xxxxxxxx  very big\n            #meta = pd.read_parquet(self.meta_dir / f'meta_{chunk_id}.parquet')\n            #meta = pd.read_parquet(INPUT_PATH / f'meta_{chunk_id}.parquet')\n            #angles = meta[['azimuth', 'zenith']].values\n            #angles = torch.from_numpy(angles).float()\n            #xyzs = angle_to_xyz(angles)\n            #meta = {eid: xyz for eid, xyz in zip(meta['event_id'].tolist(), xyzs)}\n\n            # Take all eventi_ids and split them into batches\n            #eids = list(meta.keys())\n            eids = self.g_event_ids\n            #print(eids)\n            if self.shuffle:\n                random.shuffle(eids)\n            eids_batches = [\n                eids[i : i + self.batch_size]\n                for i in range(0, len(eids), self.batch_size)\n            ]\n\n            for batch_eids in eids_batches:\n                batch = []\n\n                # For each sample, extract features\n                for eid in batch_eids:\n                    df = data.loc[eid]\n                    if len(df) > self.max_pulses:\n                        df = df.sample(n=self.max_pulses)\n                    df = df.sort_values(['time'])\n                    t = torch.from_numpy(df['time'].values).float()\n                    c = torch.from_numpy(df['charge'].values).float()\n                    s = torch.from_numpy(df['sensor_id'].values).long()\n                    p = sensor_xyz[s]\n                    a = torch.from_numpy(df['auxiliary'].values).float()\n                    feat = torch.stack([p[:, 0], p[:, 1], p[:, 2], t, c, a], dim=1)\n\n                    batch.append(\n                        Data(\n                            x=feat,\n                            #gt=meta[eid],\n                            n_pulses=len(feat),\n                            eid=torch.tensor([eid]).long(),\n                        )\n                    )\n\n                yield Batch.from_data_list(batch)\n\n            del data\n            #del meta\n            gc.collect()\n\nclass IceCubeV3(IterableDataset):\n    def __init__(\n        self,\n        #parquet_dir,\n        #meta_dir,\n        g_batchid,\n        g_event_ids,\n        mode,\n        batch_size,\n        #max_pulses=200,\n        max_pulses=300,\n        shuffle=False,\n    ):\n        #self.parquet_dir = parquet_dir\n        #self.meta_dir = meta_dir\n        self.g_batchid=g_batchid\n        self.g_event_ids=g_event_ids\n        #self.sensers=sensers\n        self.mode = mode\n        #self.chunk_ids = chunk_ids\n        self.batch_size = batch_size\n        self.max_pulses = max_pulses\n        self.shuffle = shuffle\n\n        #if self.shuffle:\n        #    random.shuffle(self.chunk_ids)\n\n    def __iter__(self):\n\n        chunk_ids = self.g_batchid\n        # Sensor data\n        #sensor_xyz = pd.read_csv('sensor_geometry.csv')[['x', 'y', 'z']]\n        sensor_xyz = pd.read_csv(INPUT_PATH / \"sensor_geometry.csv\")[['x', 'y', 'z']]\n\n        \n        sensor_xyz = torch.from_numpy(sensor_xyz.values).float()\n        \n        # Read each chunk and meta iteratively into memory and build mini-batch\n        #or c, chunk_id in enumerate(chunk_ids):\n        for chunk_id in [chunk_ids]:#661 ->[661]  stupid haha~\n            print(chunk_id)\n            #data = pd.read_parquet(self.parquet_dir / f'batch_{chunk_id}.parquet')\n            data = pd.read_parquet(INPUT_PATH / self.mode / f\"batch_{chunk_id}.parquet\")\n\n\n            # Take all eventi_ids and split them into batches\n            #eids = list(meta.keys())\n            eids = self.g_event_ids\n            #print(eids)\n            if self.shuffle:\n                random.shuffle(eids)\n            eids_batches = [\n                eids[i : i + self.batch_size]\n                for i in range(0, len(eids), self.batch_size)\n            ]\n\n            for batch_eids in eids_batches:\n                batch = []\n\n                # For each sample, extract features\n                for eid in batch_eids:\n                    df = data.loc[eid]\n                    original_pulse_len = len(df) \n                    #if len(df) > self.max_pulses and self.shuffle:\n                    if len(df) > self.max_pulses:    \n                        df = df.sample(n=self.max_pulses)\n                    df = df.sort_values(['time'])\n                    t = torch.from_numpy(df['time'].values).float()\n                    c = torch.from_numpy(df['charge'].values).float()\n                    s = torch.from_numpy(df['sensor_id'].values).long()\n                    p = sensor_xyz[s]\n                    a = torch.from_numpy(df['auxiliary'].values).float()\n                    feat = torch.stack([p[:, 0], p[:, 1], p[:, 2], t, c, a], dim=1)\n\n                    batch.append(\n                        Data(\n                            x=feat,\n                            #gt=meta[eid],\n                            #n_pulses=len(feat),\n                            n_pulses=original_pulse_len,\n                            eid=torch.tensor([eid]).long(),\n                        )\n                    )\n\n                yield Batch.from_data_list(batch)\n\n            del data\n            #del meta\n            gc.collect()            \n            \n            \n\nclass MLP(nn.Sequential):\n    def __init__(self, feats):\n        layers = []\n        for i in range(1, len(feats)):\n            layers.append(nn.Linear(feats[i - 1], feats[i]))\n            layers.append(nn.LeakyReLU())\n        super().__init__(*layers)\n\n\nclass Model(pl.LightningModule):\n    def __init__(\n        self, max_lr=1e-3, min_lr=1e-5, num_warmup_step=1_000, num_total_step=20_000\n    ):\n        super().__init__()\n        self.save_hyperparameters()\n        self.conv0 = EdgeConv(MLP([34, 128, 256]), aggr='add')\n        self.conv1 = EdgeConv(MLP([512, 336, 256]), aggr='add')\n        self.conv2 = EdgeConv(MLP([512, 336, 256]), aggr='add')\n        self.conv3 = EdgeConv(MLP([512, 336, 256]), aggr='add')\n        self.post = MLP([1041, 336, 256])\n        self.readout = MLP([768, 128])\n        self.pred = nn.Linear(128, 3)\n\n    def forward(self, data: Batch):\n        vert_feat = data.x\n        batch = data.batch\n\n        vert_feat[:, 0] /= 500.0  # x\n        vert_feat[:, 1] /= 500.0  # y\n        vert_feat[:, 2] /= 500.0  # z\n        vert_feat[:, 3] = (vert_feat[:, 3] - 1.0e04) / 3.0e4  # time\n        vert_feat[:, 4] = torch.log10(vert_feat[:, 4]) / 3.0  # charge\n\n        edge_index = knn_graph(vert_feat[:, :3], 8, batch)\n\n        # Construct global features\n        hx = homophily(edge_index, vert_feat[:, 0], batch).reshape(-1, 1)\n        hy = homophily(edge_index, vert_feat[:, 1], batch).reshape(-1, 1)\n        hz = homophily(edge_index, vert_feat[:, 2], batch).reshape(-1, 1)\n        ht = homophily(edge_index, vert_feat[:, 3], batch).reshape(-1, 1)\n        means = scatter_mean(vert_feat, batch, dim=0)\n        n_p = torch.log10(data.n_pulses).reshape(-1, 1)\n        global_feats = torch.cat([means, hx, hy, hz, ht, n_p], dim=1)  # [B, 11]\n\n        # Distribute global_feats to each vertex\n        _, cnts = torch.unique_consecutive(batch, return_counts=True)\n        global_feats = torch.repeat_interleave(global_feats, cnts, dim=0)\n        vert_feat = torch.cat((vert_feat, global_feats), dim=1)\n\n        # Convolutions\n        feats = [vert_feat]\n        # Conv 0\n        vert_feat = self.conv0(vert_feat, edge_index)\n        feats.append(vert_feat)\n        # Conv 1\n        edge_index = knn_graph(vert_feat[:, :3], k=8, batch=batch)\n        vert_feat = self.conv1(vert_feat, edge_index)\n        feats.append(vert_feat)\n        # Conv 2\n        edge_index = knn_graph(vert_feat[:, :3], k=8, batch=batch)\n        vert_feat = self.conv2(vert_feat, edge_index)\n        feats.append(vert_feat)\n        # Conv 3\n        edge_index = knn_graph(vert_feat[:, :3], k=8, batch=batch)\n        vert_feat = self.conv3(vert_feat, edge_index)\n        feats.append(vert_feat)\n\n        # Postprocessing\n        post_inp = torch.cat(feats, dim=1)\n        post_out = self.post(post_inp)\n\n        # Readout\n        readout_inp = torch.cat(\n            [\n                scatter_min(post_out, batch, dim=0)[0],\n                scatter_max(post_out, batch, dim=0)[0],\n                scatter_mean(post_out, batch, dim=0),\n            ],\n            dim=1,\n        )\n        readout_out = self.readout(readout_inp)\n\n        # Predict\n        pred = self.pred(readout_out)\n        kappa = pred.norm(dim=1, p=2) + 1e-8\n        pred_x = pred[:, 0] / kappa\n        pred_y = pred[:, 1] / kappa\n        pred_z = pred[:, 2] / kappa\n        pred = torch.stack([pred_x, pred_y, pred_z, kappa], dim=1)\n\n        return pred\n\n    def train_or_valid_step(self, data, prefix):\n        pred_xyzk = self.forward(data)  # [B, 4]\n        true_xyz = data.gt.view(-1, 3)  # [B, 3]\n        loss = VonMisesFisher3DLoss()(pred_xyzk, true_xyz).mean()\n        error = angular_error(pred_xyzk[:, :3], true_xyz).mean()\n        self.log(f'loss/{prefix}', loss, batch_size=len(true_xyz))\n        self.log(f'error/{prefix}', error, batch_size=len(true_xyz))\n        return loss\n\n    def training_step(self, data, _):\n        return self.train_or_valid_step(data, 'train')\n\n    def validation_step(self, data, _):\n        self.train_or_valid_step(data, 'valid')\n\n    def configure_optimizers(self):\n        optimizer = torch.optim.Adam(self.parameters(), lr=self.hparams.max_lr)\n        scheduler = torch.optim.lr_scheduler.SequentialLR(\n            optimizer,\n            schedulers=[\n                torch.optim.lr_scheduler.LinearLR(\n                    optimizer, 1e-12, 1.0, self.hparams.num_warmup_step\n                ),\n                torch.optim.lr_scheduler.CosineAnnealingLR(\n                    optimizer, self.hparams.num_total_step, self.hparams.min_lr\n                ),\n            ],\n            milestones=[self.hparams.num_warmup_step],\n        )\n        return {\n            'optimizer': optimizer,\n            'lr_scheduler': {\n                'scheduler': scheduler,\n                'interval': 'step',\n            },\n        }\n\n\nclass Model2(pl.LightningModule):\n    def __init__(\n        self, max_lr=1e-3, min_lr=1e-5, num_warmup_step=1_000, num_total_step=20_000\n    ):\n        super().__init__()\n        self.save_hyperparameters()\n        self.conv0 = EdgeConv(MLP([34, 128, 256]), aggr='add')\n        self.conv1 = EdgeConv(MLP([512, 336, 256]), aggr='add')\n        self.conv2 = EdgeConv(MLP([512, 336, 256]), aggr='add')\n        self.conv3 = EdgeConv(MLP([512, 336, 256]), aggr='add')\n        self.post = MLP([1041, 336, 256])\n        self.readout = MLP([768, 128])\n        self.pred = nn.Linear(128, 3)\n\n    def forward(self, data: Batch):\n        vert_feat = data.x\n        batch = data.batch\n\n        vert_feat[:, 0] /= 500.0  # x\n        vert_feat[:, 1] /= 500.0  # y\n        vert_feat[:, 2] /= 500.0  # z\n        #vert_feat[:, 3] = (vert_feat[:, 3] - 1.0e04) / 3.0e4  # time\n        vert_feat[:, 3] = (vert_feat[:, 3] - 1.0e04) / 1.5e4  # time\n\n        vert_feat[:, 4] = torch.log10(vert_feat[:, 4]) / 3.0  # charge\n\n        edge_index = knn_graph(vert_feat[:, :3], 8, batch)\n\n        # Construct global features\n        hx = homophily(edge_index, vert_feat[:, 0], batch).reshape(-1, 1)\n        hy = homophily(edge_index, vert_feat[:, 1], batch).reshape(-1, 1)\n        hz = homophily(edge_index, vert_feat[:, 2], batch).reshape(-1, 1)\n        ht = homophily(edge_index, vert_feat[:, 3], batch).reshape(-1, 1)\n        means = scatter_mean(vert_feat, batch, dim=0)\n        n_p = torch.log10(data.n_pulses).reshape(-1, 1)\n        global_feats = torch.cat([means, hx, hy, hz, ht, n_p], dim=1)  # [B, 11]\n\n        # Distribute global_feats to each vertex\n        _, cnts = torch.unique_consecutive(batch, return_counts=True)\n        global_feats = torch.repeat_interleave(global_feats, cnts, dim=0)\n        vert_feat = torch.cat((vert_feat, global_feats), dim=1)\n\n        # Convolutions\n        feats = [vert_feat]\n        # Conv 0\n        vert_feat = self.conv0(vert_feat, edge_index)\n        feats.append(vert_feat)\n        # Conv 1\n        edge_index = knn_graph(vert_feat[:, :3], k=8, batch=batch)\n        vert_feat = self.conv1(vert_feat, edge_index)\n        feats.append(vert_feat)\n        # Conv 2\n        edge_index = knn_graph(vert_feat[:, :3], k=8, batch=batch)\n        vert_feat = self.conv2(vert_feat, edge_index)\n        feats.append(vert_feat)\n        # Conv 3\n        edge_index = knn_graph(vert_feat[:, :3], k=8, batch=batch)\n        vert_feat = self.conv3(vert_feat, edge_index)\n        feats.append(vert_feat)\n\n        # Postprocessing\n        post_inp = torch.cat(feats, dim=1)\n        post_out = self.post(post_inp)\n\n        # Readout\n        readout_inp = torch.cat(\n            [\n                scatter_min(post_out, batch, dim=0)[0],\n                scatter_max(post_out, batch, dim=0)[0],\n                scatter_mean(post_out, batch, dim=0),\n            ],\n            dim=1,\n        )\n        readout_out = self.readout(readout_inp)\n\n        # Predict\n        pred = self.pred(readout_out)\n        kappa = pred.norm(dim=1, p=2) + 1e-8\n        pred_x = pred[:, 0] / kappa\n        pred_y = pred[:, 1] / kappa\n        pred_z = pred[:, 2] / kappa\n        pred = torch.stack([pred_x, pred_y, pred_z, kappa], dim=1)\n\n        return pred\n\n    def train_or_valid_step(self, data, prefix):\n        pred_xyzk = self.forward(data)  # [B, 4]\n        true_xyz = data.gt.view(-1, 3)  # [B, 3]\n        loss = VonMisesFisher3DLoss()(pred_xyzk, true_xyz).mean()\n        error = angular_error(pred_xyzk[:, :3], true_xyz).mean()\n        self.log(f'loss/{prefix}', loss, batch_size=len(true_xyz))\n        self.log(f'error/{prefix}', error, batch_size=len(true_xyz))\n        return loss\n\n    def training_step(self, data, _):\n        return self.train_or_valid_step(data, 'train')\n\n    def validation_step(self, data, _):\n        self.train_or_valid_step(data, 'valid')\n\n    def configure_optimizers(self):\n        optimizer = torch.optim.Adam(self.parameters(), lr=self.hparams.max_lr)\n        scheduler = torch.optim.lr_scheduler.SequentialLR(\n            optimizer,\n            schedulers=[\n                torch.optim.lr_scheduler.LinearLR(\n                    optimizer, 1e-12, 1.0, self.hparams.num_warmup_step\n                ),\n                torch.optim.lr_scheduler.CosineAnnealingLR(\n                    optimizer, self.hparams.num_total_step, self.hparams.min_lr\n                ),\n            ],\n            milestones=[self.hparams.num_warmup_step],\n        )\n        return {\n            'optimizer': optimizer,\n            'lr_scheduler': {\n                'scheduler': scheduler,\n                'interval': 'step',\n            },\n        }    \n    \n    \n    \ndef collate_fn(x):\n    return x[0]\n\ndef prepare_sensors():\n    sensors = pd.read_csv(INPUT_PATH / \"sensor_geometry.csv\").astype(\n        {\n            \"sensor_id\": np.int16,\n            \"x\": np.float32,\n            \"y\": np.float32,\n            \"z\": np.float32,\n        }\n    )\n    sensors[\"string\"] = 0\n    sensors[\"qe\"] = 1\n\n    for i in range(len(sensors) // 60):\n        start, end = i * 60, (i * 60) + 60\n        sensors.loc[start:end, \"string\"] = i\n\n        # High Quantum Efficiency in the lower 50 DOMs - https://arxiv.org/pdf/2209.03042.pdf (Figure 1)\n        if i in range(78, 86):\n            start_veto, end_veto = i * 60, (i * 60) + 10\n            start_core, end_core = end_veto + 1, (i * 60) + 60\n            sensors.loc[start_core:end_core, \"qe\"] = 1.35\n\n    # https://github.com/graphnet-team/graphnet/blob/b2bad25528652587ab0cdb7cf2335ee254cfa2db/src/graphnet/models/detector/icecube.py#L33-L41\n    # Assume that \"rde\" (relative dom efficiency) is equivalent to QE\n    sensors[\"x\"] /= 500\n    sensors[\"y\"] /= 500\n    sensors[\"z\"] /= 500\n    sensors[\"qe\"] -= 1.25\n    sensors[\"qe\"] /= 0.25\n\n    return sensors\n\n\ndef infer(model, loader, batch_size=32, device=\"cuda\"):\n    model.to(device)\n    model.eval()\n    #model = TTAWrapper(model, device)\n    #loader = DataLoader(dataset, batch_size=batch_size, num_workers=2)\n\n    predictions = []\n    with torch.no_grad():\n        for batch in loader:\n            batch = batch.to(device)\n            pred_azi, pred_zen = model(batch)\n            pred_angles = torch.stack([pred_azi[:, 0], pred_zen[:, 0]], dim=1)\n            predictions.append(pred_angles.cpu())\n\n    return torch.cat(predictions, 0)\n\ndef infer_xyz(model, loader, batch_size=32, device=\"cuda\"):\n    model.to(device)\n    model.eval()\n    print(\"start prediction\")\n    predictions = []\n    with torch.no_grad():\n        for batch in loader:\n            #print(batch.shape)\n            batch = batch.to(device)\n            #pred_angles = model(batch)\n            pred_xyzk = model (batch)\n            #print(pred_xyzk.shape) #(e,4) -> (3event,4)\n            #print(pred_xyzk) \n            #raise\n            #predictions.append(pred_angles.cpu())\n            predictions.append(pred_xyzk.cpu())            \n                        \n    return torch.cat(predictions, 0)\n\n\n\n\ndef xyz_to_angle(xyz_b):\n    x, y, z = xyz_b.t()\n    az = torch.arccos(x / torch.sqrt(x**2 + y**2)) * torch.sign(y)\n    zen = torch.arccos(z / torch.sqrt(x**2 + y**2 + z**2))\n    return torch.stack([az, zen], dim=1)\n\n\n#def make_predictions(model, device=\"cpu\", mode=\"test\", batch_size=32):\n#def make_predictions(model, device=\"cpu\", mode=\"test\"):\n#def make_predictions(model, device=\"cpu\", mode=\"test\",weight_path = \"\"):   \n#def make_predictions(model, device=\"cpu\", mode=\"test\"): \ndef make_predictions(model, device=\"cpu\", mode=\"test\" , dataiter = \"V2\"):    \n    import pyarrow\n    import pyarrow.parquet\n    import pandas as pd\n    from tqdm import tqdm\n    from pathlib import Path\n\n    #model = Model(num_total_step=num_total_step, num_warmup_step=num_warmup_step)\n    #model = model.to(\"cuda\")\n    #weights =  torch.load(weight_path)[\"state_dict\"]\n    \n    \n#     new_weights = dict()\n#     for k, v in weights.items():\n#         k = k.replace('_gnn._conv_layers.0', 'conv0')\n#         k = k.replace('_gnn._conv_layers.1', 'conv1')\n#         k = k.replace('_gnn._conv_layers.2', 'conv2')\n#         k = k.replace('_gnn._conv_layers.3', 'conv3')\n#         k = k.replace('_gnn._post_processing', 'post')\n#         k = k.replace('_gnn._readout', 'readout')\n#         k = k.replace('_tasks.0._affine', 'pred')\n#         new_weights[k] = v\n#     print(model.load_state_dict(new_weights))\n    #print(trainer.validate(model, valid_loader))    \n    \n    CheckBatchID =659 #51\n    \n    meta = pd.read_parquet(\n        INPUT_PATH / f'{mode}_meta.parquet', columns=[\"batch_id\", \"event_id\"]\n    ).astype(_dtype)\n    batch_ids = meta[\"batch_id\"].unique()\n\n    if mode == \"train\":\n        #batch_ids = batch_ids[:6]\n        batch_ids = [CheckBatchID]\n    #print(batch_ids)\n    batch_preds = []\n    for b in batch_ids:\n        event_ids = meta[meta[\"batch_id\"] == b][\"event_id\"].tolist()\n        \n        #dataset = IceCubeDataset(\n        #    b, event_ids, sensors, mode=mode,\n        #)\n        #loader = DataLoader(dataset, batch_size=batch_size, num_workers=1)\n        #valid_set = IceCube(parquet_dir, meta_dir, [51], batch_size=100)\n        #valid_set = IceCube(parquet_dir, meta_dir, [CheckBatchID], batch_size=100)\n        #valid_set = IceCubeV2(b, event_ids,mode, batch_size=100) \n        if(dataiter==\"V2\"):\n            valid_set = IceCubeV2(b, event_ids,mode, batch_size=200) \n        elif(dataiter==\"V3\"):\n            valid_set = IceCubeV3(b, event_ids,mode, batch_size=200) \n        else:\n            valid_set = IceCubeV2(b, event_ids,mode, batch_size=200)        \n        \n        \n        print(\"start loader\")\n        valid_loader = DataLoader(\n            valid_set,\n            batch_size=1,\n            num_workers=1,\n            collate_fn=collate_fn,\n        ) \n        #ytt modify her\n        #pred_xyzk = model (data)  # [B, 4]\n        model.to(device)\n        model.eval()\n        print(\"start prediction\")\n        predictions = []\n        with torch.no_grad():\n            for batch in valid_loader:\n                #print(batch.shape)\n                batch = batch.to(device)\n                #pred_angles = model(batch)\n                pred_xyzk = model (batch)\n                #print(pred_xyzk.shape) #(e,4) -> (3event,4)\n                #print(pred_xyzk) \n                #raise\n                #predictions.append(pred_angles.cpu())\n                predictions.append(pred_xyzk.cpu())\n        \n        batch_preds.append(torch.cat(predictions, 0))\n        #batch_preds.append(infer_xyz(model, valid_loader, device=device))\n        #print(\"Finished batch\", b)\n\n        #if mode == \"train\" and b == 6:\n        #    break\n\n    output = torch.cat(batch_preds, 0)\n    print(\"output~~~\")\n    print(output.shape)\n    #print(output)\n    \n    #xyz to angel here\n    #output2 = xyz_to_angle(output[:,:3])\n    #print(output2.shape)\n    #print(output2)\n    return output\n    #raise\n#     event_id_labels = []\n#     for b in batch_ids:\n#         event_id_labels.extend(meta[meta[\"batch_id\"] == b][\"event_id\"].tolist())\n\n#     sub = {\n#         \"event_id\": event_id_labels,\n#         \"azimuth\": output2[:, 0],\n#         \"zenith\": output2[:, 1],\n#     }\n\n#     sub = pd.DataFrame(sub)\n    \n#     #all positive\n#     sub[\"azimuth\"] = sub[\"azimuth\"].apply(lambda x: x if x>=0 else x+2*np.pi)\n    \n#     sub.to_csv(\"submission.csv\", index=False)\n\n\n\nif __name__ == '__main__':\n    pl.seed_everything(123)\n    #torch.set_float32_matmul_precision('medium')\n\n    # Memory seems to be leaking if not set to 'spawn'\n    # Maybe due to https://github.com/pytorch/pytorch/issues/13246#issuecomment-905703662\n    # It also decrease the memory consumption\n    torch.multiprocessing.set_start_method('spawn', force=True)\n\n    # Config\n    num_total_step = 100_000\n    num_warmup_step = 1_000\n    parquet_dir = Path('data/train')\n    meta_dir = Path('meta')\n    \n    #assert parquet_dir.exists() and meta_dir.exists()\n    log_dir = Path('log') / Path(__file__).stem\n    log_dir.mkdir(parents=True, exist_ok=True)\n\n    #model = Model(num_total_step=num_total_step, num_warmup_step=num_warmup_step)\n    #model = model.to(\"cuda\")\n\n\n    # Verify with offical pretrained weight\n    #weights = torch.load('state_dict.pth')\n\n    mode=\"test\"\n    #-----model1------- \n    model = Model(num_total_step=num_total_step, num_warmup_step=num_warmup_step)\n    model = model.to(\"cuda\")\n\n    #trainall tune5 r1 lb1001\n    #model.load_state_dict(torch.load('/kaggle/input/icecube/trainall_tune5_epoch02_loss_valid1.546_error_valid1.005_r1.ckpt')[\"state_dict\"])\n    #model.load_state_dict(torch.load('/kaggle/input/icecube/trainall_tune5_epoch04_loss_valid1.556_error_valid1.003_r8.ckpt')[\"state_dict\"])\n    \n    #model.load_state_dict(torch.load('/kaggle/input/icecube/trainall_tune5_epoch05_loss_valid1.537_error_valid1.002_r9.ckpt')[\"state_dict\"])\n    #train all tune7  \n    #model.load_state_dict(torch.load('/kaggle/input/icecube/trainall_tune7_epoch00_loss_valid1.521_error_valid1.001_r1.ckpt')[\"state_dict\"])\n    #train all tune8\n    model.load_state_dict(torch.load('/kaggle/input/icecube/trainall_tune8_gac20_epoch01_loss_valid1.475_error_valid0.998_r1.ckpt')[\"state_dict\"])\n    \n    result = make_predictions(model, device=\"cuda\", mode=mode)\n    del model\n    #-----model2-------\n    #model2 = Model(num_total_step=num_total_step, num_warmup_step=num_warmup_step)\n    #t1.5\n    model2 = Model2(num_total_step=num_total_step, num_warmup_step=num_warmup_step)\n     \n    model2 = model2.to(\"cuda\")\n\n    #trainall tune5 r2 \n    #model2.load_state_dict(torch.load('/kaggle/input/icecube/trainall_tune5_epoch02_loss_valid1.636_error_valid1.005_r2.ckpt')[\"state_dict\"])\n    #model2.load_state_dict(torch.load('/kaggle/input/icecube/trainall_tune5_epoch05_loss_valid1.537_error_valid1.002_r9.ckpt')[\"state_dict\"])\n\n    #model2.load_state_dict(torch.load('/kaggle/input/icecube/trainall_tune5_epoch08_loss_valid1.499_error_valid1.002_r18.ckpt')[\"state_dict\"])\n    #train all tune7\n    #model2.load_state_dict(torch.load('/kaggle/input/icecube/trainall_tune7_epoch00_loss_valid1.502_error_valid1.001_r2.ckpt')[\"state_dict\"])\n\n    #result2 = make_predictions(model2, device=\"cuda\", mode=mode)\n    \n    #t1.5 en 99007\n    model2.load_state_dict(torch.load('/kaggle/input/icecube/gac20_pulselen_von_mse3d_t1p5_epoch00_loss_valid1.822_error_valid0.9942.ckpt')[\"state_dict\"])\n    result2 = make_predictions(model2, device=\"cuda\", mode=mode,dataiter=\"V3\")    \n    \n    del model2\n    \n    #-----model3-------\n    model3 = Model(num_total_step=num_total_step, num_warmup_step=num_warmup_step)\n    model3 = model3.to(\"cuda\")\n\n    #trainall tune5 r4 lb1001\n    #model3.load_state_dict(torch.load('/kaggle/input/icecube/trainall_tune5_epoch02_loss_valid1.555_error_valid1.005_r3.ckpt')[\"state_dict\"])\n    #model3.load_state_dict(torch.load('/kaggle/input/icecube/trainall_tune5_epoch05_loss_valid1.525_error_valid1.003_r10.ckpt')[\"state_dict\"])\n    #model3.load_state_dict(torch.load('/kaggle/input/icecube/trainall_tune6_epoch00_loss_valid1.479_error_valid0.999_r2.ckpt')[\"state_dict\"])\n    #new v5 p300\n    model3.load_state_dict(torch.load('/kaggle/input/icecube/gac20_pulselen_von_mse3d_maxpulse300_epoch00_loss_valid1.813_error_valid0.9934_r1.ckpt')[\"state_dict\"])\n\n    #train all tune7\n    #model3.load_state_dict(torch.load('/kaggle/input/icecube/trainall_tune7_epoch00_loss_valid1.501_error_valid1.001_r3.ckpt')[\"state_dict\"])\n\n    #result3 = make_predictions(model3, device=\"cuda\", mode=mode)\n    result3 = make_predictions(model3, device=\"cuda\", mode=mode,dataiter=\"V3\")\n    del model3\n    \n    #-----model4-------\n    #model4 = Model(num_total_step=num_total_step, num_warmup_step=num_warmup_step)\n    #t1.5\n    model4 = Model2(num_total_step=num_total_step, num_warmup_step=num_warmup_step)\n    model4 = model4.to(\"cuda\")\n\n    #trainall tune5 r4 lb1001\n    #model4.load_state_dict(torch.load('/kaggle/input/icecube/trainall_tune5_epoch02_loss_valid1.517_error_valid1.004_r4.ckpt')[\"state_dict\"])\n    #model4.load_state_dict(torch.load('/kaggle/input/icecube/trainall_tune5_epoch05_loss_valid1.560_error_valid1.003_r11.ckpt')[\"state_dict\"])\n    #train all tune7\n    #model4.load_state_dict(torch.load('/kaggle/input/icecube/trainall_tune7_epoch01_loss_valid1.496_error_valid1.001_r4.ckpt')[\"state_dict\"])\n    #train all tune8\n    #model4.load_state_dict(torch.load('/kaggle/input/icecube/trainall_tune8_gac20_epoch01_loss_valid1.476_error_valid0.998_r2.ckpt')[\"state_dict\"])\n    #train all tune10\n    #model4.load_state_dict(torch.load('/kaggle/input/icecube/trainall_tune10_gac20_new_pluselen_epoch00_loss_valid1.435_error_valid0.995_r4.ckpt')[\"state_dict\"])\n    #t1.5 v2\n    model4.load_state_dict(torch.load('/kaggle/input/icecube/gac20_pulselen_von_mse3d_t1p5_epoch01_loss_valid1.822_error_valid0.9943.ckpt')[\"state_dict\"])\n    \n    result4 = make_predictions(model4, device=\"cuda\", mode=mode,dataiter=\"V3\")\n    del model4\n    \n    #-----model5-------\n    model5 = Model(num_total_step=num_total_step, num_warmup_step=num_warmup_step)\n    model5 = model5.to(\"cuda\")\n\n    #trainall tune5 r5\n    #model5.load_state_dict(torch.load('/kaggle/input/icecube/trainall_tune5_epoch03_loss_valid1.562_error_valid1.004_r5.ckpt')[\"state_dict\"])\n    #model5.load_state_dict(torch.load('/kaggle/input/icecube/trainall_tune5_epoch05_loss_valid1.543_error_valid1.003_r12.ckpt')[\"state_dict\"])\n    #train all tune7\n    #model5.load_state_dict(torch.load('/kaggle/input/icecube/trainall_tune7_epoch01_loss_valid1.522_error_valid1.001_r5.ckpt')[\"state_dict\"])\n    #train all tune8\n    #model5.load_state_dict(torch.load('/kaggle/input/icecube/trainall_tune8_gac20_epoch01_loss_valid1.474_error_valid0.998_r3.ckpt')[\"state_dict\"])\n    #tune11 von mse\n    #model5.load_state_dict(torch.load('/kaggle/input/icecube/gac20_pulselen_von_mse3d_epoch00_loss_valid1.821_error_valid0.9944_r1.ckpt')[\"state_dict\"])\n    model5.load_state_dict(torch.load('/kaggle/input/icecube/gac20_pulselen_von_mse3d2_epoch00_loss_valid1.591_error_valid0.9936_r1.ckpt')[\"state_dict\"])\n    \n    result5 = make_predictions(model5, device=\"cuda\", mode=mode,dataiter=\"V3\")\n    del model5\n    #-----model6-------\n    model6 = Model(num_total_step=num_total_step, num_warmup_step=num_warmup_step)\n    model6 = model6.to(\"cuda\")\n\n    #trainall tune5 r\n    #model6.load_state_dict(torch.load('/kaggle/input/icecube/trainall_tune5_epoch03_loss_valid1.531_error_valid1.004_r6.ckpt')[\"state_dict\"])\n    #model6.load_state_dict(torch.load('/kaggle/input/icecube/trainall_tune5_epoch05_loss_valid1.558_error_valid1.003_r13.ckpt')[\"state_dict\"])\n    #train all tune7\n    #model6.load_state_dict(torch.load('/kaggle/input/icecube/trainall_tune7_epoch01_loss_valid1.506_error_valid1.001_r6.ckpt')[\"state_dict\"])\n    #train all tune8\n    #model6.load_state_dict(torch.load('/kaggle/input/icecube/trainall_tune8_gac20_epoch02_loss_valid1.474_error_valid0.998_r4.ckpt')[\"state_dict\"])\n    #tune11 vonmse\n    #model6.load_state_dict(torch.load('/kaggle/input/icecube/gac20_pulselen_von_mse3d_epoch00_loss_valid1.821_error_valid0.9941_r3.ckpt')[\"state_dict\"])\n    model6.load_state_dict(torch.load('/kaggle/input/icecube/gac20_pulselen_von_mse3d2_epoch00_loss_valid1.594_error_valid0.9937_r2.ckpt')[\"state_dict\"])  \n    \n    \n    result6 = make_predictions(model6, device=\"cuda\", mode=mode,dataiter=\"V3\")\n    del model6\n    #-----model7-------\n    model7 = Model(num_total_step=num_total_step, num_warmup_step=num_warmup_step)\n    model7 = model7.to(\"cuda\")\n\n    #trainall tune5 r\n    #model7.load_state_dict(torch.load('/kaggle/input/icecube/trainall_tune5_epoch03_loss_valid1.539_error_valid1.005_r7.ckpt')[\"state_dict\"])\n    #model7.load_state_dict(torch.load('/kaggle/input/icecube/trainall_tune5_epoch06_loss_valid1.588_error_valid1.003_r14.ckpt')[\"state_dict\"])\n    #train all tune7\n    #model7.load_state_dict(torch.load('/kaggle/input/icecube/trainall_tune7_epoch01_loss_valid1.518_error_valid1.000_r7.ckpt')[\"state_dict\"])\n    #result7 = make_predictions(model7, device=\"cuda\", mode=mode)\n    model7.load_state_dict(torch.load('/kaggle/input/icecube/gac20_pulselen_von_mse3d_maxpulse300_epoch01_loss_valid1.815_error_valid0.9932_r2.ckpt')[\"state_dict\"])\n    result7 = make_predictions(model7, device=\"cuda\", mode=mode,dataiter=\"V3\")\n    del model7\n    \n     #-----result ensembling-------\n     #en2    \n#     result[:,0]=  (result[:,0]+result2[:,0])/2 \n#     result[:,1]=  (result[:,1]+result2[:,1])/2\n#     result[:,2]=  (result[:,2]+result2[:,2])/2\n  \n    #en7\n    #result[:,0]=  (result[:,0]+result2[:,0]+ result3[:,0]+result4[:,0] +result5[:,0]+result6[:,0]+result7[:,0])/7 \n    #result[:,1]=  (result[:,1]+result2[:,1]+ result3[:,1]+result4[:,1] +result5[:,1]+result6[:,1]+result7[:,1])/7 \n    #result[:,2]=  (result[:,2]+result2[:,2]+ result3[:,2]+result4[:,2] +result5[:,2]+result6[:,2]+result7[:,2])/7 \n\n    #en7 weighted\n    result[:,0]=  (result[:,0]+result2[:,0]*1.5+ result3[:,0]+result4[:,0] +result5[:,0]+result6[:,0]+result7[:,0])/7.5 \n    result[:,1]=  (result[:,1]+result2[:,1]*1.5+ result3[:,1]+result4[:,1] +result5[:,1]+result6[:,1]+result7[:,1])/7.5 \n    result[:,2]=  (result[:,2]+result2[:,2]*1.5+ result3[:,2]+result4[:,2] +result5[:,2]+result6[:,2]+result7[:,2])/7.5     \n    \n    result = xyz_to_angle(result[:,:3])\n    #print(result.shape)\n    #xyz to angel here\n    #result = xyz_to_angle(result[:,:3])\n    #print(result.shape)\n    #print(result)    \n    #raise\n    \n    meta = pd.read_parquet(\n        INPUT_PATH / f'{mode}_meta.parquet', columns=[\"batch_id\", \"event_id\"]\n    ).astype(_dtype)\n    \n\n    batch_ids = meta[\"batch_id\"].unique()\n       \n    event_id_labels = []\n    for b in batch_ids:\n        event_id_labels.extend(meta[meta[\"batch_id\"] == b][\"event_id\"].tolist())\n\n    sub = {\n        \"event_id\": event_id_labels,\n        \"azimuth\": result[:, 0],\n        \"zenith\": result[:, 1],\n    }\n\n    sub = pd.DataFrame(sub)\n    \n    #all positive\n    sub[\"azimuth\"] = sub[\"azimuth\"].apply(lambda x: x if x>=0 else x+2*np.pi)\n    \n    sub.to_csv(\"graph_net.csv\", index=False)    \n    sub.to_csv(\"submission.csv\", index=False) \n    \n    \n    print(pd.read_csv(\"graph_net.csv\"))","metadata":{"execution":{"iopub.status.busy":"2023-04-18T13:26:15.541959Z","iopub.execute_input":"2023-04-18T13:26:15.542654Z","iopub.status.idle":"2023-04-18T13:26:15.567004Z","shell.execute_reply.started":"2023-04-18T13:26:15.54262Z","shell.execute_reply":"2023-04-18T13:26:15.565882Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"_uuid":"5aec2103-12c1-4836-9e48-61dc5db48c1f","_cell_guid":"beab4c67-d8c8-443d-9558-7876db16a5d1","collapsed":false,"jupyter":{"outputs_hidden":false},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!python exp-gnn.py","metadata":{"_uuid":"3b77e91b-1d20-498d-9310-a803d2c7e881","_cell_guid":"dbace73d-eb59-4b17-886c-ecf769c3b2d5","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-04-18T13:26:15.568821Z","iopub.execute_input":"2023-04-18T13:26:15.569357Z","iopub.status.idle":"2023-04-18T13:27:04.111929Z","shell.execute_reply.started":"2023-04-18T13:26:15.569243Z","shell.execute_reply":"2023-04-18T13:27:04.110727Z"},"trusted":true},"execution_count":null,"outputs":[]}]}