{"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":"## Outline\n### Browse ICNI datasets\n>read head of dataset /show examples\n  \n### CFG_Configuration parameters\n> \n  \n### load ICNI datasets\n>use 2,000 events, which is 1% of 1 batch 200,000 events\n  \n### batch data preprocess\n>add elapsed time, event aggregate amount, etc.\n  \n### construct (Graph_)label(targets i.e. azimuth & zenith) from meta data\n> \n  \n### construct Graph(node/edge/label) from batch dataframe\n>edge(link) formation by k-nearest neighbors to pulse observation time series  \n>set limit to the first 220 pulses (covering 95% of events) in large events with many pulses  \n>the number of (sensor) neighbors used should be optimized  \n  \n### G(C)NN:Graph(Convolutional)NeuralNetwork model design\n>G(C)NN with 3GCN layers + 3MLP layers (network structure needs optimization)  \n>200 epochs with cpu for trial learning  \n  \n### neutrino_direction prediction results\n>summarize and visualize train/valid observed predictions after the set epoch (not earlystopping)  \n>summarize and extract events hard/easy to predict  \n  \n### neutrino_direction sensor/event visualization\n>3D plot output of hard/easy-to-predict events  \n  \n### test data prediction/inference\n>not included yet (for future work)\n  ","metadata":{}},{"cell_type":"markdown","source":"## References and Acknowledgments  \n### Thanks for these excellent works and notebooks!  \n  \n❄️IceCube Neutrinos - Domain & EDA for DS folks  \nhttps://www.kaggle.com/code/mvvppp/icecube-neutrinos-domain-eda-for-ds-folks  \nIcecube - Neutrino trajectory 3D projection  \nhttps://www.kaggle.com/code/diegoasuarezg/icecube-neutrino-trajectory-3d-projection  \nIceCube EDA and Domain of Nuetrinos  \nhttps://www.kaggle.com/code/yoshikuwano/icecube-eda-and-domain-of-nuetrinos  \nIceCube🧊: Neutrino🎆 fitting 3D🌪️points cloud  \nhttps://www.kaggle.com/code/jirkaborovec/icecube-neutrino-fitting-3d-points-cloud  \n  \n⚡🧊⚡LB 1.183 Lightning Fast Baseline with Polars  \nhttps://www.kaggle.com/code/roberthatch/lb-1-183-lightning-fast-baseline-with-polars  \nMake your own GNN  \nhttps://www.kaggle.com/code/kilogrand/make-your-own-gnn","metadata":{}},{"cell_type":"code","source":"%%time\nimport os\nif not os.path.exists('/opt/conda/lib/python3.7/site-packages/polars'):\n    ! pip install --quiet polars","metadata":{"execution":{"iopub.status.busy":"2023-03-27T04:08:56.478662Z","iopub.execute_input":"2023-03-27T04:08:56.479164Z","iopub.status.idle":"2023-03-27T04:09:11.010478Z","shell.execute_reply.started":"2023-03-27T04:08:56.479022Z","shell.execute_reply":"2023-03-27T04:09:11.008953Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport glob\nimport random\nimport math\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\nimport polars as pl\nimport math\nimport time\nimport copy\nimport gc\nfrom tqdm import tqdm\nfrom tqdm.notebook import tqdm\n\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\npl.Config.set_tbl_rows(8)\npl.Config.set_tbl_cols(-1)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-03-27T04:09:11.013744Z","iopub.execute_input":"2023-03-27T04:09:11.014212Z","iopub.status.idle":"2023-03-27T04:09:11.7342Z","shell.execute_reply.started":"2023-03-27T04:09:11.014174Z","shell.execute_reply":"2023-03-27T04:09:11.732222Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## browse ICNI datasets","metadata":{}},{"cell_type":"code","source":"PATH_DATASET= '/kaggle/input/icecube-neutrinos-in-deep-ice'","metadata":{"execution":{"iopub.status.busy":"2023-03-27T04:09:11.735325Z","iopub.execute_input":"2023-03-27T04:09:11.735674Z","iopub.status.idle":"2023-03-27T04:09:11.740712Z","shell.execute_reply.started":"2023-03-27T04:09:11.735644Z","shell.execute_reply":"2023-03-27T04:09:11.739435Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## Browse meta data [train/test]_meta.parquet\nn_rows = 10\n# n_rows = None\nmeta_train = pl.read_parquet(os.path.join(PATH_DATASET, \"train_meta.parquet\"), n_rows=n_rows)\ndisplay(meta_train.head(10))\nmeta_test = pl.read_parquet(os.path.join(PATH_DATASET, \"test_meta.parquet\"), n_rows=n_rows)\ndisplay(meta_test.head())","metadata":{"execution":{"iopub.status.busy":"2023-03-27T04:09:11.743229Z","iopub.execute_input":"2023-03-27T04:09:11.743534Z","iopub.status.idle":"2023-03-27T04:09:11.902535Z","shell.execute_reply.started":"2023-03-27T04:09:11.743507Z","shell.execute_reply":"2023-03-27T04:09:11.901533Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## Browse sensor geometry\nsensor = pl.read_csv(os.path.join(PATH_DATASET, \"sensor_geometry.csv\"))\ndisplay(sensor.head())","metadata":{"execution":{"iopub.status.busy":"2023-03-27T04:09:11.903797Z","iopub.execute_input":"2023-03-27T04:09:11.904127Z","iopub.status.idle":"2023-03-27T04:09:11.919233Z","shell.execute_reply.started":"2023-03-27T04:09:11.904098Z","shell.execute_reply":"2023-03-27T04:09:11.918165Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## Browse training/test data\nn_rows = 10\n# n_rows = None\ntrain = pl.read_parquet(os.path.join(PATH_DATASET, \"train/batch_1.parquet\"), \n                        columns=['event_id','sensor_id','time','charge','auxiliary'],\n                        n_rows=n_rows)\nprint(f\"counts of events in loaded batch: {len(train['event_id'].unique())}\")\ntrain","metadata":{"execution":{"iopub.status.busy":"2023-03-27T04:09:11.920574Z","iopub.execute_input":"2023-03-27T04:09:11.920906Z","iopub.status.idle":"2023-03-27T04:09:11.963583Z","shell.execute_reply.started":"2023-03-27T04:09:11.920879Z","shell.execute_reply":"2023-03-27T04:09:11.962742Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## CFG_Configuration parameters","metadata":{}},{"cell_type":"code","source":"## Configuration parameters\nMODE = 'train'\n# MODE = 'test'\n\nTRAIN_MAX_EVENTS = 2000\n#TRAIN_MAX_EVENTS = None\nTRAIN_BATCH_START = 1\nTRAIN_N_BATCHES = 1\n\n## Constants\nPATH_DATASET= '/kaggle/input/icecube-neutrinos-in-deep-ice'","metadata":{"execution":{"iopub.status.busy":"2023-03-27T04:09:11.964797Z","iopub.execute_input":"2023-03-27T04:09:11.965414Z","iopub.status.idle":"2023-03-27T04:09:11.970032Z","shell.execute_reply.started":"2023-03-27T04:09:11.96538Z","shell.execute_reply":"2023-03-27T04:09:11.969231Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Load ICNI datasets","metadata":{}},{"cell_type":"code","source":"print(f'Mode Setting :{MODE}')\nprint('load meta data...')\nmeta = pl.scan_parquet(f'{PATH_DATASET}/{MODE}_meta.parquet')\n\nprint('load sensor data...')\nsensor = (pl.scan_csv(f'{PATH_DATASET}/sensor_geometry.csv')\n            .with_columns([\n                pl.col('sensor_id').cast(pl.Int16)              \n            ])\n         )\nprint(sensor)","metadata":{"execution":{"iopub.status.busy":"2023-03-27T04:09:11.971201Z","iopub.execute_input":"2023-03-27T04:09:11.972301Z","iopub.status.idle":"2023-03-27T04:09:11.989062Z","shell.execute_reply.started":"2023-03-27T04:09:11.972269Z","shell.execute_reply":"2023-03-27T04:09:11.987271Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n%%time\nif MODE == 'train':\n    batch_id_start = TRAIN_BATCH_START\n    batch_id_end = batch_id_start + TRAIN_N_BATCHES\n    print(f'load data from batch {batch_id_start} to {batch_id_end-1}')\n\n    batch = pl.DataFrame()\n    for batch_id in range(batch_id_start, batch_id_end):\n        print(f'loading data of batch {batch_id}')\n        max_events = TRAIN_MAX_EVENTS\n        batch_i = pl.scan_parquet(f'{PATH_DATASET}/{MODE}/batch_{batch_id}.parquet')\n\n        ## set/load max_events\n        if max_events is not None:\n            batch_i = batch_i.collect()\n            last_event_id = batch_i.select(pl.col('event_id')).unique()[TRAIN_MAX_EVENTS-1, 0]\n            batch_i = batch_i.filter(pl.col('event_id') <= last_event_id)\n        else:\n            max_events = \"all\"\n            batch_i = batch_i.collect()\n        ## stack loaded batch data\n        batch = pl.concat([batch, batch_i], how=\"vertical\")\n\n    ## Merge in sensor x,y,z data\n    batch = batch.lazy().join(sensor, on='sensor_id', how='left').collect()\n    del batch_i\n    print(f'data of batch {batch_id_start} to {batch_id_end-1} every {max_events} events loaded')\n    print(f\"counts of events in loaded batch: {len(batch['event_id'].unique())}\")\n\nif MODE == 'test':\n    batch_id = 661\n    print(f'load data from test batch_661')\n    batch = pl.scan_parquet(f'{PATH_DATASET}/{MODE}/batch_{batch_id}.parquet')\n    ## Merge in sensor x,y,z data\n    batch = batch.join(sensor, on='sensor_id', how='left').collect()\n    print(f'data of batch_661 (test) loaded')\n    print(f\"counts of events in loaded batch: {len(batch['event_id'].unique())}\")","metadata":{"execution":{"iopub.status.busy":"2023-03-27T04:09:11.991577Z","iopub.execute_input":"2023-03-27T04:09:11.992052Z","iopub.status.idle":"2023-03-27T04:09:14.145076Z","shell.execute_reply.started":"2023-03-27T04:09:11.992007Z","shell.execute_reply":"2023-03-27T04:09:14.143858Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## original/loaded batch data\nprint(f'data size :{batch.estimated_size(unit=\"mb\"):.2f} MB')\nbatch","metadata":{"execution":{"iopub.status.busy":"2023-03-27T04:09:14.149781Z","iopub.execute_input":"2023-03-27T04:09:14.15013Z","iopub.status.idle":"2023-03-27T04:09:14.160335Z","shell.execute_reply.started":"2023-03-27T04:09:14.150082Z","shell.execute_reply":"2023-03-27T04:09:14.159364Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## batch data preprocess","metadata":{}},{"cell_type":"code","source":"%%time\n## process batch data\nbatch = (batch.lazy()\n         # elapsed time in event\n         .with_columns([\n             (pl.col(\"time\") - pl.min(\"time\").over(\"event_id\")).alias(\"elapsed_time\"),\n         ])\n         # label/aggrigate event/sensor data\n         .with_columns([\n             pl.col(\"sensor_id\").n_unique().over(\"event_id\").alias(\"n_uniq_sens_per_evt\"),\n             pl.col(\"sensor_id\").count().over(\"event_id\").alias(\"n_pls_per_evt\"),\n         ])\n         # sort columns of event/sensor/geom data\n         .select(['event_id','n_pls_per_evt','n_uniq_sens_per_evt',\n                  'sensor_id','time','elapsed_time','charge','auxiliary','x','y','z'])\n        ).collect()","metadata":{"execution":{"iopub.status.busy":"2023-03-27T04:09:14.161802Z","iopub.execute_input":"2023-03-27T04:09:14.162499Z","iopub.status.idle":"2023-03-27T04:09:14.191665Z","shell.execute_reply.started":"2023-03-27T04:09:14.162416Z","shell.execute_reply":"2023-03-27T04:09:14.190399Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## preprocessed batch data\nprint(f'data size :{batch.estimated_size(unit=\"mb\"):.2f} MB')\nbatch","metadata":{"execution":{"iopub.status.busy":"2023-03-27T04:09:14.193408Z","iopub.execute_input":"2023-03-27T04:09:14.193853Z","iopub.status.idle":"2023-03-27T04:09:14.205372Z","shell.execute_reply.started":"2023-03-27T04:09:14.193808Z","shell.execute_reply":"2023-03-27T04:09:14.204074Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## construct (Graph_)label(targets i.e. azimuth & zenith) from meta data","metadata":{}},{"cell_type":"code","source":"\n%%time\ndf_label = (meta\n            .select([\"event_id\",\"azimuth\",\"zenith\"])\n            .filter(pl.col(\"event_id\").is_in(batch[\"event_id\"].unique()))\n            .with_columns([\n                pl.col([\"event_id\"]).cast(pl.Int32),\n                pl.col([\"azimuth\", \"zenith\"]).cast(pl.Float32),\n            ])\n           ).collect()\n\n## The generated \"df_label\" is the targets(azimuth & zenith) information of the graph for each event in loaded batch\nprint(f'df_label data size :{df_label.estimated_size(unit=\"mb\"):.2f} MB')\ndisplay(df_label)","metadata":{"execution":{"iopub.status.busy":"2023-03-27T04:09:14.207688Z","iopub.execute_input":"2023-03-27T04:09:14.208312Z","iopub.status.idle":"2023-03-27T04:09:31.235205Z","shell.execute_reply.started":"2023-03-27T04:09:14.208199Z","shell.execute_reply":"2023-03-27T04:09:31.233685Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## construct Graph(node/edge/label) from batch dataframe","metadata":{}},{"cell_type":"markdown","source":"### install whl from kaggle/input  \n### CPU version\n!pip install torch-scatter torch-sparse torch-cluster torch-spline-conv torch-geometric --no-index --find-links=file:///kaggle/input/pytorch-geometric/PyTorch-Geometric  \n  \n### GPU version  \n!pip install torch-scatter torch-sparse torch-cluster torch-spline-conv torch-geometric --no-index --find-links=file:///kaggle/input/torch-geometric  ","metadata":{}},{"cell_type":"code","source":"%%time\n!pip install torch-scatter torch-sparse torch-cluster torch-spline-conv torch-geometric --no-index --find-links=file:///kaggle/input/pytorch-geometric/PyTorch-Geometric","metadata":{"execution":{"iopub.status.busy":"2023-03-27T04:09:31.236686Z","iopub.execute_input":"2023-03-27T04:09:31.237031Z","iopub.status.idle":"2023-03-27T04:09:43.596096Z","shell.execute_reply.started":"2023-03-27T04:09:31.236998Z","shell.execute_reply":"2023-03-27T04:09:43.594322Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn.functional as F\nfrom torch.nn import ModuleList, Linear, BatchNorm1d\nfrom torch_geometric.nn import knn_graph\nfrom torch_geometric.nn import GCNConv, NNConv, SAGEConv\nfrom torch_geometric.nn import global_add_pool, global_mean_pool, global_max_pool\nfrom torch_geometric.data import Data\nfrom torch_geometric.loader import DataLoader\nfrom torch_scatter import scatter_max\nimport networkx as nx\nimport sklearn\nfrom sklearn.metrics import mean_squared_error, r2_score\nfrom sklearn.model_selection import KFold, train_test_split","metadata":{"execution":{"iopub.status.busy":"2023-03-27T04:09:43.598748Z","iopub.execute_input":"2023-03-27T04:09:43.599241Z","iopub.status.idle":"2023-03-27T04:09:49.521088Z","shell.execute_reply.started":"2023-03-27T04:09:43.599186Z","shell.execute_reply":"2023-03-27T04:09:49.519802Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## library version\nfor use_library in torch, pd, np, pl, nx, sklearn:\n    print(f\"{use_library} ver : {use_library.__version__}\")","metadata":{"execution":{"iopub.status.busy":"2023-03-27T04:09:49.522931Z","iopub.execute_input":"2023-03-27T04:09:49.523643Z","iopub.status.idle":"2023-03-27T04:09:49.529745Z","shell.execute_reply.started":"2023-03-27T04:09:49.523607Z","shell.execute_reply":"2023-03-27T04:09:49.528563Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seed_settings(seed=57):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    # torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n\nseed_settings(seed=57)","metadata":{"execution":{"iopub.status.busy":"2023-03-27T04:09:49.531458Z","iopub.execute_input":"2023-03-27T04:09:49.531946Z","iopub.status.idle":"2023-03-27T04:09:49.54657Z","shell.execute_reply.started":"2023-03-27T04:09:49.531864Z","shell.execute_reply":"2023-03-27T04:09:49.545528Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## function to generate graph_dataset\n## construct edge from the observed time series of pulses by knn_graph()\n## each graph(_id) corresponds to an event(_id)\n## each node(_id) corresponds to the pulse observation order(row_index)\ndef get_graph_datasets(graph_ids, batch, pulse_limit, n_neighbors, df_label=None):\n    \"\"\"\n    A data object describing a homogeneous graph\n    (generate/return test dataset if df_label=None)\n    \"\"\"\n    datasets = []\n    pulse_limit = pulse_limit\n    for graph_id in tqdm(graph_ids):\n        node_feature = (batch\n                        .filter(pl.col(\"event_id\") == graph_id)\n                        .select([\"x\", \"y\", \"z\", \"elapsed_time\", \"charge\", \"auxiliary\"])\n                       ).to_numpy()\n        # Limit number of pulses in the large events\n        if node_feature.shape[0] > pulse_limit:\n            node_feature = node_feature[0:pulse_limit, :]\n        else:\n            pass\n        \n        if df_label is None:\n            data = Data(x=torch.tensor(node_feature, dtype=torch.float),\n                        n_pulses=torch.tensor(node_feature.shape[0], dtype=torch.int),\n                       )\n            # construct edge from the k-nearest neighbours.\n            data.edge_index = knn_graph(\n                data.x,  # node(sensor) features\n                k=n_neighbors,  # The number of neighbors\n                cosine=False,  # cos_distance or euclidean_distance\n                loop=False\n            )\n        else:\n            y = df_label.filter(pl.col(\"event_id\") == graph_id).select([\"azimuth\",\"zenith\"]).to_numpy()[0]\n            \n            data = Data(x=torch.tensor(node_feature, dtype=torch.float),\n                        n_pulses=torch.tensor(node_feature.shape[0], dtype=torch.int),\n                        y=torch.tensor(y, dtype=torch.float))\n            # construct edge from the k-nearest neighbours.\n            data.edge_index = knn_graph(\n                data.x,  # node(sensor) features\n                k=n_neighbors,  # The number of neighbors\n                cosine=False,  # cos_distance or euclidean_distance\n                loop=False\n            )\n        \n        datasets.append(data)\n    return datasets","metadata":{"execution":{"iopub.status.busy":"2023-03-27T04:09:49.547932Z","iopub.execute_input":"2023-03-27T04:09:49.548269Z","iopub.status.idle":"2023-03-27T04:09:49.56123Z","shell.execute_reply.started":"2023-03-27T04:09:49.548241Z","shell.execute_reply":"2023-03-27T04:09:49.55995Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n%%time\n## generate graph_dataset from event/sensor_time-series dataframe\n## Each graph(_id) corresponds to an event(_id)\n## Each node(_id) corresponds to the pulse observation order(row_index)\ngraph_ids = df_label[\"event_id\"].to_numpy()\npulse_limit = 220  # up to 220 pulses covering 95% of events\nprint('Converting to graph data...')\ndatasets = get_graph_datasets(graph_ids=graph_ids,\n                                 batch=batch,\n                                 pulse_limit=pulse_limit,\n                                 n_neighbors=5,\n                                 df_label=df_label)\nprint('...converted')","metadata":{"execution":{"iopub.status.busy":"2023-03-27T04:09:49.5629Z","iopub.execute_input":"2023-03-27T04:09:49.563235Z","iopub.status.idle":"2023-03-27T04:09:53.63278Z","shell.execute_reply.started":"2023-03-27T04:09:49.563205Z","shell.execute_reply":"2023-03-27T04:09:53.631621Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## display the head/sample of event converted to graph data\n# Data(x=, y=, n_pulses(of raw observations)=, edge_index=)\nfor i in range(10):\n    print(i, datasets[i])","metadata":{"execution":{"iopub.status.busy":"2023-03-27T04:09:53.63415Z","iopub.execute_input":"2023-03-27T04:09:53.634481Z","iopub.status.idle":"2023-03-27T04:09:53.641466Z","shell.execute_reply.started":"2023-03-27T04:09:53.634448Z","shell.execute_reply":"2023-03-27T04:09:53.640496Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## display (2D)graph of event data sample\nfrom torch_geometric.utils import to_networkx\n\nnxg = to_networkx(datasets[0])\nnx.draw(nxg, with_labels=True)","metadata":{"execution":{"iopub.status.busy":"2023-03-27T04:09:53.642893Z","iopub.execute_input":"2023-03-27T04:09:53.643186Z","iopub.status.idle":"2023-03-27T04:09:55.885852Z","shell.execute_reply.started":"2023-03-27T04:09:53.643157Z","shell.execute_reply":"2023-03-27T04:09:55.884522Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## graph data/shape summarize\nn_graph_node, n_graph_edge = [], [] \nfor i in range(len(datasets)):\n    n_graph_node.append(datasets[i].num_nodes)\n    n_graph_edge.append(datasets[i].num_edges)\ndf_graphs = pl.DataFrame([n_graph_node, n_graph_edge],\n                         schema=[(\"n_graph_node\", pl.Int16), (\"n_graph_edge\", pl.Int16)])\ndf_graphs = pl.concat([df_label, df_graphs], how=\"horizontal\")\ndisplay(df_graphs)\n\n# graph size visualization\nfig, axes = plt.subplots(1, 2, figsize=(12,4), tight_layout=True)\naxes[0].hist(df_graphs[\"n_graph_node\"], color='black', alpha=0.5, bins=50, log=True)\naxes[0].set_title('number of graph_nodes')\naxes[0].grid()\naxes[1].hist(df_graphs[\"n_graph_edge\"], color='black', alpha=0.5, bins=50, log=True)\naxes[1].set_title('number of graph_edges')\naxes[1].grid()\nplt.show()\n\ndisplay(df_graphs.describe())","metadata":{"execution":{"iopub.status.busy":"2023-03-27T04:09:55.887574Z","iopub.execute_input":"2023-03-27T04:09:55.887926Z","iopub.status.idle":"2023-03-27T04:09:56.962872Z","shell.execute_reply.started":"2023-03-27T04:09:55.887889Z","shell.execute_reply":"2023-03-27T04:09:56.961927Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## G(C)NN:Graph(Convolutional)NeuralNetwork model design","metadata":{}},{"cell_type":"code","source":"class domNet(torch.nn.Module):\n    def __init__(self):\n        super(domNet, self).__init__()\n        self.n_features = 6  # node(sensor) feature(observed charge & geometry) value\n        self.n_conv_hidden = 2\n        self.n_mlp_hidden = 3\n        self.hidden_dim = 64\n        self.n_outputs = 2  # targets values\n        self.graphconv1 = GCNConv(self.n_features, self.hidden_dim)\n        self.bn1 = BatchNorm1d(self.hidden_dim)\n        self.graphconv_hidden = ModuleList(\n            [GCNConv(self.hidden_dim, self.hidden_dim, cached=False) for _ in range(self.n_conv_hidden)]\n        )\n        self.bn_conv = ModuleList(\n            [BatchNorm1d(self.hidden_dim) for _ in range(self.n_conv_hidden)]\n        )\n        self.mlp_hidden =  ModuleList(\n            [Linear(self.hidden_dim, self.hidden_dim) for _ in range(self.n_mlp_hidden)]\n        )\n        self.bn_mlp = ModuleList(\n            [BatchNorm1d(self.hidden_dim) for _ in range(self.n_mlp_hidden)]\n        )\n        self.mlp_out = Linear(self.hidden_dim, self.n_outputs)\n\n    def forward(self, data):\n        x, edge_index = data.x, data.edge_index  # (pulse_in_batch, n_feature)\n        x = F.relu(self.graphconv1(x, edge_index))\n        x = self.bn1(x)\n        for graphconv, bn_conv in zip(self.graphconv_hidden, self.bn_conv):\n            x = graphconv(x, edge_index)\n            x = bn_conv(x)\n        x = global_add_pool(x, data.batch)\n        for fc_mlp, bn_mlp in zip(self.mlp_hidden, self.bn_mlp):\n            x = F.relu(fc_mlp(x))\n            x = bn_mlp(x)\n            x = F.dropout(x, p=0.1, training=self.training)\n        x = self.mlp_out(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2023-03-27T04:09:56.964289Z","iopub.execute_input":"2023-03-27T04:09:56.964803Z","iopub.status.idle":"2023-03-27T04:09:56.977191Z","shell.execute_reply.started":"2023-03-27T04:09:56.964771Z","shell.execute_reply":"2023-03-27T04:09:56.97587Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## split/load dataset\n# train_val_split(train:val=8:2)\ndatasets_train = [datasets[i] for i in range(len(datasets)) if i % 5 != 0]\ndatasets_val = [datasets[i] for i in range(len(datasets)) if i % 5 == 0]\nprint(f\"n_train={len(datasets_train)}, n_val={len(datasets_val)}\")\n\ndataloader_train = DataLoader(datasets_train, batch_size=256, shuffle=False)\ndataloader_val = DataLoader(datasets_val, batch_size=256, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2023-03-27T04:09:56.978335Z","iopub.execute_input":"2023-03-27T04:09:56.97864Z","iopub.status.idle":"2023-03-27T04:09:56.99492Z","shell.execute_reply.started":"2023-03-27T04:09:56.978613Z","shell.execute_reply":"2023-03-27T04:09:56.993832Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## Scoring functions\n# mean-angular-error\ndef angular_dist_score(az_true, zen_true, az_pred, zen_pred):\n    if not (np.all(np.isfinite(az_true)) and\n            np.all(np.isfinite(zen_true)) and\n            np.all(np.isfinite(az_pred)) and\n            np.all(np.isfinite(zen_pred))):\n        raise ValueError(\"All arguments must be finite\")\n    sa1 = np.sin(az_true)\n    ca1 = np.cos(az_true)\n    sz1 = np.sin(zen_true)\n    cz1 = np.cos(zen_true)\n    sa2 = np.sin(az_pred)\n    ca2 = np.cos(az_pred)\n    sz2 = np.sin(zen_pred)\n    cz2 = np.cos(zen_pred)\n    scalar_prod = sz1*sz2*(ca1*ca2 + sa1*sa2) + (cz1*cz2)\n    scalar_prod =  np.clip(scalar_prod, -1, 1)\n    return np.average(np.abs(np.arccos(scalar_prod)))","metadata":{"execution":{"iopub.status.busy":"2023-03-27T04:09:56.996621Z","iopub.execute_input":"2023-03-27T04:09:56.996967Z","iopub.status.idle":"2023-03-27T04:09:57.013479Z","shell.execute_reply.started":"2023-03-27T04:09:56.996938Z","shell.execute_reply":"2023-03-27T04:09:57.012264Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## train/valid loop settings\ndef train(model, \n          dataloader_train, \n          dataloader_val, \n          optimizer,\n          criterion, \n          epoch):\n    \n    ret = {}\n    \n    # ====================\n    # training\n    # ====================\n    y_a, y_a_outputs, y_z, y_z_outputs = [], [], [], []\n    epoch_loss = 0\n    model.train()\n    for batch in dataloader_train:\n        batch = batch.to(device)\n        optimizer.zero_grad()\n        output = model(batch)\n        loss = criterion(output, batch.y.reshape(-1, 2))\n        loss.backward()\n        optimizer.step()\n        epoch_loss += loss.item()\n        \n        y_a.extend(batch.y.reshape(-1, 2)[:,0].tolist())\n        y_a_outputs.extend(output.detach()[:,0].tolist())\n        y_z.extend(batch.y.reshape(-1, 2)[:,1].tolist())\n        y_z_outputs.extend(output.detach()[:,1].tolist())\n    ang_dist_score = angular_dist_score(y_a, y_z, y_a_outputs, y_z_outputs)\n    \n    ret[\"train_loss\"] = epoch_loss / len(dataloader_train)\n    ret[\"mean_angular_error_train\"] = ang_dist_score\n    ret[\"azimuth_true_train\"] = y_a\n    ret[\"azimuth_pred_train\"] = y_a_outputs\n    ret[\"zenith_true_train\"] = y_z\n    ret[\"zenith_pred_train\"] = y_z_outputs\n\n    # ====================\n    # validation\n    # ====================\n    y_a, y_a_preds, y_z, y_z_preds = [], [], [], []\n    epoch_loss = 0\n    model.eval()\n    for batch in dataloader_val:\n        batch = batch.to(device)\n        pred = model(batch)\n        loss = criterion(pred, batch.y.reshape(-1, 2))\n        epoch_loss += loss.item()\n        \n        y_a.extend(batch.y.reshape(-1, 2)[:,0].tolist())\n        y_a_preds.extend(pred.detach()[:,0].tolist())\n        y_z.extend(batch.y.reshape(-1, 2)[:,1].tolist())\n        y_z_preds.extend(pred.detach()[:,1].tolist())\n    ang_dist_score = angular_dist_score(y_a, y_z, y_a_preds, y_z_preds)\n\n    ret[\"val_loss\"] = epoch_loss / len(dataloader_val)\n    ret[\"mean_angular_error_valid\"] = ang_dist_score\n    ret[\"azimuth_true_valid\"] = y_a\n    ret[\"azimuth_pred_valid\"] = y_a_preds\n    ret[\"zenith_true_valid\"] = y_z\n    ret[\"zenith_pred_valid\"] = y_z_preds\n    \n    return ret","metadata":{"execution":{"iopub.status.busy":"2023-03-27T04:09:57.014877Z","iopub.execute_input":"2023-03-27T04:09:57.01531Z","iopub.status.idle":"2023-03-27T04:09:57.033184Z","shell.execute_reply.started":"2023-03-27T04:09:57.015264Z","shell.execute_reply":"2023-03-27T04:09:57.032009Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## train/valid settings/parameters\nn_epochs = 200\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint('using ', device)\nmodel = domNet().to(device)\ncriterion = torch.nn.L1Loss()\n#criterion = torch.nn.MSELoss()\noptimizer = torch.optim.Adam(model.parameters(), lr=0.01)","metadata":{"execution":{"iopub.status.busy":"2023-03-27T04:09:57.034917Z","iopub.execute_input":"2023-03-27T04:09:57.035871Z","iopub.status.idle":"2023-03-27T04:09:57.056557Z","shell.execute_reply.started":"2023-03-27T04:09:57.035823Z","shell.execute_reply":"2023-03-27T04:09:57.055343Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n%%time\n## train/valid domNET(GCN+MLP)\n# parameter initialization reset\nfinal_model = None\nbest_angular_error = 10\nbest_epoch = -1\n\nfor epoch in range(n_epochs):\n    t = time.time()\n    ret = train(model=model, \n                dataloader_train=dataloader_train,\n                dataloader_val=dataloader_val,\n                optimizer=optimizer, \n                criterion=criterion, \n                epoch=epoch)\n    consumed_time = time.time() - t\n    if (epoch % 50 == 0) or (epoch == n_epochs-1):\n        print(\"epoch:{:4} ({:.1f}s)[train] loss={:.4f}, mean_angular_error={:.4f} [valid] loss={:.4f}, mean_angular_error={:.4f}\".format(\n            epoch,\n            consumed_time,\n            ret[\"train_loss\"],\n            ret[\"mean_angular_error_train\"],\n            ret[\"val_loss\"],\n            ret[\"mean_angular_error_valid\"],\n        )\n             )\n\n    if ret[\"mean_angular_error_valid\"] < best_angular_error:\n        final_model = copy.deepcopy(model)\n        best_angular_error = ret[\"mean_angular_error_valid\"]\n        best_epoch = epoch\n\nprint(\"result: best_angular_error={:.4f}(epoch={})\".format(best_angular_error, best_epoch))","metadata":{"execution":{"iopub.status.busy":"2023-03-27T04:09:57.063222Z","iopub.execute_input":"2023-03-27T04:09:57.063572Z","iopub.status.idle":"2023-03-27T04:15:37.842193Z","shell.execute_reply.started":"2023-03-27T04:09:57.063541Z","shell.execute_reply":"2023-03-27T04:15:37.841323Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## neutrino_direction prediction results","metadata":{}},{"cell_type":"code","source":"## observed/predicted angles(azimuth & zenith)\nazm_ob_tr, azm_pr_tr = np.array(ret['azimuth_true_train']), np.array(ret['azimuth_pred_train'])\nzen_ob_tr, zen_pr_tr = np.array(ret['zenith_true_train']), np.array(ret['zenith_pred_train'])\nazm_ob_va, azm_pr_va = np.array(ret['azimuth_true_valid']), np.array(ret['azimuth_pred_valid'])\nzen_ob_va, zen_pr_va = np.array(ret['zenith_true_valid']), np.array(ret['zenith_pred_valid'])","metadata":{"execution":{"iopub.status.busy":"2023-03-27T04:15:37.843602Z","iopub.execute_input":"2023-03-27T04:15:37.844166Z","iopub.status.idle":"2023-03-27T04:15:37.851071Z","shell.execute_reply.started":"2023-03-27T04:15:37.844133Z","shell.execute_reply":"2023-03-27T04:15:37.850005Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## event_level prediction(train/valid)\n# train data prediction\nang_dist_score_tr = np.empty(0)\nfor i in range(azm_ob_tr.shape[0]):\n    ang_dist_score_tr = np.append(ang_dist_score_tr, angular_dist_score(azm_ob_tr[i], zen_ob_tr[i], azm_pr_tr[i], zen_pr_tr[i]))\n# valid data prediction\nang_dist_score_va = np.empty(0)\nfor i in range(azm_ob_va.shape[0]):\n    ang_dist_score_va = np.append(ang_dist_score_va, angular_dist_score(azm_ob_va[i], zen_ob_va[i], azm_pr_va[i], zen_pr_va[i]))\n\n# histgram of mean_angular_error\nplt.figure(figsize=(12,4))\nplt.hist(ang_dist_score_tr, bins=50, color='blue', alpha=0.5)\nplt.hist(ang_dist_score_va, bins=25, color='red', alpha=0.5)\nplt.title('mean_angular_error (train/valid)')\nplt.grid()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-03-27T04:15:37.852641Z","iopub.execute_input":"2023-03-27T04:15:37.852978Z","iopub.status.idle":"2023-03-27T04:15:38.337251Z","shell.execute_reply.started":"2023-03-27T04:15:37.852933Z","shell.execute_reply":"2023-03-27T04:15:38.335877Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, axes = plt.subplots(2, 3, figsize=(18,12), tight_layout=True)\n## observed/predicted azimuth\naxes[0][0].scatter(azm_ob_tr, azm_pr_tr, color='blue', alpha=0.2)\naxes[0][0].plot([azm_ob_tr.min(), azm_ob_tr.max()], [azm_ob_tr.min(), azm_ob_tr.max()])\naxes[0][0].set_title('azimuth_train')\naxes[0][0].set_xlabel(\"Observed\")\naxes[0][0].set_ylabel(\"Predicted\")\naxes[0][0].grid()\n\naxes[0][1].scatter(azm_ob_va, azm_pr_va, color='firebrick', alpha=0.2)\naxes[0][1].plot([azm_ob_va.min(), azm_ob_va.max()], [azm_ob_va.min(), azm_ob_va.max()])\naxes[0][1].set_title('azimuth_valid')\naxes[0][1].set_xlabel(\"Observed\")\naxes[0][1].set_ylabel(\"Predicted\")\naxes[0][1].grid()\n\naxes[0][2].scatter(azm_ob_tr, azm_pr_tr, color='cyan', alpha=0.2)\naxes[0][2].scatter(azm_ob_va, azm_pr_va, color='tomato', alpha=0.2)\naxes[0][2].plot([azm_ob_va.min(), azm_ob_va.max()], [azm_ob_va.min(), azm_ob_va.max()])\naxes[0][2].set_title('azimuth_train/valid')\naxes[0][2].set_xlabel(\"Observed\")\naxes[0][2].set_ylabel(\"Predicted\")\naxes[0][2].grid()\n\n## observed/predicted zenith\naxes[1][0].scatter(zen_ob_tr, zen_pr_tr, color='green', alpha=0.2)\naxes[1][0].plot([zen_ob_tr.min(), zen_ob_tr.max()], [zen_ob_tr.min(), zen_ob_tr.max()])\naxes[1][0].set_title('zenith_train')\naxes[1][0].set_xlabel(\"Observed\")\naxes[1][0].set_ylabel(\"Predicted\")\naxes[1][0].grid()\n\naxes[1][1].scatter(zen_ob_va, zen_pr_va, color='orange', alpha=0.2)\naxes[1][1].plot([zen_ob_va.min(), zen_ob_va.max()], [zen_ob_va.min(), zen_ob_va.max()])\naxes[1][1].set_title('zenith_valid')\naxes[1][1].set_xlabel(\"Observed\")\naxes[1][1].set_ylabel(\"Predicted\")\naxes[1][1].grid()\n\naxes[1][2].scatter(zen_ob_tr, zen_pr_tr, color='lime', alpha=0.2)\naxes[1][2].scatter(zen_ob_va, zen_pr_va, color='gold', alpha=0.2)\naxes[1][2].plot([zen_ob_va.min(), zen_ob_va.max()], [zen_ob_va.min(), zen_ob_va.max()])\naxes[1][2].set_title('zenith_train/valid')\naxes[1][2].set_xlabel(\"Observed\")\naxes[1][2].set_ylabel(\"Predicted\")\naxes[1][2].grid()\n\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-03-27T04:15:38.339393Z","iopub.execute_input":"2023-03-27T04:15:38.34023Z","iopub.status.idle":"2023-03-27T04:15:39.37262Z","shell.execute_reply.started":"2023-03-27T04:15:38.340183Z","shell.execute_reply":"2023-03-27T04:15:39.371542Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n%%time\n## train/predicted results\n#model = final_model\nmodel.eval()\nazimuth_preds, zenith_preds = [], []\nfor batch in datasets:\n    batch = batch.to(device)\n   #pred = final_model(batch)\n    pred = model(batch)\n    azimuth_preds.extend(pred.detach()[:,0].tolist())\n    zenith_preds.extend(pred.detach()[:,1].tolist())\n\n# merge graph/event info\ndf_preds = pl.DataFrame([azimuth_preds, zenith_preds],\n                         schema=[(\"predicted_azimuth\", pl.Float32), (\"predicted_zenith\", pl.Float32)])\ndf_preds = (pl.concat([df_label, df_preds], how=\"horizontal\")\n            .with_columns([\n                (pl.col(\"azimuth\") - pl.col(\"predicted_azimuth\")).abs().alias(\"error_az\"),\n                (pl.col(\"zenith\") - pl.col(\"predicted_zenith\")).abs().alias(\"error_ze\"),\n            ])\n           )\n\nprint(\"df_preds:\")\ndisplay(df_preds)\nprint(\"df_preds_describe():\")\ndisplay(df_preds.describe())","metadata":{"execution":{"iopub.status.busy":"2023-03-27T04:15:39.373902Z","iopub.execute_input":"2023-03-27T04:15:39.374243Z","iopub.status.idle":"2023-03-27T04:15:44.326699Z","shell.execute_reply.started":"2023-03-27T04:15:39.374216Z","shell.execute_reply":"2023-03-27T04:15:44.325597Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# hard/easy predicting events\nhard_to_pred_event = df_preds.filter((pl.col(\"error_az\") > 4) & (pl.col(\"error_ze\") > 1))\neasy_to_pred_event = df_preds.filter((pl.col(\"error_az\") < 0.05) & (pl.col(\"error_ze\") < 0.02))\nprint(\"hard_to_pred_event\")\ndisplay(hard_to_pred_event)\nprint(\"easy_to_pred_event\")\ndisplay(easy_to_pred_event)\n# hard/easy event list\nhard_event_id = hard_to_pred_event[\"event_id\"].to_list()\neasy_event_id = easy_to_pred_event[\"event_id\"].to_list()\n# hard/easy graphs\nhard_graphs = df_graphs.filter(pl.col(\"event_id\").is_in(hard_event_id))\neasy_graphs = df_graphs.filter(pl.col(\"event_id\").is_in(easy_event_id))\nprint(\"hard_to_pred_eventgraph\")\ndisplay(hard_graphs)\nprint(\"easy_to_pred_eventgraph\")\ndisplay(easy_graphs)","metadata":{"execution":{"iopub.status.busy":"2023-03-27T04:15:44.32818Z","iopub.execute_input":"2023-03-27T04:15:44.328539Z","iopub.status.idle":"2023-03-27T04:15:44.354053Z","shell.execute_reply.started":"2023-03-27T04:15:44.328497Z","shell.execute_reply":"2023-03-27T04:15:44.353139Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## neutrino_direction sensor/event visualization","metadata":{}},{"cell_type":"code","source":"%%time\n! pip install --quiet plotly==5.11.0","metadata":{"execution":{"iopub.status.busy":"2023-03-27T04:15:44.355187Z","iopub.execute_input":"2023-03-27T04:15:44.355534Z","iopub.status.idle":"2023-03-27T04:16:35.623419Z","shell.execute_reply.started":"2023-03-27T04:15:44.355503Z","shell.execute_reply":"2023-03-27T04:16:35.621883Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import plotly\nimport plotly.express as px\nimport plotly.graph_objects as go\nfrom plotly.subplots import make_subplots","metadata":{"execution":{"iopub.status.busy":"2023-03-27T04:16:35.625441Z","iopub.execute_input":"2023-03-27T04:16:35.626247Z","iopub.status.idle":"2023-03-27T04:16:36.968227Z","shell.execute_reply.started":"2023-03-27T04:16:35.626197Z","shell.execute_reply":"2023-03-27T04:16:36.966585Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n%%time\n# reload sensor geometry w/ pandas\nsensor_geometry = pd.read_csv(f'{PATH_DATASET}/sensor_geometry.csv')\n# reload selected batch w/ pandas\nbatch_id = 1\ntrain_batch = pd.read_parquet(f'{PATH_DATASET}/{MODE}/batch_{batch_id}.parquet')\nlist_tr_batch = df_label[\"event_id\"].to_list()\ntrain_batch.query('event_id in @list_tr_batch')\n# reload meta data(train) from polars to pandas\ntrain_meta = df_preds.to_pandas()","metadata":{"execution":{"iopub.status.busy":"2023-03-27T04:16:36.969715Z","iopub.execute_input":"2023-03-27T04:16:36.970052Z","iopub.status.idle":"2023-03-27T04:16:38.960706Z","shell.execute_reply.started":"2023-03-27T04:16:36.970018Z","shell.execute_reply":"2023-03-27T04:16:38.959467Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# function for events visualization\ndef visualize_event(pulse_id: str, mode: str, sensor_geometry: pd.DataFrame, pulses_data: pd.DataFrame, meta_data: pd.DataFrame):\n    event_pulses = pulses_data[pulses_data.index == pulse_id]\n    meta_of_event = train_meta[train_meta[\"event_id\"] == pulse_id]\n    \n    fig = make_subplots(\n        rows=1, cols=2,\n        specs=[[{\"type\": \"scatter3d\"}, {\"type\": \"scatter3d\"}]],\n        subplot_titles=(\"All events\", \"Not auxiliary events\")\n    )\n    \n    # data_frame auxiliary or not\n    aux_pulses_data = event_pulses\n   #aux_pulses_data = event_pulses[event_pulses[\"auxiliary\"] == True]\n    not_aux_pulses_data = event_pulses[event_pulses[\"auxiliary\"] == False]\n    aux_df = aux_pulses_data.merge(sensor_geometry, left_on='sensor_id', right_on='sensor_id')[[\"x\", \"y\", \"z\", \"charge\", \"time\"]]\n    not_aux_df = not_aux_pulses_data.merge(sensor_geometry, left_on='sensor_id', right_on='sensor_id')[[\"x\", \"y\", \"z\", \"charge\", \"time\"]]\n\n    # plot all caharge~plot_size, time~plot_color\n    fig.add_trace(\n        go.Scatter3d(x=aux_df[\"x\"], y=aux_df[\"y\"], z=aux_df[\"z\"], opacity=0.75, \n                     mode='markers', marker_size=aux_df[\"charge\"]*10, text=aux_df[\"charge\"],\n                     marker=dict(color=aux_df[\"time\"], cmin=0, cmax=aux_df.iloc[-1][\"time\"], colorscale='Temps', showscale=True)),\n        row=1, col=1)\n    # plot not auxiliary only caharge~plot_size, time~plot_color\n    fig.add_trace(\n        go.Scatter3d(x=not_aux_df[\"x\"], y=not_aux_df[\"y\"], z=not_aux_df[\"z\"], opacity=0.75,\n                     mode='markers', marker_size=not_aux_df[\"charge\"]*10, text=aux_df[\"charge\"],\n                     marker=dict(color=not_aux_df[\"time\"], cmin=0, cmax=not_aux_df.iloc[-1][\"time\"], colorscale='Temps', showscale=True)),\n        row=1, col=2)\n    # display sensors by gray_dot\n    fig.add_trace(\n        go.Scatter3d(x=sensor_geometry[\"x\"], y=sensor_geometry[\"y\"], z=sensor_geometry[\"z\"], mode='markers',\n                     opacity=0.3, marker=dict(size=1, color=\"gray\")), row=1, col=1)\n    # display sensors by gray_dot\n    fig.add_trace(\n        go.Scatter3d(x=sensor_geometry[\"x\"], y=sensor_geometry[\"y\"], z=sensor_geometry[\"z\"], mode='markers',\n                     opacity=0.3, marker=dict(size=1, color=\"gray\")), row=1, col=2)\n    # convert angle to coordinates(observed)\n    azimuth, zenith = meta_of_event[\"azimuth\"].values[0], meta_of_event[\"zenith\"].values[0]\n    true_x = math.cos(azimuth) * math.sin(zenith)\n    true_y = math.sin(azimuth) * math.sin(zenith)\n    true_z = math.cos(zenith)\n    # convert angle to coordinates(predicted)\n    p_azimuth, p_zenith = meta_of_event[\"predicted_azimuth\"].values[0], meta_of_event[\"predicted_zenith\"].values[0]\n    pred_x = math.cos(p_azimuth) * math.sin(p_zenith)\n    pred_y = math.sin(p_azimuth) * math.sin(p_zenith)\n    pred_z = math.cos(p_zenith)\n    \n    # display direction trajectory by red_line\n    if mode == \"hard\":\n        # plot observed direction_line\n        fig.add_trace(\n            go.Scatter3d(\n                x=[-true_x * 500, true_x * 500], y=[-true_y * 500, true_y * 500], z=[-true_z * 500, true_z * 500],\n                opacity=0.8, mode='lines', line=dict(color='red', width=5)\n            ), row=1, col=1)\n        fig.add_trace(\n            go.Scatter3d(\n                x=[-true_x * 500, true_x * 500], y=[-true_y * 500, true_y * 500], z=[-true_z * 500, true_z * 500],\n                opacity=0.8, mode='lines', line=dict(color='red', width=5)\n            ), row=1, col=2)\n        # plot predicted direction_line\n        fig.add_trace(\n            go.Scatter3d(\n                x=[-pred_x * 500, pred_x * 500], y=[-pred_y * 500, pred_y * 500], z=[-pred_z * 500, pred_z * 500],\n                opacity=0.8, mode='lines', line=dict(color='gold', width=5)\n            ), row=1, col=1)\n        fig.add_trace(\n            go.Scatter3d(\n                x=[-pred_x * 500, pred_x * 500], y=[-pred_y * 500, pred_y * 500], z=[-pred_z * 500, pred_z * 500],\n                opacity=0.8, mode='lines', line=dict(color='gold', width=5)\n            ), row=1, col=2)\n    elif mode == \"easy\":\n        # plot observed direction_line\n        fig.add_trace(\n            go.Scatter3d(\n                x=[-true_x * 500, true_x * 500], y=[-true_y * 500, true_y * 500], z=[-true_z * 500, true_z * 500],\n                opacity=0.8, mode='lines', line=dict(color='green', width=5)\n            ), row=1, col=1)\n        fig.add_trace(\n            go.Scatter3d(\n                x=[-true_x * 500, true_x * 500], y=[-true_y * 500, true_y * 500], z=[-true_z * 500, true_z * 500],\n                opacity=0.8, mode='lines', line=dict(color='green', width=5)\n            ), row=1, col=2)\n        # plot predicted direction_line\n        fig.add_trace(\n            go.Scatter3d(\n                x=[-pred_x * 500, pred_x * 500], y=[-pred_y * 500, pred_y * 500], z=[-pred_z * 500, pred_z * 500],\n                opacity=0.8, mode='lines', line=dict(color='lime', width=5)\n            ), row=1, col=1)\n        fig.add_trace(\n            go.Scatter3d(\n                x=[-pred_x * 500, pred_x * 500], y=[-pred_y * 500, pred_y * 500], z=[-pred_z * 500, pred_z * 500],\n                opacity=0.8, mode='lines', line=dict(color='lime', width=5)\n            ), row=1, col=2)\n    else:\n        pass\n    \n    fig.update_layout(title_text=f\"Event {pulse_id}, {mode} to fit_predict\")\n    \n    fig.show()","metadata":{"execution":{"iopub.status.busy":"2023-03-27T04:16:38.962423Z","iopub.execute_input":"2023-03-27T04:16:38.962751Z","iopub.status.idle":"2023-03-27T04:16:39.000074Z","shell.execute_reply.started":"2023-03-27T04:16:38.962723Z","shell.execute_reply":"2023-03-27T04:16:38.998822Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# event list\nprint(hard_event_id)\nprint(easy_event_id)","metadata":{"execution":{"iopub.status.busy":"2023-03-27T04:16:39.001432Z","iopub.execute_input":"2023-03-27T04:16:39.001902Z","iopub.status.idle":"2023-03-27T04:16:39.014613Z","shell.execute_reply.started":"2023-03-27T04:16:39.001868Z","shell.execute_reply":"2023-03-27T04:16:39.013666Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 3dplot visualize event list samples\n# true label: red line / predicted: orange line\nfor hard_evt in hard_event_id[0:4]:\n    visualize_event(hard_evt, \"hard\", sensor_geometry, train_batch, train_meta)","metadata":{"execution":{"iopub.status.busy":"2023-03-27T04:16:39.015711Z","iopub.execute_input":"2023-03-27T04:16:39.016013Z","iopub.status.idle":"2023-03-27T04:16:39.301122Z","shell.execute_reply.started":"2023-03-27T04:16:39.015987Z","shell.execute_reply":"2023-03-27T04:16:39.300082Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 3dplot visualize event list samples\n# true label: green line / predicted: light green line\nfor easy_evt in easy_event_id[0:4]:\n    visualize_event(easy_evt, \"easy\", sensor_geometry, train_batch, train_meta)","metadata":{"execution":{"iopub.status.busy":"2023-03-27T04:16:39.302729Z","iopub.execute_input":"2023-03-27T04:16:39.303618Z","iopub.status.idle":"2023-03-27T04:16:39.387391Z","shell.execute_reply.started":"2023-03-27T04:16:39.303577Z","shell.execute_reply":"2023-03-27T04:16:39.386259Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}