{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# # Move software to working disk\n# import os\n\n# os.system('rm  -r software')\n# os.system('scp -r /kaggle/input/graphnet-and-dependencies/software ./')\n\n# os.system('pip install /kaggle/working/software/dependencies/torch-1.11.0+cu115-cp37-cp37m-linux_x86_64.whl')\n# os.system('pip install /kaggle/working/software/dependencies/torch_cluster-1.6.0-cp37-cp37m-linux_x86_64.whl')\n# os.system('pip install /kaggle/working/software/dependencies/torch_scatter-2.0.9-cp37-cp37m-linux_x86_64.whl')\n# os.system('pip install /kaggle/working/software/dependencies/torch_sparse-0.6.13-cp37-cp37m-linux_x86_64.whl')\n# os.system('pip install /kaggle/working/software/dependencies/torch_geometric-2.0.4.tar.gz')\n\n# os.system('cd software/graphnet;pip install --no-index --find-links=\"/kaggle/working/software/dependencies\" -e .[torch]')\n\n# os.system('pip install torchvision==0.12.0')\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# # !pip install torchvision==0.12.0\n\n# # Append to PATH\n# import sys\n# sys.path.append('/kaggle/working/software/graphnet/src')\n\n# import graphnet","metadata":{"_uuid":"68e2918f-b333-492b-ba03-db719a74b114","_cell_guid":"3286684a-52d1-47fd-8d36-4ccb0ac0e51d","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-03-17T17:40:14.801849Z","iopub.execute_input":"2023-03-17T17:40:14.803255Z","iopub.status.idle":"2023-03-17T17:40:14.838925Z","shell.execute_reply.started":"2023-03-17T17:40:14.803198Z","shell.execute_reply":"2023-03-17T17:40:14.837923Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nos.system('ls /kaggle/input/batch-52-56')\nos.system('cp /kaggle/input/batch-52-56/batch_52_56.db ./')\n\n# !ls /kaggle/input/batch-52\n# !cp /kaggle/input/batch-52/batch_52.db .","metadata":{"_uuid":"31e31a57-1c3c-40a8-a2f5-426e178c3ed9","_cell_guid":"e6ecb974-3e8d-46cc-9e0f-936a3fef3c6b","collapsed":false,"execution":{"iopub.status.busy":"2023-03-16T19:43:59.414588Z","iopub.execute_input":"2023-03-16T19:43:59.417141Z","iopub.status.idle":"2023-03-16T19:44:18.15674Z","shell.execute_reply.started":"2023-03-16T19:43:59.417084Z","shell.execute_reply":"2023-03-16T19:44:18.155305Z"},"jupyter":{"outputs_hidden":false},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%writefile make_selection.py\n\n\nimport matplotlib.pyplot as plt\nimport pandas as pd\nimport numpy as np\nimport sqlite3\nfrom tqdm import tqdm\nfrom 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\n\n\npulsemap = 'pulse_table'\ndatabase = '/kaggle/working/batch_52_56.db'\n\ndf = count_pulses(database, pulsemap)\nmake_selection(df = df, pulse_threshold =  150)","metadata":{"_uuid":"844bf9fa-8113-4285-a02b-dd100bf9cb38","_cell_guid":"70606716-1884-4c2a-b826-a79b2bdf1616","collapsed":false,"execution":{"iopub.status.busy":"2023-03-16T19:44:18.159015Z","iopub.execute_input":"2023-03-16T19:44:18.159478Z","iopub.status.idle":"2023-03-16T19:44:19.170561Z","shell.execute_reply.started":"2023-03-16T19:44:18.159413Z","shell.execute_reply":"2023-03-16T19:44:19.168582Z"},"jupyter":{"outputs_hidden":false},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!python /kaggle/working/make_selection.py","metadata":{"_uuid":"472a74ab-4a50-48e9-8128-c8c3ce1443fd","_cell_guid":"50535ed3-3b30-464a-8971-61c0c37b986e","collapsed":false,"execution":{"iopub.status.busy":"2023-03-16T19:44:19.174038Z","iopub.execute_input":"2023-03-16T19:44:19.174443Z"},"jupyter":{"outputs_hidden":false},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%writefile train.py\n\n\nimport os\nfrom typing import Any, Dict, List, Optional\nimport pytorch_lightning\nfrom 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 Logger\nfrom pytorch_lightning import Trainer\nfrom pytorch_lightning.loggers import WandbLogger\nfrom pytorch_lightning.utilities import rank_zero_only\nimport pandas as pd\n\n\n# Constants\nfeatures = FEATURES.KAGGLE\ntruth = TRUTH.KAGGLE\n\n# Configuration\nconfig = {\n        \"use_wandb\": True,\n        \"state_dict_path\": '/kaggle/input/dynedge-pretrained/dynedge_pretrained_batch_1_to_50/state_dict.pth',\n        \"path\": '/kaggle/working/batch_52_56.db',\n        \"inference_database_path\": '/kaggle/working/batch_52_56.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\": 10,\n                \"gpus\": [0, 1],\n                \"distribution_strategy\": 'ddp',\n                },\n        'train_selection': '/kaggle/working/train_selection_max_150_pulses.csv',\n        'validate_selection': '/kaggle/working/validate_selection_max_150_pulses.csv',\n        'test_selection': None,\n        'base_dir': 'training'\n}\n\nprint(config)\n    \n# Make sure W&B output directory exists\nWANDB_DIR = \"/kaggle/working/wandb\"\nos.makedirs(WANDB_DIR, exist_ok=True)\n\nprint(WANDB_DIR)\n\nlogger = Logger()\n\n\ndef build_model(config: Dict[str,Any], train_dataloader: Any) -> StandardModel:\n    \"\"\"Builds GNN from config\"\"\"\n    # from torch.optim.adam import Adam\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    \n    model.prediction_columns = prediction_columns\n    model.additional_attributes = additional_attributes\n    \n    return model\n\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\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\n\ndef train_dynedge(config: Dict[str, Any]) -> StandardModel:\n    \"\"\"Builds and trains GNN according to config.\"\"\"\n    \n    if config[\"use_wandb\"]:\n        wandb_logger = WandbLogger(\n            project=\"icecube-train\",\n            entity=\"tz_li\",\n            save_dir=WANDB_DIR,\n            log_model=True,\n        )\n    \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    if config[\"use_wandb\"]:\n        if rank_zero_only.rank == 0:\n            wandb_logger.experiment.config.update(config)\n    \n\n    train_dataloader, validate_dataloader = make_dataloaders(config = config)\n\n    model = build_model(config, train_dataloader)\n    model.load_state_dict(config[\"state_dict_path\"])\n    \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        logger=wandb_logger,\n        **config[\"fit\"],\n    )\n    \n    wandb.finish()\n    return model\n\n\nprint('start training!')\n\nfrom kaggle_secrets import UserSecretsClient\nuser_secrets = UserSecretsClient()\nkey = user_secrets.get_secret(\"wandb\")\n\n    \nimport wandb\nwandb.login(key=key)\n\n\nmodel = train_dynedge(config)\nmodel.save_state_dict(\"/kaggle/working/finetune_batch52_56_epch20_state_dict.pth\")\nmodel.save(\"/kaggle/working/finetune_batch52_56_epch20_model.pth\")","metadata":{"_uuid":"a128aa45-6917-4421-82c9-efaa2974b965","_cell_guid":"55c47900-2711-40d9-8b32-b93f4165d49b","collapsed":false,"jupyter":{"outputs_hidden":false},"trusted":true},"execution_count":null,"outputs":[]}]}