{"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":"\n# Graph Neural Networks on Competition Data\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\nGraph Neural Networks (GNNs) are a class of neural networks that work on graph representation of data. In recent years, development and application of GNNs have been of high interest in the machine learning community, and includes areas of research such as protein folding, computer vision and knowledge graphs. Part of the appeal of GNNs is the data representation itself - the graphs - as they provide an abstract and flexible format to represent a wide range of issues as GNN problems. Specifically for IceCube, GNNs are interesting because graphs allow us to naturally represent the irregular geometry. \n\nIn a recent [publication from IceCube](https://iopscience.iop.org/article/10.1088/1748-0221/17/11/P11003), we show how a GNN compares against traditional methods on a series of reconstruction and classification tasks that are of common interest in neutrino physics in the low energy range of IceCube. While the paper explores applications in the low energy range, the GNN is applicable to higher energies too.\n\nThe GNN is called *dynedge* and is implemented in [GraphNeT](https://github.com/graphnet-team/graphnet). GraphNeT is an open-source python framework aimed at providing high quality, user friendly, end-to-end functionality to perform reconstruction tasks at neutrino telescopes using graph neural networks (GNNs). This [paper](https://arxiv.org/abs/2210.12194) (in review) supplements the details found in the github.\n\nThis notebook contains everything you need 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.\n\n\nWe hope that this is enough to get you started with GNNs on the competition data. Once you've gotten the code running, you'll see that it doesn't require much to introduce a completely different GNN or to make changes to dynedge. Please feel free to tinker with all of the code - and if you want to - you're very much welcome to add your contributions via pull request on our github. This way, your work might benefit researchers across the globe! **But remember to read our [contribution guide](https://github.com/graphnet-team/graphnet/blob/main/CONTRIBUTING.md) first**!  \n\nFinally, please note that the first stable release of GraphNeT is aimed for first half of 2023 - so you might still encounter bugs - and if you do, [please create an issue.](https://github.com/graphnet-team/graphnet/issues)\n\n\n","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) which contains guides for both CPU and GPU installation. However, I had to take a few extra steps to get the library installed in a Kaggle notebook. Run the cell below - 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')","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-02-03T08:28:40.556433Z","iopub.execute_input":"2023-02-03T08:28:40.556880Z","iopub.status.idle":"2023-02-03T08:33:24.879030Z","shell.execute_reply.started":"2023-02-03T08:28:40.556778Z","shell.execute_reply":"2023-02-03T08:33:24.877780Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import graphnet","metadata":{"execution":{"iopub.status.busy":"2023-02-03T08:33:24.882607Z","iopub.execute_input":"2023-02-03T08:33:24.883272Z","iopub.status.idle":"2023-02-03T08:33:24.965010Z","shell.execute_reply.started":"2023-02-03T08:33:24.883233Z","shell.execute_reply":"2023-02-03T08:33:24.963646Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Converting The Parquet Files to SQLite\n\nWhile GraphNeT have some support for parquet, the majority of functionality is tied to the SQLite data format. Therefore, to make GraphNeT compatible with the data provided in this competition, I've included a small converter 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 you 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\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    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-02-03T08:33:24.967092Z","iopub.execute_input":"2023-02-03T08:33:24.967514Z","iopub.status.idle":"2023-02-03T08:33:29.466709Z","shell.execute_reply.started":"2023-02-03T08:33:24.967470Z","shell.execute_reply":"2023-02-03T08:33:29.465454Z"},"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":"!cp /kaggle/input/batch-1/batch_1.db .\n!cp /kaggle/input/batch-51/batch_51.db .","metadata":{"execution":{"iopub.status.busy":"2023-02-03T08:33:29.469852Z","iopub.execute_input":"2023-02-03T08:33:29.470673Z","iopub.status.idle":"2023-02-03T08:34:33.177122Z","shell.execute_reply.started":"2023-02-03T08:33:29.470630Z","shell.execute_reply":"2023-02-03T08:34:33.175468Z"},"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\nIn the following few cells, I'll introduce a simple selection based on the number of pulses. We'll then use this selection for training a GNN later.","metadata":{}},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\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-02-03T08:34:33.180456Z","iopub.execute_input":"2023-02-03T08:34:33.181030Z","iopub.status.idle":"2023-02-03T08:34:33.774680Z","shell.execute_reply.started":"2023-02-03T08:34:33.180983Z","shell.execute_reply":"2023-02-03T08:34:33.773312Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The cell below will produce three csv files; training and validation selections, and a third files called *counts.csv*. ","metadata":{}},{"cell_type":"code","source":"pulsemap = 'pulse_table'\ndatabase = '/kaggle/working/batch_1.db'\n\ndf = count_pulses(database, pulsemap)\nmake_selection(df = df, pulse_threshold =  200)","metadata":{"execution":{"iopub.status.busy":"2023-02-03T08:34:33.776400Z","iopub.execute_input":"2023-02-03T08:34:33.776821Z","iopub.status.idle":"2023-02-03T08:35:49.835165Z","shell.execute_reply.started":"2023-02-03T08:34:33.776775Z","shell.execute_reply":"2023-02-03T08:35:49.834123Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport pandas as pd\nfig = plt.figure(figsize=(6,4), constrained_layout = True)\nplt.hist(df['n_pulses'], histtype = 'step', label = 'batch_1', bins = np.arange(0,400,1))\nplt.xlabel('# of Pulses', size = 15);\nplt.xticks(size = 15);\nplt.yticks(size = 15);\nplt.plot(np.repeat(200,2), [0, 4000], label = f'Selection\\n{np.round((sum(df[\"n_pulses\"]<= 200)/len(df))*100, 1)} % pass' ) \nplt.legend(frameon = False, fontsize = 15);","metadata":{"execution":{"iopub.status.busy":"2023-02-03T08:35:49.836752Z","iopub.execute_input":"2023-02-03T08:35:49.837371Z","iopub.status.idle":"2023-02-03T08:35:50.256450Z","shell.execute_reply.started":"2023-02-03T08:35:49.837324Z","shell.execute_reply":"2023-02-03T08:35:50.255463Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f'Event with highest number of pulses counted: {df[\"n_pulses\"].max()}')","metadata":{"execution":{"iopub.status.busy":"2023-02-03T08:35:50.257832Z","iopub.execute_input":"2023-02-03T08:35:50.258850Z","iopub.status.idle":"2023-02-03T08:35:50.266848Z","shell.execute_reply.started":"2023-02-03T08:35:50.258772Z","shell.execute_reply":"2023-02-03T08:35:50.265678Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"As visible from the plot, most events have less or equal to 200 pulses. The print message shows us that the outliers can have tens of thousands of pulses (20000 is the max we let it count in the code above). In principle, one could train the GNN without descrimination on this collection of events, but events with large pulse counts will make the memory usuage volatile and force a low batch size. Therefore (and feel free to challenge this) this notebook will train on events with 200 or less pulses.  ","metadata":{}},{"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.adam import Adam\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\n\nlogger = 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=Adam,\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._inference_trainer = Trainer(config['fit'])\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    # Save predictions and model to file\n    archive = os.path.join(config['base_dir'], \"train_model_without_configs\")\n    run_name = f\"dynedge_{config['target']}_{config['run_name_tag']}\"\n    db_name = config['path'].split(\"/\")[-1].split(\".\")[0]\n    path = os.path.join(archive, db_name, run_name)\n    logger.info(f\"Writing results to {path}\")\n    os.makedirs(path, exist_ok=True)\n\n    results.to_csv(f\"{path}/results.csv\")\n    return results","metadata":{"execution":{"iopub.status.busy":"2023-02-03T08:35:50.268509Z","iopub.execute_input":"2023-02-03T08:35:50.269123Z","iopub.status.idle":"2023-02-03T08:35:51.121812Z","shell.execute_reply.started":"2023-02-03T08:35:50.269085Z","shell.execute_reply":"2023-02-03T08:35:51.120342Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Constants\nfeatures = FEATURES.KAGGLE\ntruth = TRUTH.KAGGLE\n\n# Configuration\nconfig = {\n        \"path\": '/kaggle/working/batch_1.db',\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': '/kaggle/working/train_selection_max_200_pulses.csv',\n        'validate_selection': '/kaggle/working/validate_selection_max_200_pulses.csv',\n        'test_selection': None,\n        'base_dir': 'training'\n}","metadata":{"execution":{"iopub.status.busy":"2023-02-03T08:35:51.127064Z","iopub.execute_input":"2023-02-03T08:35:51.127385Z","iopub.status.idle":"2023-02-03T08:35:51.137833Z","shell.execute_reply.started":"2023-02-03T08:35:51.127354Z","shell.execute_reply":"2023-02-03T08:35:51.136719Z"},"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_1 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\n# Load state-dict from pre-trained model (faster)\nmodel = load_pretrained_model(config = config)","metadata":{"execution":{"iopub.status.busy":"2023-02-03T08:35:51.139859Z","iopub.execute_input":"2023-02-03T08:35:51.140351Z","iopub.status.idle":"2023-02-03T08:35:51.806054Z","shell.execute_reply.started":"2023-02-03T08:35:51.140312Z","shell.execute_reply":"2023-02-03T08:35:51.805023Z"},"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 batch_51. 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":"# Inference\nresults = inference(model, config)","metadata":{"execution":{"iopub.status.busy":"2023-02-03T08:35:51.807346Z","iopub.execute_input":"2023-02-03T08:35:51.807746Z","iopub.status.idle":"2023-02-03T08:46:55.296296Z","shell.execute_reply.started":"2023-02-03T08:35:51.807704Z","shell.execute_reply":"2023-02-03T08:46:55.295092Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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    \"\"\"Calcualtes 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-02-03T08:46:55.298831Z","iopub.execute_input":"2023-02-03T08:46:55.299249Z","iopub.status.idle":"2023-02-03T08:46:55.307103Z","shell.execute_reply.started":"2023-02-03T08:46:55.299205Z","shell.execute_reply":"2023-02-03T08:46:55.306088Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"results = convert_to_3d(results)\nresults = calculate_angular_error(results)","metadata":{"execution":{"iopub.status.busy":"2023-02-03T08:46:55.308430Z","iopub.execute_input":"2023-02-03T08:46:55.309457Z","iopub.status.idle":"2023-02-03T08:46:55.360992Z","shell.execute_reply.started":"2023-02-03T08:46:55.309420Z","shell.execute_reply":"2023-02-03T08:46:55.359991Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig = plt.figure(figsize = (6,6))\nplt.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)}')\nplt.xlabel('Angular Error [rad.]', size = 15)\nplt.ylabel('Counts', size = 15)\nplt.title('Angular Error Distribution (Batch 51)', size = 15)\nplt.legend(frameon = False, fontsize = 15)","metadata":{"execution":{"iopub.status.busy":"2023-02-03T08:56:19.735174Z","iopub.execute_input":"2023-02-03T08:56:19.735528Z","iopub.status.idle":"2023-02-03T08:56:19.981355Z","shell.execute_reply.started":"2023-02-03T08:56:19.735497Z","shell.execute_reply":"2023-02-03T08:56:19.980402Z"},"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":"code","source":"cut_threshold = 0.5\nfig = plt.figure(figsize = (6,6))\nplt.hist(results['angular_error'][1/np.sqrt(results['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(results[\"angular_error\"][1/np.sqrt(results[\"direction_kappa\"]) <= cut_threshold].mean(),2)}')\n\nplt.hist(results['angular_error'][1/np.sqrt(results['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(results[\"angular_error\"][1/np.sqrt(results[\"direction_kappa\"]) > cut_threshold].mean(),2)}')\nplt.xlabel('Angular Error [rad.]', size = 15)\nplt.ylabel('Counts', size = 15)\nplt.title('Angular Error Distribution (Batch 51)', size = 15)\nplt.legend(frameon = False, fontsize = 15)","metadata":{"execution":{"iopub.status.busy":"2023-02-03T08:56:24.975565Z","iopub.execute_input":"2023-02-03T08:56:24.976250Z","iopub.status.idle":"2023-02-03T08:56:25.236386Z","shell.execute_reply.started":"2023-02-03T08:56:24.976211Z","shell.execute_reply":"2023-02-03T08:56:25.235418Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{}}]}