{"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":"# Info\n\n\nHere I develop a transformer architecture for this comp.\n\nInput: sequence of xyz, q, t, auxiliary.","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nfrom torch.nn import functional as F\nfrom tqdm import tqdm\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport random\nimport math\nimport wandb\nimport gc","metadata":{"execution":{"iopub.status.busy":"2023-09-15T10:08:07.999326Z","iopub.execute_input":"2023-09-15T10:08:07.999699Z","iopub.status.idle":"2023-09-15T10:08:08.006538Z","shell.execute_reply.started":"2023-09-15T10:08:07.999668Z","shell.execute_reply":"2023-09-15T10:08:08.005408Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# set seed for reproducibility\ntorch.manual_seed(1337);","metadata":{"execution":{"iopub.status.busy":"2023-09-15T10:08:08.008558Z","iopub.execute_input":"2023-09-15T10:08:08.009501Z","iopub.status.idle":"2023-09-15T10:08:08.019634Z","shell.execute_reply.started":"2023-09-15T10:08:08.009461Z","shell.execute_reply":"2023-09-15T10:08:08.018619Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"RUN_ON_KAGGLE = True\nLOG_WANDB = False\n\nif RUN_ON_KAGGLE:\n    root_dir = \"/kaggle/input\"\nelse:\n    root_dir = \"./\"","metadata":{"execution":{"iopub.status.busy":"2023-09-15T10:08:08.021817Z","iopub.execute_input":"2023-09-15T10:08:08.023056Z","iopub.status.idle":"2023-09-15T10:08:08.031974Z","shell.execute_reply.started":"2023-09-15T10:08:08.023028Z","shell.execute_reply":"2023-09-15T10:08:08.030948Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Configuration class\n\nclass CFG:\n    # training hyperparameters\n    max_iters = int(1e4) # number of training iterations\n    n_opt_distance = 1e3 # before this you're optimizing for the distance, after for the comp metric\n    change_batch_every = 256 # change batch of data every change_batch_every steps\n    device = 'cuda' if torch.cuda.is_available() else 'cpu'\n    evaluate_every_step = 200 \n    learning_rate = 1e-4\n    eval_iters = 16 # number of batches to process for evaluation\n\n    ## network hyperparameters\n    # --------------------\n    batch_size= 4\n    block_size = 192 # maximum context length\n    n_embd = 384\n    n_layers = 12\n    num_heads = 8\n    head_size = block_size // num_heads\n    drop_path = 0.2\n    # --------------------\n\n    # preprocecssing hyperparameters\n    max_xyz = 500 # maximum value for xyz: xyz will be preprocessed as x -> x/max_xyz\n    clip_charge = 5 # maximum charge: charge will be preprocessed as charge -> np.clip(charge, 0, clip_charge)\n    charge_rescale = 5 # after charge clipping, you'll rescale the charge by 5: charge -> charge / charge_rescale\n    max_t = 16000 # maximum cumulative time\n\n\n","metadata":{"execution":{"iopub.status.busy":"2023-09-15T10:08:08.033468Z","iopub.execute_input":"2023-09-15T10:08:08.034172Z","iopub.status.idle":"2023-09-15T10:08:08.046063Z","shell.execute_reply.started":"2023-09-15T10:08:08.034134Z","shell.execute_reply":"2023-09-15T10:08:08.045146Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load sensor geometry","metadata":{}},{"cell_type":"code","source":"df_sensor_geometry = pd.read_csv(f\"{root_dir}/icecube-neutrinos-in-deep-ice/sensor_geometry.csv\")\nsensor_ids = sorted(list(set(df_sensor_geometry[\"sensor_id\"].to_list())))\nsensor_ids[:10], len(sensor_ids)","metadata":{"execution":{"iopub.status.busy":"2023-09-15T10:08:08.050115Z","iopub.execute_input":"2023-09-15T10:08:08.050467Z","iopub.status.idle":"2023-09-15T10:08:08.070474Z","shell.execute_reply.started":"2023-09-15T10:08:08.050440Z","shell.execute_reply":"2023-09-15T10:08:08.069436Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Load metadata\n\nThis code defines a class called `EventDataLoader` that is responsible for loading and preprocessing data for a neutrino particle detection problem. \n\nUpon initialization, the class loads the training and test metadata files in Parquet format and splits the training data into a training and validation set. It also sets some internal variables for tracking the current batch and the last batch switch time. Additionally, it loads the events from Parquet files into memory and preprocesses them using several helper functions.\n\nThe main functions in the class are `get_single_event` and `get_xy`. The former takes as input the indices of the first and last pulse for a given event and a split type (train, test, or eval), preprocesses the event data, and returns the processed data as tensors. The `get_xy` function returns a batch of batch_size preprocessed events and their associated target values for either the training or test split.\n\nThe class also has a method `change_event_batch` that shuffles the current batch of events and associated metadata and reloads the new batch of events from disk. This function is called every change_batch_every iterations of the get_xy function to allow for iterating through multiple batches of data.","metadata":{}},{"cell_type":"code","source":"class EventDataLoader:\n    def __init__(self, CFG):\n        self.change_batch_every = CFG.change_batch_every\n        # train-test split in meta\n        self.batch_ids = {\n            \"train\": np.arange(1, 660),\n            \"test\": np.array([660]),\n        }\n        self.batch_id_current = {\n            'train': 0,\n            'test': 0\n        }\n        # loading metadata\n        print(\"Loading metadata ...\")\n        self.meta_data_current = {\n            \"train\": self.shuffle_metadata_batch(\"train\"),\n            \"test\": self.shuffle_metadata_batch(\"test\"),\n        }\n        # events\n        print(\"Loading events ...\")\n        self.events = {\n            \"train\": self.load_batch_events(\"train\"),\n            \"test\": self.load_batch_events(\"test\"),\n        }\n        self.counter=0 # used to keep track of the current batch\n        \n        self.CFG = CFG\n    \n    # Helper functions\n    def shuffle_metadata_batch(self, split):\n        \n        assert split in [\"train\", \"test\"]\n        \n        if split in [\"train\", \"test\"]:\n            # take random element from self.batch_ids[split]\n            random_batch_id = np.random.choice(self.batch_ids[split])\n            self.batch_id_current[split] = random_batch_id\n        else:\n            random_batch_id = 660\n            self.batch_id_current[split] = 660\n        return pd.read_parquet(f\"{root_dir}/train-meta-parquet/train_meta_{random_batch_id}.parquet\")\n    \n    \n    def load_batch_events(self,split):\n        \"\"\" load a batch of events from the parquet file \"\"\"\n        \n        \n        special_dir = \"train\" if split in [\"train\", \"test\"] else \"test\"\n        batch_id = self.batch_id_current[split]\n        return pd.read_parquet(f\"{root_dir}/icecube-neutrinos-in-deep-ice/{special_dir}/batch_{batch_id}.parquet\").reset_index()\n        \n    \n    # --- Main functions ---\n    \n\n\n    def change_event_batch(self):\n        \"\"\" change both the train and test event batches \"\"\"\n        \n        self.meta_data_current = {\n            \"train\":self.shuffle_metadata_batch(\"train\"),\n            \"test\":self.shuffle_metadata_batch(\"test\"),\n        }\n        self.events[\"train\"] = self.load_batch_events(\"train\")\n        self.events[\"test\"] = self.load_batch_events(\"test\")\n        self.counter = 0\n    \n    \n    def get_single_event(self, first_pulse_index, last_pulse_index, split):\n        \"\"\" get a single event from the dataframe (train or test) \"\"\"\n\n        assert split in [\"train\", \"test\"], \"split must be either 'train', 'test'\"\n\n        # get the event\n        event = self.events[split].iloc[first_pulse_index:last_pulse_index+1]\n        \n        \n        \n        # merge event with df_sensor_geometry using sensor_id to get x, y, z\n        event = pd.merge(event, df_sensor_geometry, on=\"sensor_id\")\n\n        # preprocess xyz\n        event[\"x\"] = (event[\"x\"] - np.average(event[\"x\"])) / self.CFG.max_xyz\n        event[\"y\"] = (event[\"y\"] - np.average(event[\"y\"])) / self.CFG.max_xyz\n        event[\"z\"] = (event[\"z\"] - np.average(event[\"z\"])) / self.CFG.max_xyz\n\n        # pad with zeros if the event is smaller than the block size\n        event_size = len(event)\n        if event_size < self.CFG.block_size:\n            event = event.append(pd.DataFrame(np.zeros((self.CFG.block_size - event_size, len(event.columns))), columns=event.columns))\n        event = event[:self.CFG.block_size]\n\n        xyz = event[['x', 'y', 'z']].values\n        xyz[event_size:] = 0\n        \n        # preprocess time\n        time = event[\"time\"].values / self.CFG.max_t\n        time[event_size:] = 0\n        \n\n        # preprocess charge\n        charge = event[\"charge\"].values\n        charge = np.clip(charge, 0, self.CFG.clip_charge)\n        charge = charge / self.CFG.charge_rescale\n        charge[event_size:] = 0\n        \n        # preprocess auxiliary\n        auxiliary = np.array(event[\"auxiliary\"].values, dtype=np.float32)\n        auxiliary -= 0.5\n        auxiliary[event_size:] = 0\n\n        # convert to tensors\n        xyz = torch.tensor(xyz, dtype=torch.float32)\n        time = torch.tensor(time, dtype=torch.float32)\n        charge = torch.tensor(charge, dtype=torch.float32)\n        auxiliary = torch.tensor(auxiliary, dtype=torch.float32)\n        event_size = torch.tensor(event_size, dtype=torch.float32)\n        \n        out = {\n            \"xyz\": xyz, \n            \"time\": time, \n            \"charge\": charge,\n            'auxiliary':auxiliary,\n            \"event_size\": event_size\n        }\n\n        return out\n\n    def get_xy(self, split):\n        assert split in [\"train\", \"test\"], \"split must be either 'train' or 'test'\"\n\n        df_meta_sample = self.meta_data_current[split].sample(n=self.CFG.batch_size).reset_index(drop=True)\n        \n\n        # in dataframe df_meta_sample, loop over first_pulse_index and last_pulse_index columns, and slice the corresponding rows from df_events_batch\n        first_pulse_indices = df_meta_sample[\"first_pulse_index\"].to_list()\n        last_pulse_indices = df_meta_sample[\"last_pulse_index\"].to_list()\n\n        # preprocessed data as tensors\n        preprocessed_data = [self.get_single_event(first_pulse_index, last_pulse_index, split) for first_pulse_index, last_pulse_index in zip(first_pulse_indices, last_pulse_indices)]\n\n        # create batch of preprocessed data\n        xyz_batch = torch.stack([data[\"xyz\"] for data in preprocessed_data])\n        time_batch = torch.stack([data[\"time\"] for data in preprocessed_data])\n        charge_batch = torch.stack([data[\"charge\"] for data in preprocessed_data])\n        auxiliary_batch = torch.stack([data[\"auxiliary\"] for data in preprocessed_data])\n        event_size_batch = torch.stack([data[\"event_size\"] for data in preprocessed_data])\n\n        # Targets\n        theta_target = np.squeeze(np.array(df_meta_sample[[\"zenith\"]].values, dtype=np.float32))\n        phi_target = np.squeeze(np.array(df_meta_sample[[\"azimuth\"]].values, dtype=np.float32))\n        y = np.array([\n            np.sin(theta_target)*np.cos(phi_target), \n            np.sin(theta_target)*np.sin(phi_target), \n            np.cos(theta_target)], dtype=np.float32)\n        y = torch.tensor(y.T)\n\n        # send them to the device\n        device = self.CFG.device\n        xyz_batch = xyz_batch.to(device)\n        time_batch = time_batch.to(device)\n        charge_batch = charge_batch.to(device)\n        auxiliary_batch = auxiliary_batch.to(device)\n        event_size_batch = event_size_batch.to(device)\n        y = y.to(device)\n\n        self.counter += 1\n        if self.counter == self.change_batch_every:\n            self.change_event_batch()\n            self.counter = 0\n        \n        X = {\n            \"xyz\": xyz_batch,\n            \"time\": time_batch,\n            \"charge\": charge_batch,\n            \"auxiliary\": auxiliary_batch,\n            \"event_size\": event_size_batch\n        }\n        \n        return X, y\n","metadata":{"execution":{"iopub.status.busy":"2023-09-15T10:08:08.093664Z","iopub.execute_input":"2023-09-15T10:08:08.093944Z","iopub.status.idle":"2023-09-15T10:08:08.133711Z","shell.execute_reply.started":"2023-09-15T10:08:08.093919Z","shell.execute_reply":"2023-09-15T10:08:08.132685Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_loader = EventDataLoader(CFG)","metadata":{"execution":{"iopub.status.busy":"2023-09-15T10:08:08.135703Z","iopub.execute_input":"2023-09-15T10:08:08.136175Z","iopub.status.idle":"2023-09-15T10:08:14.446924Z","shell.execute_reply.started":"2023-09-15T10:08:08.136138Z","shell.execute_reply":"2023-09-15T10:08:14.445906Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Let's check how one example looks like\nX, y = data_loader.get_xy(split=\"train\")\nX['xyz'][0][:3], X['time'][0][:3], X['charge'][0][:3], X['auxiliary'][0][:3]","metadata":{"execution":{"iopub.status.busy":"2023-09-15T10:08:14.448416Z","iopub.execute_input":"2023-09-15T10:08:14.448849Z","iopub.status.idle":"2023-09-15T10:08:14.498985Z","shell.execute_reply.started":"2023-09-15T10:08:14.448801Z","shell.execute_reply":"2023-09-15T10:08:14.495249Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"X['xyz'].shape, X['time'].shape, X['charge'].shape, X['auxiliary'].shape\")\nX['xyz'].shape, X['time'].shape, X['charge'].shape, X['auxiliary'].shape","metadata":{"execution":{"iopub.status.busy":"2023-09-15T10:08:14.502475Z","iopub.execute_input":"2023-09-15T10:08:14.503237Z","iopub.status.idle":"2023-09-15T10:08:14.513527Z","shell.execute_reply.started":"2023-09-15T10:08:14.503183Z","shell.execute_reply":"2023-09-15T10:08:14.512255Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Angular Loss function\n\nThis is the metric of this competition","metadata":{}},{"cell_type":"code","source":"def angular_dist_score(az_true, zen_true, az_pred, zen_pred):\n    '''\n    calculate the MAE of the angular distance between two directions.\n    The two vectors are first converted to cartesian unit vectors,\n    and then their scalar product is computed, which is equal to\n    the cosine of the angle between the two vectors. The inverse \n    cosine (arccos) thereof is then the angle between the two input vectors\n    \n    Parameters:\n    -----------\n    \n    az_true : float (or array thereof)\n        true azimuth value(s) in radian\n    zen_true : float (or array thereof)\n        true zenith value(s) in radian\n    az_pred : float (or array thereof)\n        predicted azimuth value(s) in radian\n    zen_pred : float (or array thereof)\n        predicted zenith value(s) in radian\n    \n    Returns:\n    --------\n    \n    dist : float\n        mean over the angular distance(s) in radian\n    '''\n    \n    if not (torch.all(torch.isfinite(az_true))  and\n            torch.all(torch.isfinite(zen_true)) and\n            torch.all(torch.isfinite(az_pred)) and\n            torch.all(torch.isfinite(zen_pred))\n           ):\n        raise ValueError(\"All arguments must be finite\")\n    \n    # pre-compute all sine and cosine values\n    sa1 = torch.sin(az_true)\n    ca1 = torch.cos(az_true)\n    sz1 = torch.sin(zen_true)\n    cz1 = torch.cos(zen_true)\n    \n    sa2 = torch.sin(az_pred)\n    ca2 = torch.cos(az_pred)\n    sz2 = torch.sin(zen_pred)\n    cz2 = torch.cos(zen_pred)\n    \n    # scalar product of the two cartesian vectors (x = sz*ca, y = sz*sa, z = cz)\n    scalar_prod = sz1*sz2*(ca1*ca2 + sa1*sa2) + (cz1*cz2)\n    \n    # scalar product of two unit vectors is always between -1 and 1, this is against nummerical instability\n    # that might otherwise occure from the finite precision of the sine and cosine functions\n    scalar_prod =  torch.clamp(scalar_prod, -1, 1)\n    \n    # convert back to an angle (in radian)\n    return torch.abs(torch.acos(scalar_prod))","metadata":{"execution":{"iopub.status.busy":"2023-09-15T10:08:14.515209Z","iopub.execute_input":"2023-09-15T10:08:14.515615Z","iopub.status.idle":"2023-09-15T10:08:14.528441Z","shell.execute_reply.started":"2023-09-15T10:08:14.515576Z","shell.execute_reply":"2023-09-15T10:08:14.526943Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# estimates the loss over eval_iters batches.\n@torch.no_grad()\ndef estimate_loss():\n    out = {'train': {}, 'test': {}}\n    model.eval()\n    for split in ['train', 'test']:\n        losses_distance = []\n        losses_angular = []\n        for k in range(CFG.eval_iters):\n            X, y = data_loader.get_xy(split)\n            model_out = model(X, y)\n            loss_distance = model_out['loss_distance']\n            loss_angular = model_out['loss_angular']\n            losses_distance.append(loss_distance.detach().cpu().numpy())\n            losses_angular.append(loss_angular.detach().cpu().numpy())\n        out[split]['distance'] = np.average(losses_distance)\n        out[split]['angular'] = np.average(losses_angular)\n    model.train()\n    return out","metadata":{"execution":{"iopub.status.busy":"2023-09-15T10:08:14.530799Z","iopub.execute_input":"2023-09-15T10:08:14.531102Z","iopub.status.idle":"2023-09-15T10:08:14.543729Z","shell.execute_reply.started":"2023-09-15T10:08:14.531076Z","shell.execute_reply":"2023-09-15T10:08:14.542790Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.utils.checkpoint as checkpoint\nimport math\nfrom timm.models.layers import drop_path, trunc_normal_\nimport math\nimport torch.utils.checkpoint as checkpoint\nfrom typing import Any, Callable, List, Optional, Sequence, Tuple, Union\nfrom torch import Tensor, LongTensor\nfrom timm.models.layers import drop_path, trunc_normal_\n\nclass DropPath(nn.Module):\n    def __init__(self, drop_prob=None):\n        super(DropPath, self).__init__()\n        self.drop_prob = drop_prob\n\n    def forward(self, x):\n        return drop_path(x, self.drop_prob, self.training)\n\n    def extra_repr(self) -> str:\n        return f\"p={self.drop_prob}\"\n\nclass SinusoidalPosEmb(nn.Module):\n    def __init__(self, dim):\n        super().__init__()\n        self.dim = dim\n\n    def forward(self, x):\n        device = x.device\n        pos = torch.arange(x.size(1), device=device).unsqueeze(1)\n        div_term = torch.exp(torch.arange(0, self.dim, 2, device=device) * -(math.log(10000.2) / self.dim))\n        emb = pos * div_term\n        emb = torch.cat([emb.sin(), emb.cos()], dim=1)\n        return emb\n    \nclass EmbeddingExtractor(nn.Module):\n    def __init__(self, input_features=6, embd_dim=384):\n        super().__init__()\n        self.emb = SinusoidalPosEmb(dim=embd_dim)\n        self.proj = nn.Sequential(\n            nn.Linear(input_features, embd_dim),\n            nn.LayerNorm(embd_dim),\n            nn.GELU(),\n            nn.Linear(embd_dim, embd_dim),\n        )\n\n    def forward(self, x):\n        # x is of shape (B, T, C)\n        \n        # Create positional embeddings of shape (T, dim_base)\n        pos_emb = self.emb(x)\n        \n        # Project x to shape (B, T, dim_base)\n        x = self.proj(x)\n        \n        # Add positional embeddings to x\n        x = x + pos_emb.unsqueeze(0)\n        \n        \n        return x\n\nclass MLP(nn.Module):\n    def __init__(\n        self,\n        in_features,\n        hidden_features=None,\n        out_features=None,\n        act_layer=nn.GELU,\n        drop=0.2,\n    ):\n        super().__init__()\n        out_features = out_features or in_features\n        hidden_features = hidden_features or in_features\n        self.fc1 = nn.Linear(in_features, hidden_features)\n        self.act = act_layer()\n        self.fc2 = nn.Linear(hidden_features, out_features)\n        self.drop = nn.Dropout(drop)\n\n    def forward(self, x):\n        x = self.fc1(x)\n        x = self.act(x)\n        # x = self.drop(x)\n        # commit this for the orignal BERT implement\n        x = self.fc2(x)\n        x = self.drop(x)\n        return x\n      \nclass SingleHeadAttention(nn.Module):\n    def __init__(self, embed_dim):\n        super(SingleHeadAttention, self).__init__()\n        self.embed_dim = embed_dim\n        self.query = nn.Linear(embed_dim, embed_dim, bias=False)\n        self.key = nn.Linear(embed_dim, embed_dim, bias=False)\n        self.value = nn.Linear(embed_dim, embed_dim, bias=False)\n        \n    def forward(self, x, mask=None):\n        # x is of shape [batch_size, block_size, embed_dim]\n        # mask is of shape [batch_size, block_size]\n        q = self.query(x)\n        k = self.key(x)\n        v = self.value(x)\n        attn_weights = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.embed_dim)\n        if mask is not None:\n            mask = mask.unsqueeze(1).expand(-1, x.size(1), -1)\n            attn_weights.masked_fill_(mask, float('-inf'))\n            # attn_weights.masked_fill_(mask, -1e9)\n        attn_weights = F.softmax(attn_weights, dim=-1)\n        output = torch.matmul(attn_weights, v)\n        return output\n\nclass CustomMultiheadAttention(nn.Module):\n    def __init__(self, embed_dim, num_heads):\n        super(CustomMultiheadAttention, self).__init__()\n        self.heads = nn.ModuleList([SingleHeadAttention(embed_dim) for _ in range(num_heads)])\n        self.fc = nn.Linear(num_heads * embed_dim, embed_dim)\n        \n    def forward(self, x, mask=None):\n        # x is of shape [batch_size, block_size, embed_dim]\n        # mask is of shape [batch_size, block_size]\n        outputs = [head(x, mask) for head in self.heads]\n        outputs = torch.cat(outputs, dim=-1)\n        output = self.fc(outputs)\n        return output\n    \n# BEiTv2 block (modified to have fixed number of tokens, to be used for CoreML)\nclass Block(nn.Module):\n    def __init__(\n        self,\n        embd_dim,\n        num_heads,\n        mlp_ratio=4.0,\n        qkv_bias=False,\n        qk_scale=None,\n        drop=0.2,\n        attn_drop=0.2,\n        drop_path=0.2,\n        init_values=None,\n        act_layer=nn.GELU,\n        # act_layer=nn.Tanh,\n        norm_layer=nn.LayerNorm,\n        window_size=None,\n        attn_head_dim=None,\n        **kwargs,\n    ):\n        super().__init__()\n        self.norm1 = norm_layer(embd_dim)\n        self.attn = CustomMultiheadAttention(embd_dim, num_heads)\n        self.drop_path = DropPath(drop_path) if drop_path > 0.2 else nn.Identity()\n        self.norm2 = norm_layer(embd_dim)\n        mlp_hidden_dim = int(embd_dim * mlp_ratio)\n        self.mlp = MLP(\n            in_features=embd_dim,\n            hidden_features=mlp_hidden_dim,\n            act_layer=act_layer,\n            drop=drop,\n        )\n\n        if init_values is not None:\n            self.gamma_1 = nn.Parameter(\n                init_values * torch.ones((embd_dim)), requires_grad=True\n            )\n            self.gamma_2 = nn.Parameter(\n                init_values * torch.ones((embd_dim)), requires_grad=True\n            )\n        else:\n            self.gamma_1, self.gamma_2 = None, None\n\n    def forward(self, x, attn_mask=None, key_padding_mask=None):\n        if self.gamma_1 is None:\n            xn = self.norm1(x)\n            x = x + self.drop_path(self.attn(xn, attn_mask))\n            x = x + self.drop_path(self.mlp(self.norm2(x)))\n        else:\n            xn = self.norm1(x)\n            x = x + self.drop_path(self.gamma_1 * self.attn(xn, attn_mask))\n            x = x + self.drop_path(self.gamma_2 * self.mlp(self.norm2(x)))\n        return x\n\n\nclass NeutrinoGPT(nn.Module):\n    def __init__(\n        self,\n        input_features=108,\n        embd_dim=384,\n        depth=12,\n        head_size=32,\n        drop_path=0.2,\n        **kwargs,\n    ):\n        super().__init__()\n        self.extractor = EmbeddingExtractor(input_features, embd_dim)\n        self.blocks = nn.ModuleList(\n            [\n                Block(\n                    embd_dim=embd_dim,\n                    num_heads=embd_dim // head_size,\n                    mlp_ratio=4,\n                    drop_path=drop_path * (i / (depth - 1)),\n                    init_values=1,\n                )\n                for i in range(depth)\n            ]\n        )\n        self.num_heads=embd_dim // head_size\n        self.cls_token = nn.Linear(embd_dim, 1, bias=False)\n        self.proj_out = nn.Linear(embd_dim, 3, bias=False) \n        \n        self.apply(self._init_weights)\n        trunc_normal_(self.cls_token.weight, std=0.02)\n\n    def fix_init_weight(self):\n        def rescale(param, layer_id):\n            param.div_(math.sqrt(2.0 * layer_id))\n\n        for layer_id, layer in enumerate(self.blocks):\n            rescale(layer.attn.proj.weight.data, layer_id + 1)\n            rescale(layer.mlp.fc2.weight.data, layer_id + 1)\n\n    def _init_weights(self, m):\n        if isinstance(m, nn.Linear):\n            trunc_normal_(m.weight, std=0.02)\n            if isinstance(m, nn.Linear) and m.bias is not None:\n                nn.init.constant_(m.bias, 0)\n        elif isinstance(m, nn.LayerNorm):\n            nn.init.constant_(m.bias, 0)\n            nn.init.constant_(m.weight, 1.0)\n\n    def init_weights(self, pretrained=None):\n        def _init_weights(m):\n            if isinstance(m, nn.Linear):\n                trunc_normal_(m.weight, std=0.02)\n                if isinstance(m, nn.Linear) and m.bias is not None:\n                    nn.init.constant_(m.bias, 0)\n            elif isinstance(m, nn.LayerNorm):\n                nn.init.constant_(m.bias, 0)\n                nn.init.constant_(m.weight, 1.0)\n\n        self.apply(_init_weights)\n\n    @torch.jit.ignore\n    def no_weight_decay(self):\n        return {\"cls_token\"}\n\n    def forward(self, X, targets=None):\n        \n        xyz, time, charge, event_size, auxiliary = X['xyz'], X['time'], X['charge'], X['event_size'], X['auxiliary']\n        x = torch.cat([xyz, time.unsqueeze(-1), charge.unsqueeze(-1), auxiliary.unsqueeze(-1)], dim=-1)\n        \n        x = self.extractor(x)\n        \n        # Attach a CLS token\n        B, T, C = x.shape\n        cls_token = self.cls_token.weight.unsqueeze(0).expand(B, -1, -1) # (B, 1, C)\n        x = torch.cat([cls_token, x], 1) # (B, T+1, C)\n\n        for i, blk in enumerate(self.blocks):\n            x = blk(x)\n\n        #  Get the CLS token\n        x = self.proj_out(x[:, 0])  # cls token\n        \n        # normalize according to the norm\n        x = x / torch.norm(x, dim=1, keepdim=True)\n        \n        # Get the angle theta and phi from the 3D vector\n        theta = torch.acos(x[:, 2])\n        phi = torch.atan2(x[:, 1], x[:, 0])\n        # if phi is negative, add 2pi\n        phi = torch.where(phi < 0, phi + 2*np.pi, phi)\n        \n        \n        \n        if targets is None:\n            loss_distance = None\n            loss_angular = None\n        else: \n            \n            vec_target = targets\n            loss_distance = torch.sum(event_size * torch.sum((vec_target - x)**2, axis=1)) / torch.sum(event_size)\n            \n            \n            # Get the angle theta and phi from the 3D vector\n            theta_target = torch.acos(vec_target[:, 2])\n            phi_target = torch.atan2(vec_target[:, 1], vec_target[:, 0])\n            \n            \n            \n            # if phi is negative, add 2pi\n            phi_target = torch.where(phi_target < 0, phi_target + 2*np.pi, phi_target)\n            \n            \n            losses_angular = angular_dist_score(phi_target, theta_target, phi, theta)\n            # losses_angular = angular_dist_score(phi_target, theta_target, phi, theta)\n            loss_angular = torch.sum(event_size * losses_angular) / torch.sum(event_size)\n            # loss_angular = torch.mean(losses_angular)\n\n        \n        out = {\n            'x': x,\n            'loss_distance': loss_distance,\n            'theta': theta, \n            'phi': phi, \n            'loss_angular': loss_angular\n        }\n        \n        return out\n        \n    \ndef load_weights(model, weights_path, strict=True):\n\n    \n    # Save the original state dict for comparison\n    original_state_dict = model.state_dict().copy()\n\n\n    # Load weights\n    loaded_weights = torch.load(weights_path, map_location=torch.device('cpu'))\n    \n    # If the saved model was saved with DataParallel\n    loaded_weights = {k.replace('module.', ''): v for k, v in loaded_weights.items()}\n    \n    # Verify if weights are the same\n    for ((layer_name, original_weight), (_, loaded_weight)) in zip(original_state_dict.items(), loaded_weights.items()):\n        if torch.equal(original_weight, loaded_weight):\n            raise Exception(f\"[ERROR] Layer {layer_name} weights remain the same.\")\n    \n    # Check if the model architecture and loaded weights match\n    if set(model.state_dict().keys()) != set(loaded_weights.keys()):\n        raise Exception(\"Model's architecture does not match with the saved weights. Ensure they are compatible.\")\n    \n    # Verify if weights are loaded correctly, ie that they are not the same as current weights\n    for ((layer_name, original_weight), (_, loaded_weight)) in zip(original_state_dict.items(), loaded_weights.items()):\n        if torch.equal(original_weight, loaded_weight):\n            print(f\"Original: {original_weight}\")\n            print(f\"Loaded: {loaded_weight}\")\n            raise Exception(f\"[ERROR] Layer {layer_name} weights remain the same.\")\n            \n    # custom implementation of model.load_state_dict(loaded_weights)\n    for name, param in model.named_parameters():\n        if name in loaded_weights.keys():\n            try:\n                param.data.copy_(loaded_weights[name])\n            except:\n                if strict:\n                    raise Exception(f\"Layer {name} weights could not be loaded.\")\n                else:\n                    print(f\"Layer {name} weights could not be loaded.\")\n        else:\n            raise Exception(f\"Layer {name} weights not found in saved weights.\")\n        \n    print(\"Weights loaded successfully.\")\n    return model","metadata":{"execution":{"iopub.status.busy":"2023-09-15T10:08:14.545664Z","iopub.execute_input":"2023-09-15T10:08:14.546265Z","iopub.status.idle":"2023-09-15T10:08:14.631662Z","shell.execute_reply.started":"2023-09-15T10:08:14.546220Z","shell.execute_reply":"2023-09-15T10:08:14.630718Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = NeutrinoGPT(input_features=6, \n                    embd_dim=CFG.n_embd, \n                    depth=CFG.n_layers, \n                    head_size=CFG.head_size, \n                    drop_path=CFG.drop_path)\n\nmodel = model.to(CFG.device)\n# print the number of parameters in the model\nprint(sum(p.numel() for p in model.parameters())/1e6, 'M parameters')","metadata":{"execution":{"iopub.status.busy":"2023-09-15T10:08:14.634017Z","iopub.execute_input":"2023-09-15T10:08:14.634378Z","iopub.status.idle":"2023-09-15T10:08:25.173403Z","shell.execute_reply.started":"2023-09-15T10:08:14.634342Z","shell.execute_reply":"2023-09-15T10:08:25.172281Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ndef print_params(model):\n    print(\"looping over model params to see if there is some weird behaviour\")\n    for name, param in model.named_parameters():\n        m, s = torch.mean(param), torch.std(param)\n        if s > 0.04 or m > 0.04:\n            print(f'{name}: {m:.4f} +/- {s:.4f}')\n    print(\"loop finished\")\n\nprint_params(model)","metadata":{"execution":{"iopub.status.busy":"2023-09-15T10:08:25.175089Z","iopub.execute_input":"2023-09-15T10:08:25.175475Z","iopub.status.idle":"2023-09-15T10:08:25.455278Z","shell.execute_reply.started":"2023-09-15T10:08:25.175438Z","shell.execute_reply":"2023-09-15T10:08:25.454222Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X, y = data_loader.get_xy(\"train\")\nwith torch.no_grad():\n    out = model(X, y)\n\nx = out['x']\nprint(torch.mean(x), torch.std(x))\nprint(\"x[0].shape, x[1].shape\", x[0].shape, x[1].shape)","metadata":{"execution":{"iopub.status.busy":"2023-09-15T10:08:25.459473Z","iopub.execute_input":"2023-09-15T10:08:25.459786Z","iopub.status.idle":"2023-09-15T10:08:25.737723Z","shell.execute_reply.started":"2023-09-15T10:08:25.459756Z","shell.execute_reply":"2023-09-15T10:08:25.736685Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%time x = model(X)","metadata":{"execution":{"iopub.status.busy":"2023-09-15T10:08:25.742192Z","iopub.execute_input":"2023-09-15T10:08:25.744619Z","iopub.status.idle":"2023-09-15T10:08:26.063917Z","shell.execute_reply.started":"2023-09-15T10:08:25.744579Z","shell.execute_reply":"2023-09-15T10:08:26.062910Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%time x = model(X, y)","metadata":{"execution":{"iopub.status.busy":"2023-09-15T10:08:26.066260Z","iopub.execute_input":"2023-09-15T10:08:26.066907Z","iopub.status.idle":"2023-09-15T10:08:26.370905Z","shell.execute_reply.started":"2023-09-15T10:08:26.066865Z","shell.execute_reply":"2023-09-15T10:08:26.369775Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"# start a new wandb run to track this script\nif LOG_WANDB:\n    from kaggle_secrets import UserSecretsClient\n    user_secrets = UserSecretsClient()\n    secret_value = user_secrets.get_secret(\"wandb_key\")\n    \n    wandb.login(key=secret_value)\n    \n    wandb.init(\n        # set the wandb project where this run will be logged\n        project=\"neutrino-detection\",\n\n        # track hyperparameters and run metadata\n        config={\n        \"n_layers\":n_layers,\n        \"num_heads\":num_heads\n        }\n    )\n    ","metadata":{"execution":{"iopub.status.busy":"2023-09-15T10:08:26.379384Z","iopub.execute_input":"2023-09-15T10:08:26.379718Z","iopub.status.idle":"2023-09-15T10:08:26.385894Z","shell.execute_reply.started":"2023-09-15T10:08:26.379688Z","shell.execute_reply":"2023-09-15T10:08:26.384708Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# create a PyTorch optimizer\noptimizer = torch.optim.Adam(model.parameters(), lr=CFG.learning_rate)\n\nfor iter in tqdm(range(CFG.max_iters)):\n\n    # every once in a while evaluate the loss on train and val sets\n    if iter % CFG.evaluate_every_step == 0 or iter == CFG.max_iters - 1:\n        \n        losses = estimate_loss()\n        \n        print(f\"step {iter}: train distance loss {losses['train']['distance']:.4f}, test distance loss {losses['test']['distance']:.4f} train angular loss {losses['train']['angular']:.4f}, test angular loss {losses['test']['angular']:.4f} \")\n        if LOG_WANDB:\n            wandb.log({\"train_dis_loc\": losses['train']['distance'],\n                       \"test_dis_loc\": losses['test']['distance'],\n                       \"train_ang_loc\": losses['train']['angular'],\n                       \"test_ang_loc\": losses['test']['angular']})\n\n            # log the model parameters\n            wandb.log({\"model\": model.state_dict()})\n            torch.save(model.state_dict(), 'model.pth')\n            \n        \n\n        # send model weights to wandb\n        \n        \n    X, y = data_loader.get_xy('train')\n\n    # evaluate the loss\n    out = model(X, y)\n    \n    x = out['x']\n    loss_distance = out['loss_distance']\n    theta = out['theta']\n    phi = out['phi']\n    loss_angular = out['loss_angular']\n    \n    optimizer.zero_grad(set_to_none=True)\n    if iter < CFG.n_opt_distance:\n        loss_distance.backward()\n    else:\n        loss_angular.backward()\n        \n    \n    optimizer.step()\nif LOG_WANDB:\n    wandb.finish()","metadata":{"execution":{"iopub.status.busy":"2023-09-15T10:08:26.387435Z","iopub.execute_input":"2023-09-15T10:08:26.388143Z","iopub.status.idle":"2023-09-15T10:09:16.384260Z","shell.execute_reply.started":"2023-09-15T10:08:26.388077Z","shell.execute_reply":"2023-09-15T10:09:16.382796Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(model.state_dict(), 'model.pth')","metadata":{"execution":{"iopub.status.busy":"2023-09-15T10:09:16.385802Z","iopub.status.idle":"2023-09-15T10:09:16.386569Z","shell.execute_reply.started":"2023-09-15T10:09:16.386300Z","shell.execute_reply":"2023-09-15T10:09:16.386326Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Make predictions","metadata":{}},{"cell_type":"code","source":"# model.eval()\n\n# batch_size = 64\n# predictions = []\n# for i in range(0, len(df), batch_size):\n#     first_pulse_indices = df[\"first_pulse_index\"][i:i+batch_size].tolist()\n#     last_pulse_indices = df[\"last_pulse_index\"][i:i+batch_size].tolist()\n\n#     # preprocessed data as tensors\n#     preprocessed_data = [data_loader.get_single_event(first_pulse_index, last_pulse_index, 'eval') for first_pulse_index, last_pulse_index in zip(first_pulse_indices, last_pulse_indices)]\n\n#     # create batch of preprocessed data\n#     xyz = torch.stack([data[\"xyz\"] for data in preprocessed_data])\n#     time = torch.stack([data[\"time\"] for data in preprocessed_data])\n#     charge = torch.stack([data[\"charge\"] for data in preprocessed_data])\n    \n#     xyz = xyz.to(device)\n#     time = time.to(device)\n#     charge = charge.to(device)\n\n#     # make predictions\n#     with torch.no_grad():\n#         (x, _), (theta, phi, _) = model(xyz, time, charge)\n    \n    \n#     theta_predictions = theta.detach().cpu().numpy()\n#     phi_predictions = phi.detach().cpu().numpy()\n    \n#     print(theta_predictions, phi_predictions)\n#     pred = np.array([theta_predictions, phi_predictions])\n    \n#     # predictions.append(pred.T.tolist()[)\n    \n\n# print(predictions)\n# # df[\"prediction\"] = predictions\n# # df[\"azimuth\"] = df[\"prediction\"].apply(lambda x: x[0])\n# # df[\"zenith\"] = df[\"prediction\"].apply(lambda x: x[1])\n\n# # event_id,azimuth,zenith\n# # df = df[[\"event_id\", \"azimuth\", \"zenith\"]]\n","metadata":{"execution":{"iopub.status.busy":"2023-09-15T10:09:16.388074Z","iopub.status.idle":"2023-09-15T10:09:16.388854Z","shell.execute_reply.started":"2023-09-15T10:09:16.388571Z","shell.execute_reply":"2023-09-15T10:09:16.388597Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# df = df.sort_values([\"event_id\"])\n# df.to_csv('submission.csv', index=False)\n# df","metadata":{"execution":{"iopub.status.busy":"2023-09-15T10:09:16.390258Z","iopub.status.idle":"2023-09-15T10:09:16.391022Z","shell.execute_reply.started":"2023-09-15T10:09:16.390746Z","shell.execute_reply":"2023-09-15T10:09:16.390783Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}