{"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":"# Graph Neural Networks on Competition Data\n\n<img style=\"float: right;\" src=\"https://raw.githubusercontent.com/graphnet-team/graphnet/main/assets/identity/graphnet-logo-and-wordmark.png\" width=\"600\" height=\"600\" />\n\nThis notebook is a copy of graphnet-example and contains everything needed to tinker with GraphNeT on the competition data. It includes\n\n1. Installation instructions for GraphNeT\n\n2. Code for data conversion \n\n3. Snippets for training dynedge similarly to whats shown in the JINST paper\n\n4. A pre-trained dynedge on batch 1 to 50.\n\n5. Snippets for inference and evaluation of results.","metadata":{}},{"cell_type":"markdown","source":"## Installing GraphNeT\n\nYou can find the official installation instructions for GraphNeT [here](https://github.com/graphnet-team/graphnet#gear--install)\nThis code contains a few extra steps to get the library installed in a Kaggle notebook. It will copy a recent version of graphnet and it's dependencies to the working disk and install the GPU version of GraphNeT in /kaggle/working/software.","metadata":{}},{"cell_type":"code","source":"# Move software to working disk\n!rm  -r software\n!scp -r /kaggle/input/graphnet-and-dependencies/software .\n\n# Install dependencies\n!pip install /kaggle/working/software/dependencies/torch-1.11.0+cu115-cp37-cp37m-linux_x86_64.whl\n!pip install /kaggle/working/software/dependencies/torch_cluster-1.6.0-cp37-cp37m-linux_x86_64.whl\n!pip install /kaggle/working/software/dependencies/torch_scatter-2.0.9-cp37-cp37m-linux_x86_64.whl\n!pip install /kaggle/working/software/dependencies/torch_sparse-0.6.13-cp37-cp37m-linux_x86_64.whl\n!pip install /kaggle/working/software/dependencies/torch_geometric-2.0.4.tar.gz\n\n# Install GraphNeT\n!cd software/graphnet;pip install --no-index --find-links=\"/kaggle/working/software/dependencies\" -e .[torch]\n\n# Append to PATH\nimport sys\nsys.path.append('/kaggle/working/software/graphnet/src')\nprint ('Finished installing graphnet.')","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-04-09T22:43:08.984770Z","iopub.execute_input":"2023-04-09T22:43:08.985978Z","iopub.status.idle":"2023-04-09T22:47:04.670320Z","shell.execute_reply.started":"2023-04-09T22:43:08.985918Z","shell.execute_reply":"2023-04-09T22:47:04.666717Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import graphnet","metadata":{"execution":{"iopub.status.busy":"2023-04-09T23:01:55.938621Z","iopub.execute_input":"2023-04-09T23:01:55.940838Z","iopub.status.idle":"2023-04-09T23:01:56.086270Z","shell.execute_reply.started":"2023-04-09T23:01:55.940620Z","shell.execute_reply":"2023-04-09T23:01:56.084204Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Converting The Parquet Files to SQLite\n\nThe majority of functionality is tied to the SQLite data format. Therefore, to make GraphNeT compatible with the data provided in this competition, a small converter was included that reads a selection of batch_id's and writes them to a single database. The database will contain two tables\n\n* meta_table  : Contains the information associated with train_meta_data.parquet\n* pulse_table : Contains the information associated with the batch_n.parquet files, including detector geometry.\n\nboth tables are indexed according to *event_id* for efficiency. This will allow to extract information from the databases by simply writing\n\n```python \nimport pandas as pd\nimport sqlite3\n\nmy_event_id = 32\nmy_database = '/kaggle/working/sqlite/batch_01.db'\n\nwith sqlite3.connect(my_database) as conn:\n    # extracts meta data for event\n    meta_query = f'SELECT * FROM meta_table WHERE event_id ={my_event_id}'\n    meta_data = pd.read_sql(meta_query,conn)\n    \n    # extracts pulses / detector response for event\n    pulse_query = f'SELECT * FROM pulse_table WHERE event_id ={my_event_id}'\n    pulse_data = pd.read_sql(query,conn)\n```\n\nThe upside of the SQLite format is that you only have the events you want in memory. Downside is ... the conversion :-) \n","metadata":{}},{"cell_type":"code","source":"import pyarrow.parquet as pq\nimport sqlite3\nimport pandas as pd\nimport sqlalchemy\nfrom tqdm import tqdm\nimport os, shutil\nfrom typing import Any, Dict, List, Optional\nimport numpy as np\n\nfrom graphnet.data.sqlite.sqlite_utilities import create_table\n\ndef load_input(meta_batch: pd.DataFrame, input_data_folder: str) -> pd.DataFrame:\n        \"\"\"\n        Will load the corresponding detector readings associated with the meta data batch.\n        \"\"\"\n        batch_id = pd.unique(meta_batch['batch_id'])\n\n        assert len(batch_id) == 1, \"contains multiple batch_ids. Did you set the batch_size correctly?\"\n        \n        detector_readings = pd.read_parquet(path = f'{input_data_folder}/batch_{batch_id[0]}.parquet')\n        sensor_positions = geometry_table.loc[detector_readings['sensor_id'], ['x', 'y', 'z']]\n        sensor_positions.index = detector_readings.index\n\n        for column in sensor_positions.columns:\n            if column not in detector_readings.columns:\n                detector_readings[column] = sensor_positions[column]\n\n        detector_readings['auxiliary'] = detector_readings['auxiliary'].replace({True: 1, False: 0})\n        return detector_readings.reset_index()\n\ndef add_to_table(database_path: str,\n                      df: pd.DataFrame,\n                      table_name:  str,\n                      is_primary_key: bool,\n                      ) -> None:\n    \"\"\"Writes meta data to sqlite table. \n\n    Args:\n        database_path (str): the path to the database file.\n        df (pd.DataFrame): the dataframe that is being written to table.\n        table_name (str, optional): The name of the meta table. Defaults to 'meta_table'.\n        is_primary_key(bool): Must be True if each row of df corresponds to a unique event_id. Defaults to False.\n    \"\"\"\n    try:\n        create_table(   columns=  df.columns,\n                        database_path = database_path, \n                        table_name = table_name,\n                        integer_primary_key= is_primary_key,\n                        index_column = 'event_id')\n    except sqlite3.OperationalError as e:\n        if 'already exists' in str(e):\n            pass\n        else:\n            raise e\n    engine = sqlalchemy.create_engine(\"sqlite:///\" + database_path)\n    df.to_sql(table_name, con=engine, index=False, if_exists=\"append\", chunksize = 200000)\n    engine.dispose()\n    return\n\ndef convert_to_sqlite(meta_data_path: str,\n                      database_path: str,\n                      input_data_folder: str,\n                      batch_size: int = 200000,\n                      batch_ids: Optional[List[int]] = None,) -> None:\n    \"\"\"Converts a selection of the Competition's parquet files to a single sqlite database.\n\n    Args:\n        meta_data_path (str): Path to the meta data file.\n        batch_size (int): the number of rows extracted from meta data file at a time. Keep low for memory efficiency.\n        database_path (str): path to database. E.g. '/my_folder/data/my_new_database.db'\n        input_data_folder (str): folder containing the parquet input files.\n        batch_ids (List[int]): The batch_ids you want converted. Defaults to None (all batches will be converted)\n    \"\"\"\n    if batch_ids is None:\n        batch_ids = np.arange(1,661,1).to_list()\n    else: \n        assert isinstance(batch_ids,list), \"Variable 'batch_ids' must be list.\"\n    if not database_path.endswith('.db'):\n        database_path = database_path +'.db'\n    meta_data_iter = pq.ParquetFile(meta_data_path).iter_batches(batch_size = batch_size)\n    batch_id = 1\n    converted_batches = []\n    progress_bar = tqdm(total = len(batch_ids))\n    if batch_ids == [661]: #particular case of test batch (only one and numbered 661)\n        batch_id = 661\n    for meta_data_batch in meta_data_iter:\n        if batch_id in batch_ids:\n            meta_data_batch  = meta_data_batch.to_pandas()\n            add_to_table(database_path = database_path,\n                        df = meta_data_batch,\n                        table_name='meta_table',\n                        is_primary_key= True)\n            pulses = load_input(meta_batch=meta_data_batch, input_data_folder= input_data_folder)\n            del meta_data_batch # memory\n            add_to_table(database_path = database_path,\n                        df = pulses,\n                        table_name='pulse_table',\n                        is_primary_key= False)\n            del pulses # memory\n            progress_bar.update(1)\n            converted_batches.append(batch_id)\n        batch_id +=1\n        if len(batch_ids) == len(converted_batches):\n            break\n    progress_bar.close()\n    del meta_data_iter # memory\n    print(f'Conversion Complete!. Database available at\\n {database_path}')","metadata":{"execution":{"iopub.status.busy":"2023-04-09T23:02:01.062088Z","iopub.execute_input":"2023-04-09T23:02:01.062678Z","iopub.status.idle":"2023-04-09T23:02:10.163073Z","shell.execute_reply.started":"2023-04-09T23:02:01.062611Z","shell.execute_reply":"2023-04-09T23:02:10.161105Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"This notebook comes with both batch 1 and 51 converted to sqlite databases, so you don't have to convert them yourself. They were produced by running the following code:\n\n```python \ninput_data_folder = '/kaggle/input/icecube-neutrinos-in-deep-ice/train'\ngeometry_table = pd.read_csv('/kaggle/input/icecube-neutrinos-in-deep-ice/sensor_geometry.csv')\nmeta_data_path = '/kaggle/input/icecube-neutrinos-in-deep-ice/train_meta.parquet'\n\n#batch_1\ndatabase_path = '/kagge/working/batch_1'\nconvert_to_sqlite(meta_data_path,\n                  database_path=database_path,\n                  input_data_folder=input_data_folder,\n                  batch_ids = [1])\n\n#batch_51\ndatabase_path = '/kagge/working/batch_51'\nconvert_to_sqlite(meta_data_path,\n                  database_path=database_path,\n                  input_data_folder=input_data_folder,\n                  batch_ids = [51])\n```\n\nYou can convert multiple batches into a single database by adjusting *batch_id*. Notice that the compression of SQLite is far inferiour to parquet. The entire competition dataset will take up more than 1.5T of disk space in SQLite.\n\nInstead of producing these databases again, the next cell will copy them into /kaggle/working which is a faster disk.","metadata":{}},{"cell_type":"code","source":"for dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        filepath = os.path.join(dirname, filename)\n        if '.db' in filepath:\n            src_path = filepath\n            dest_path = f'/kaggle/working/{filename}'\n            shutil.copy(src_path, dest_path)","metadata":{"execution":{"iopub.status.busy":"2023-04-09T23:02:10.165750Z","iopub.execute_input":"2023-04-09T23:02:10.166860Z","iopub.status.idle":"2023-04-09T23:03:48.124327Z","shell.execute_reply.started":"2023-04-09T23:02:10.166806Z","shell.execute_reply":"2023-04-09T23:03:48.123175Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Defining A Selection\nThe [SQLiteDataset class](https://github.com/graphnet-team/graphnet/blob/main/src/graphnet/data/sqlite/sqlite_dataset.py) is essentially a PyTorch Dataset (read more [here](https://pytorch.org/tutorials/beginner/data_loading_tutorial.html)) where the __get_item__ function extracts a single event at a time from the specified database. If a so-called *selection* is specified, only events in the selection are used for the dataset - that allows us to sub-sample the dataset for training.\n\nThe following few cells introduce a simple selection based on the number of pulses. This selection will then be used for training a GNN later.","metadata":{}},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\n\ndef make_selection(df: pd.DataFrame, pulse_threshold: int = 200) -> None:\n    \"\"\"Creates a validation and training selection (20 - 80). All events in both selections satisfies n_pulses <= 200 by default. \"\"\"\n    n_events = np.arange(0, len(df),1)\n    train_selection, validate_selection = train_test_split(n_events, \n                                                           shuffle=True, \n                                                           random_state = 42, \n                                                           test_size=0.20) \n    df['train'] = 0\n    df['validate'] = 0\n    \n    df['train'][train_selection] = 1\n    df['validate'][validate_selection] = 1\n    \n    assert len(train_selection) == sum(df['train'])\n    assert len(validate_selection) == sum(df['validate'])\n\n    #Remove events with large pulses from training and validation sample (memory)\n    df['train'][df['n_pulses']> pulse_threshold] = 0\n    df['validate'][df['n_pulses']> pulse_threshold] = 0\n    \n    for selection in ['train', 'validate']:\n        df.loc[df[selection] == 1, :].to_csv(f'{selection}_selection_max_{pulse_threshold}_pulses.csv')\n    return\n\ndef get_number_of_pulses(db: str, event_id: int, pulsemap: str) -> int:\n    with sqlite3.connect(db) as con:\n        query = f'select event_id from {pulsemap} where event_id = {event_id} limit 20000'\n        data = con.execute(query).fetchall()\n    return len(data)\n\ndef count_pulses(database: str, pulsemap: str) -> pd.DataFrame:\n    \"\"\" Will count the number of pulses in each event and return a single dataframe that contains counts for each event_id.\"\"\"\n    with sqlite3.connect(database) as con:\n        query = 'select event_id from meta_table'\n        events = pd.read_sql(query,con)\n    counts = {'event_id': [],\n              'n_pulses': []}\n    for event_id in tqdm(events['event_id']):\n        a = get_number_of_pulses(database, event_id, pulsemap)\n        counts['event_id'].append(event_id)\n        counts['n_pulses'].append(a)\n    df = pd.DataFrame(counts)\n    df.to_csv('counts.csv')\n    return df","metadata":{"execution":{"iopub.status.busy":"2023-04-09T23:12:59.087918Z","iopub.execute_input":"2023-04-09T23:12:59.088475Z","iopub.status.idle":"2023-04-09T23:12:59.867507Z","shell.execute_reply.started":"2023-04-09T23:12:59.088425Z","shell.execute_reply":"2023-04-09T23:12:59.865989Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training DynEdge\n\nNow that both database and selection is ready, everything is in place to begin training. DynEdge is a GNN implemented in GraphNeT - it represents IceCube events as 3D point clouds and leverages techniques from segmentation analysis in computer vision to reconstruct events. You can find technical details on the model in [this paper](https://iopscience.iop.org/article/10.1088/1748-0221/17/11/P11003). The model and training configuration shown below is nearly identical to what's presented in the paper. Note that this configuration was originally meant for low energy, so it's possible that some adjustments might improve performance.","metadata":{}},{"cell_type":"code","source":"from pytorch_lightning.callbacks import EarlyStopping\nfrom torch.optim.adamw import AdamW\nfrom graphnet.data.constants import FEATURES, TRUTH\nfrom graphnet.models import StandardModel\nfrom graphnet.models.detector.icecube import IceCubeKaggle\nfrom graphnet.models.gnn import DynEdge\nfrom graphnet.models.graph_builders import KNNGraphBuilder\nfrom graphnet.models.task.reconstruction import DirectionReconstructionWithKappa, ZenithReconstructionWithKappa, AzimuthReconstructionWithKappa\nfrom graphnet.training.callbacks import ProgressBar, PiecewiseLinearLR\nfrom graphnet.training.loss_functions import VonMisesFisher3DLoss, VonMisesFisher2DLoss\nfrom graphnet.training.labels import Direction\nfrom graphnet.training.utils import make_dataloader\nfrom graphnet.utilities.logging import get_logger\nfrom pytorch_lightning import Trainer\nimport pandas as pd\nprint ('Finished with imports.')","metadata":{"execution":{"iopub.status.busy":"2023-04-09T23:13:00.957135Z","iopub.execute_input":"2023-04-09T23:13:00.957595Z","iopub.status.idle":"2023-04-09T23:13:02.618425Z","shell.execute_reply.started":"2023-04-09T23:13:00.957557Z","shell.execute_reply":"2023-04-09T23:13:02.616781Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"logger = get_logger()\n\ndef build_model(config: Dict[str,Any], train_dataloader: Any) -> StandardModel:\n    \"\"\"Builds GNN from config\"\"\"\n    # Building model\n    detector = IceCubeKaggle(\n        graph_builder=KNNGraphBuilder(nb_nearest_neighbours=8),\n    )\n    gnn = DynEdge(\n        nb_inputs=detector.nb_outputs,\n        global_pooling_schemes=[\"min\", \"max\", \"mean\"],\n    )\n\n    if config[\"target\"] == 'direction':\n        task = DirectionReconstructionWithKappa(\n            hidden_size=gnn.nb_outputs,\n            target_labels=config[\"target\"],\n            loss_function=VonMisesFisher3DLoss(),\n        )\n        prediction_columns = [config[\"target\"] + \"_x\", \n                              config[\"target\"] + \"_y\", \n                              config[\"target\"] + \"_z\", \n                              config[\"target\"] + \"_kappa\" ]\n        additional_attributes = ['zenith', 'azimuth', 'event_id']\n\n    model = StandardModel(\n        detector=detector,\n        gnn=gnn,\n        tasks=[task],\n        optimizer_class=AdamW,\n        optimizer_kwargs={\"lr\": 1e-03, \"eps\": 1e-03},\n        scheduler_class=PiecewiseLinearLR,\n        scheduler_kwargs={\n            \"milestones\": [\n                0,\n                len(train_dataloader) / 2,\n                len(train_dataloader) * config[\"fit\"][\"max_epochs\"],\n            ],\n            \"factors\": [1e-02, 1, 1e-02],\n        },\n        scheduler_config={\n            \"interval\": \"step\",\n        },\n    )\n    model.prediction_columns = prediction_columns\n    model.additional_attributes = additional_attributes\n    \n    return model\n\ndef load_pretrained_model(config: Dict[str,Any], state_dict_path: str = '/kaggle/input/dynedge-pretrained/dynedge_pretrained_batch_1_to_50/state_dict.pth') -> StandardModel:\n    train_dataloader, _ = make_dataloaders(config = config)\n    model = build_model(config = config, \n                        train_dataloader = train_dataloader)\n    model.load_state_dict(state_dict_path)\n    model.prediction_columns = [config[\"target\"] + \"_x\", \n                              config[\"target\"] + \"_y\", \n                              config[\"target\"] + \"_z\", \n                              config[\"target\"] + \"_kappa\" ]\n    model.additional_attributes = ['zenith', 'azimuth', 'event_id']\n    return model\n\ndef make_dataloaders(config: Dict[str, Any]) -> List[Any]:\n    \"\"\"Constructs training and validation dataloaders for training with early stopping.\"\"\"\n    train_dataloader = make_dataloader(db = config['path'],\n                                            selection = pd.read_csv(config['train_selection'])[config['index_column']].ravel().tolist(),\n                                            pulsemaps = config['pulsemap'],\n                                            features = features,\n                                            truth = truth,\n                                            batch_size = config['batch_size'],\n                                            num_workers = config['num_workers'],\n                                            shuffle = True,\n                                            labels = {'direction': Direction()},\n                                            index_column = config['index_column'],\n                                            truth_table = config['truth_table'],\n                                            )\n    \n    validate_dataloader = make_dataloader(db = config['path'],\n                                            selection = pd.read_csv(config['validate_selection'])[config['index_column']].ravel().tolist(),\n                                            pulsemaps = config['pulsemap'],\n                                            features = features,\n                                            truth = truth,\n                                            batch_size = config['batch_size'],\n                                            num_workers = config['num_workers'],\n                                            shuffle = False,\n                                            labels = {'direction': Direction()},\n                                            index_column = config['index_column'],\n                                            truth_table = config['truth_table'],\n                                          \n                                            )\n    return train_dataloader, validate_dataloader\n\ndef train_dynedge_from_scratch(config: Dict[str, Any]) -> StandardModel:\n    \"\"\"Builds and trains GNN according to config.\"\"\"\n    logger.info(f\"features: {config['features']}\")\n    logger.info(f\"truth: {config['truth']}\")\n    \n    #archive = os.path.join(config['base_dir'], \"train_model_without_configs\")\n    #run_name = f\"dynedge_{config['target']}_{config['run_name_tag']}\"\n\n    train_dataloader, validate_dataloader = make_dataloaders(config = config)\n\n    model = build_model(config, train_dataloader)\n\n    # Training model\n    callbacks = [\n        EarlyStopping(\n            monitor=\"val_loss\",\n            patience=config[\"early_stopping_patience\"],\n        ),\n        ProgressBar(),\n    ]\n\n    model.fit(\n        train_dataloader,\n        validate_dataloader,\n        callbacks=callbacks,\n        **config[\"fit\"]\n    )\n    return model\n\ndef inference(model, config: Dict[str, Any]) -> pd.DataFrame:\n    \"\"\"Applies model to the database specified in config['inference_database_path'] and saves results to disk.\"\"\"\n    # Make Dataloader\n    test_dataloader = make_dataloader(db = config['inference_database_path'],\n                                            selection = None, # Entire database\n                                            pulsemaps = config['pulsemap'],\n                                            features = features,\n                                            truth = truth,\n                                            batch_size = config['batch_size'],\n                                            num_workers = config['num_workers'],\n                                            shuffle = False,\n                                            labels = {'direction': Direction()},\n                                            index_column = config['index_column'],\n                                            truth_table = config['truth_table'],\n                                            )\n    \n    # Get predictions\n    results = model.predict_as_dataframe(\n        gpus = [0],\n        dataloader = test_dataloader,\n        prediction_columns=model.prediction_columns,\n        additional_attributes=model.additional_attributes,\n    )\n    return results","metadata":{"execution":{"iopub.status.busy":"2023-04-09T23:13:02.621574Z","iopub.execute_input":"2023-04-09T23:13:02.622057Z","iopub.status.idle":"2023-04-09T23:13:02.651872Z","shell.execute_reply.started":"2023-04-09T23:13:02.622003Z","shell.execute_reply":"2023-04-09T23:13:02.650018Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Constants\nfeatures = FEATURES.KAGGLE\ntruth = TRUTH.KAGGLE\n\n# Configuration\nconfig = {\n        \"path\": '',\n        \"inference_database_path\": '/kaggle/working/batch_51.db',\n        \"pulsemap\": 'pulse_table',\n        \"truth_table\": 'meta_table',\n        \"features\": features,\n        \"truth\": truth,\n        \"index_column\": 'event_id',\n        \"run_name_tag\": 'my_example',\n        \"batch_size\": 200,\n        \"num_workers\": 2,\n        \"target\": 'direction',\n        \"early_stopping_patience\": 5,\n        \"fit\": {\n                \"max_epochs\": 50,\n                \"gpus\": [0],\n                \"distribution_strategy\": None,\n                },\n        'train_selection': f'/kaggle/working/train_selection_max_180_pulses.csv',\n        'validate_selection': f'/kaggle/working/validate_selection_max_180_pulses.csv',\n        'test_selection': None,\n        'base_dir': 'training'\n}","metadata":{"execution":{"iopub.status.busy":"2023-04-09T23:13:02.654295Z","iopub.execute_input":"2023-04-09T23:13:02.654694Z","iopub.status.idle":"2023-04-09T23:13:02.670569Z","shell.execute_reply.started":"2023-04-09T23:13:02.654618Z","shell.execute_reply":"2023-04-09T23:13:02.668998Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Inference & Evaluation\n\nWith a trained model loaded into memory, we can now apply the model to a batch. The following cells will start inference (or load in a csv with predictions, if you're in a hurry) and plot the results. ","metadata":{}},{"cell_type":"code","source":"def convert_to_3d(df: pd.DataFrame) -> pd.DataFrame:\n    \"\"\"Converts zenith and azimuth to 3D direction vectors\"\"\"\n    df['true_x'] = np.cos(df['azimuth']) * np.sin(df['zenith'])\n    df['true_y'] = np.sin(df['azimuth'])*np.sin(df['zenith'])\n    df['true_z'] = np.cos(df['zenith'])\n    return df\n\ndef calculate_angular_error(df : pd.DataFrame) -> pd.DataFrame:\n    \"\"\"Calculates the opening angle (angular error) between true and reconstructed direction vectors\"\"\"\n    df['angular_error'] = np.arccos(df['true_x']*df['direction_x'] + df['true_y']*df['direction_y'] + df['true_z']*df['direction_z'])\n    return df","metadata":{"execution":{"iopub.status.busy":"2023-04-09T23:13:03.193231Z","iopub.execute_input":"2023-04-09T23:13:03.194768Z","iopub.status.idle":"2023-04-09T23:13:03.203771Z","shell.execute_reply.started":"2023-04-09T23:13:03.194709Z","shell.execute_reply":"2023-04-09T23:13:03.202029Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Running all functions for 1 to 3 training batches","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport pandas as pd\n\npulsemap = 'pulse_table'\ndatabase = '/kaggle/working/batch'\nscores = []\nfor dirname, _, filenames in os.walk('/kaggle/working'):\n    for filename in filenames:\n        filepath = os.path.join(dirname, filename)\n        if database in filepath:\n            db_filepath = filepath\n            print ('\\n',db_filepath)\n            # Counting pulses, cleaning data to eliminate outliers and creating train/validation dataloaders\n            # 3 csv files are produced; training and validation selections, and a third file called *counts.csv*. \n            df = count_pulses(db_filepath, pulsemap)\n            make_selection(df = df, pulse_threshold = 180)\n            # Plotting pulses filtering\n            fig = plt.figure(figsize=(6,4), constrained_layout = True)\n            plt.hist(df['n_pulses'], histtype = 'step', label = 'batch_1', bins = np.arange(0,400,1))\n            plt.xlabel('# of Pulses', size = 15);\n            plt.xticks(size = 15);\n            plt.yticks(size = 15);\n            plt.plot(np.repeat(200,2), [0, 4000], label = f'Selection\\n{np.round((sum(df[\"n_pulses\"]<= 180)/len(df))*100, 1)} % pass' ) \n            plt.legend(frameon = False, fontsize = 15);\n            print(f'Event with highest number of pulses counted: {df[\"n_pulses\"].max()}')\n            # Train model\n            config['path'] = db_filepath\n            batch_number = int(filename.split('_')[1].split('.')[0])\n            model_path = f'/kaggle/working/model{batch_number}.pth'\n            model_state_path = f'/kaggle/working/state_dict{batch_number}.pth'\n            model = train_dynedge_from_scratch(config = config)\n            # Save model\n            model.save(model_path)\n            model.save_state_dict(model_state_path)\n            # Inference\n            results = inference(model, config)\n            results1 = convert_to_3d(results)\n            # Scoring\n            results2 = calculate_angular_error(results1)\n            display(results2)\n            logger.info(f\"Writing results to /kaggle/working\")\n            results2.to_csv(f\"results{batch_number}.csv\")\n            score = np.round(results2[\"angular_error\"].mean(),2)\n            scores.append (dict(batch_id = batch_number, scores = score))\n            print (f'Model model{batch_number} score:', score)\n            # Angular error distribution\n            cut_threshold = 0.5\n            fig = plt.figure(figsize = (6,6))\n            plt.hist(results2['angular_error'][1/np.sqrt(results2['direction_kappa']) <= cut_threshold], \n                     bins = np.arange(0,np.pi*2, 0.05), \n                     histtype = 'step', \n                     label = f'sigma <= {cut_threshold}: {np.round(results2[\"angular_error\"][1/np.sqrt(results2[\"direction_kappa\"]) <= cut_threshold].mean(),2)}')\n\n            plt.hist(results2['angular_error'][1/np.sqrt(results2['direction_kappa']) > cut_threshold], \n                     bins = np.arange(0,np.pi*2, 0.05), \n                     histtype = 'step', \n                     label = f'sigma > {cut_threshold}: {np.round(results2[\"angular_error\"][1/np.sqrt(results2[\"direction_kappa\"]) > cut_threshold].mean(),2)}')\n            plt.xlabel('Angular Error [rad.]', size = 15)\n            plt.ylabel('Counts', size = 15)\n            plt.title(f'Angular Error Distribution (Batch {batch_number})', size = 15)\n            plt.legend(frameon = False, fontsize = 15)# saving batch vs scores into a file   \nscores_df = pd.DataFrame (scores)\nscores_df.to_csv('scores_csv')","metadata":{"execution":{"iopub.status.busy":"2023-04-09T23:13:04.203150Z","iopub.execute_input":"2023-04-09T23:13:04.203926Z","iopub.status.idle":"2023-04-09T23:15:45.952330Z","shell.execute_reply.started":"2023-04-09T23:13:04.203869Z","shell.execute_reply":"2023-04-09T23:15:45.949981Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"In the cell below, you can choose between training dynedge from scratch on the batch database or loading in a pretrained model that has trained on batches 1 to 50.","metadata":{}},{"cell_type":"code","source":"# Train from scratch (slow) - remember to save it!\n#model = train_dynedge_from_scratch(config = config)\n#model.save(model_path)\n#model.save_state_dict(model_state_path)\n\n\n# Load state-dict from pre-trained model (faster)\n#model = load_pretrained_model(config = config)","metadata":{"execution":{"iopub.status.busy":"2023-04-09T17:05:40.653777Z","iopub.execute_input":"2023-04-09T17:05:40.654747Z","iopub.status.idle":"2023-04-09T17:40:25.438017Z","shell.execute_reply.started":"2023-04-09T17:05:40.654696Z","shell.execute_reply":"2023-04-09T17:40:25.437029Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Inference\n#results = inference(model, config)","metadata":{"execution":{"iopub.status.busy":"2023-04-09T17:45:07.418849Z","iopub.execute_input":"2023-04-09T17:45:07.419624Z","iopub.status.idle":"2023-04-09T17:48:08.423924Z","shell.execute_reply.started":"2023-04-09T17:45:07.419575Z","shell.execute_reply":"2023-04-09T17:48:08.422674Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#display(results)","metadata":{"execution":{"iopub.status.busy":"2023-04-09T17:54:36.867953Z","iopub.execute_input":"2023-04-09T17:54:36.869032Z","iopub.status.idle":"2023-04-09T17:54:36.894881Z","shell.execute_reply.started":"2023-04-09T17:54:36.868959Z","shell.execute_reply":"2023-04-09T17:54:36.893683Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#fig = plt.figure(figsize = (6,6))\n#plt.hist(results['angular_error'], \n#         bins = np.arange(0,np.pi*2, 0.05), \n#         histtype = 'step', \n#         label = f'mean angular error: {np.round(results[\"angular_error\"].mean(),2)}')\n#plt.xlabel('Angular Error [rad.]', size = 15)\n#plt.ylabel('Counts', size = 15)\n#plt.title('Angular Error Distribution (Batch 51)', size = 15)\n#plt.legend(frameon = False, fontsize = 15)","metadata":{"execution":{"iopub.status.busy":"2023-04-09T17:53:36.821008Z","iopub.execute_input":"2023-04-09T17:53:36.822015Z","iopub.status.idle":"2023-04-09T17:53:37.075945Z","shell.execute_reply.started":"2023-04-09T17:53:36.821961Z","shell.execute_reply":"2023-04-09T17:53:37.074913Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"So the pre-trained dynedge seems to perform quite well. Another interesting feature of the reconstruction is that dynedge (when coupled with the[ DirectionReconstructionWithKappa](https://github.com/graphnet-team/graphnet/blob/7e857562898ebebebc9a105159fd3d4eb4994aea/src/graphnet/models/task/reconstruction.py#L45) is that dynedge estimated *kappa* the concentration parameter from the vonMisesFisher distribution. Kappa is analogus to sigma via sigma = 1/sqrt(kappa), and the quality of the direction estimate should be highly correlated with this parameter. ","metadata":{}},{"cell_type":"markdown","source":"As you can see, the variable can be used to distinguish \"good\" and \"bad\" reconstructions with some confidence. ","metadata":{}},{"cell_type":"markdown","source":"## A few hints for your neutrino data science journey!\n\n* The configuration of dynedge shown in this notebook is the so-called \"baseline\". It's not optimized for high energy neutrinos, so you might be able to squeeze out a bit more performance by tuning hyperparameters or making larger modifications; such as switching out the learning rate scheduler or choosing a different loss function, etc.\n\n* You can use the kappa variable to group events into different categories. Perhaps training a seperate reconstruction method for each performs better?\n\n* You may want to adjust the [ParquetDataset](https://github.com/graphnet-team/graphnet/blob/7e857562898ebebebc9a105159fd3d4eb4994aea/src/graphnet/data/parquet/parquet_dataset.py#L11) such that it works with the competition data. This would allow you to train / infer directly on the competition files (No conversion to sqlite needed). Feel free to contribute this to the repository!\n\n\nGood luck!","metadata":{}}]}