{"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":"Added center of charge from https://www.kaggle.com/code/roberthatch/lb-1-183-lightning-fast-baseline-with-polars#Section-1---Center-of-Charge-with-Pandas but using an encoder transformer with a sigmoid output instead of the charge.\n(using the transformer to predict the azimuth and zenith directly didn't work)\nEven so can't get the model to overfit so there must be some problem","metadata":{}},{"cell_type":"code","source":"!pip install einops","metadata":{"execution":{"iopub.status.busy":"2023-01-29T17:42:44.104107Z","iopub.execute_input":"2023-01-29T17:42:44.104506Z","iopub.status.idle":"2023-01-29T17:42:53.595769Z","shell.execute_reply.started":"2023-01-29T17:42:44.104475Z","shell.execute_reply":"2023-01-29T17:42:53.594482Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport random\nimport math\nfrom pathlib import Path\nfrom collections import Counter\nimport datetime\nimport gc\nimport time\n\nfrom tqdm.notebook import tqdm\nimport pandas as pd\nimport numpy as np\nimport pyarrow.parquet as pq\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader, TensorDataset\nfrom torch.nn.utils.rnn import pad_sequence\n\nfrom sklearn.model_selection import train_test_split, KFold\n# from sklearn.preprocessing import StandardScaler, scale, MinMaxScaler\n# from sklearn.decomposition import TruncatedSVD\n\nfrom einops import rearrange","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-01-29T17:42:53.598526Z","iopub.execute_input":"2023-01-29T17:42:53.598933Z","iopub.status.idle":"2023-01-29T17:42:53.608821Z","shell.execute_reply.started":"2023-01-29T17:42:53.598891Z","shell.execute_reply":"2023-01-29T17:42:53.606716Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{"execution":{"iopub.status.busy":"2023-01-29T17:42:53.610532Z","iopub.execute_input":"2023-01-29T17:42:53.611434Z","iopub.status.idle":"2023-01-29T17:42:53.620654Z","shell.execute_reply.started":"2023-01-29T17:42:53.611398Z","shell.execute_reply":"2023-01-29T17:42:53.619541Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"SEED = 333\ndef seedBasic(seed=SEED):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    \n    \ndef seedTorch(seed=SEED):\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n      \n# basic + torch \ndef seedEverything(seed=SEED):\n    seedBasic(seed)\n    seedTorch(seed)\n\nseedEverything()","metadata":{"execution":{"iopub.status.busy":"2023-01-29T17:42:53.624028Z","iopub.execute_input":"2023-01-29T17:42:53.624339Z","iopub.status.idle":"2023-01-29T17:42:53.633313Z","shell.execute_reply.started":"2023-01-29T17:42:53.624293Z","shell.execute_reply":"2023-01-29T17:42:53.632091Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_path = Path(\"/kaggle/input/icecube-neutrinos-in-deep-ice/\")","metadata":{"execution":{"iopub.status.busy":"2023-01-29T17:42:53.634950Z","iopub.execute_input":"2023-01-29T17:42:53.635764Z","iopub.status.idle":"2023-01-29T17:42:53.643238Z","shell.execute_reply.started":"2023-01-29T17:42:53.635729Z","shell.execute_reply":"2023-01-29T17:42:53.642120Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sensor_geometry = pd.read_csv(data_path / \"sensor_geometry.csv\")\nx_min, x_max, y_min, y_max, z_min, z_max = sensor_geometry.x.min(), sensor_geometry.x.max(), sensor_geometry.y.min(), sensor_geometry.y.max(), sensor_geometry.z.min(), sensor_geometry.z.max()\nsensor_geometry['x'] = (sensor_geometry['x'] - x_min) / (x_max - x_min)\nsensor_geometry['y'] = (sensor_geometry['y'] - y_min) / (y_max - y_min)\nsensor_geometry['z'] = (sensor_geometry['z'] - z_min) / (z_max - z_min)\nprint(f\"Shape: {sensor_geometry.shape}\")\nsensor_geometry.head(10)","metadata":{"execution":{"iopub.status.busy":"2023-01-29T17:42:53.644821Z","iopub.execute_input":"2023-01-29T17:42:53.645510Z","iopub.status.idle":"2023-01-29T17:42:53.673413Z","shell.execute_reply.started":"2023-01-29T17:42:53.645474Z","shell.execute_reply":"2023-01-29T17:42:53.672470Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def angles_from_vectors(vectors):\n    v_squared = np.square(vectors)\n    \n    ## Shortcut optimization for azimuth: calculate 2d unit vectors for x and y independent of z\n    xy_sq = np.sum(v_squared[:, 0:2], axis=1)\n    xy_d = np.sqrt(xy_sq)[:, None]\n    np.seterr(divide='ignore', invalid='ignore') ## Turn off the warning temporarily\n    vectors[:, 0:2] = np.where(xy_d == 0, xy_d, vectors[:, 0:2]/xy_d)\n\n    ## For z, use full 3d unit vector\n    d = np.sqrt(xy_sq + v_squared[:, 2])\n    vectors[:, 2] = np.where(d == 0, d, vectors[:, 2]/d)\n    np.seterr(divide='warn', invalid='warn') ## Turn back on\n\n    ## As mentioned by others, clip solely to avoid floating point errors, the unit vectors should already be within this range.\n    vectors =  np.clip(vectors, -1, 1)\n\n    azimuth = np.arccos(vectors[:, 0])\n    ## if y < 0, convert from quadrants 1 and 2 to quadrants 3 and 4\n    azimuth = np.where(vectors[:, 1] >= 0, azimuth, 2*math.pi - azimuth)\n    azimuth = np.where(np.isfinite(azimuth), azimuth, 0.0)\n\n    zenith = np.arccos(vectors[:, 2])\n    ## IMPORTANT: zenith angles are not evenly distributed, so set the error case to pi/2!\n    ## (even though x, y, z might be. It would be a fun exercise to check if random values\n    ##  for x, y, z converted to zenith angles would match the observed distribution of zenith angles in the train labels)\n    zenith = np.where(np.isfinite(zenith), zenith, math.pi/2)\n\n    return np.stack([azimuth, zenith], axis=1)","metadata":{"execution":{"iopub.status.busy":"2023-01-29T17:42:53.674952Z","iopub.execute_input":"2023-01-29T17:42:53.675649Z","iopub.status.idle":"2023-01-29T17:42:53.685796Z","shell.execute_reply.started":"2023-01-29T17:42:53.675613Z","shell.execute_reply":"2023-01-29T17:42:53.684552Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class EventDataset():\n    def __init__(self, df, train=True):\n        self.df = df\n        batch_id = df['batch_id'].unique()[0]\n        folder = 'train' if train else 'test'\n        self.batch = pq.ParquetDataset(data_path / folder / f\"batch_{batch_id}.parquet\", use_legacy_dataset=False).read().to_pandas()\n        self.train = train\n        \n    def __getitem__(self, index):\n        row = self.df.iloc[index]\n        event_id = int(row['event_id'])\n        events = self.batch.loc[event_id]\n        events = events.join(sensor_geometry.set_index('sensor_id'), on='sensor_id')\n        events['auxiliary'] = events['auxiliary'].astype(int)\n        \n#         events['zenith'] = np.arccos(events['z'] / np.sqrt(events['x']**2 + events['y']**2 + events['z']**2))\n#         events['azimuth'] = np.arctan2(events['y'], events['x'])\n#         events.loc[events['azimuth'] < 0, 'azimuth'] = events.loc[events['azimuth'] < 0, 'azimuth'] + 2 * np.pi\n        \n        min_time, max_time = events['time'].min(), events['time'].max()\n        events['time'] = (events['time'] - min_time) / (max_time - min_time)\n        \n        min_time, max_time = events['charge'].min(), events['charge'].max()\n        events['charge'] = (events['charge'] - min_time) / (max_time - min_time)\n        \n        features = [\n            'time', 'charge', 'auxiliary',\n            'x', 'y', 'z',\n#             'azimuth', 'zenith'\n        ]\n        target = ['azimuth', 'zenith']\n        \n        if self.train:\n            return events[features].values, row[target].values\n        else:\n            return events[features].values\n\n    def __len__(self):\n        return len(self.df)","metadata":{"execution":{"iopub.status.busy":"2023-01-29T17:42:53.687990Z","iopub.execute_input":"2023-01-29T17:42:53.688743Z","iopub.status.idle":"2023-01-29T17:42:53.701492Z","shell.execute_reply.started":"2023-01-29T17:42:53.688663Z","shell.execute_reply":"2023-01-29T17:42:53.700291Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_collate(max_num_pulses, train = True):\n    def proccess_inputs(inputs):\n        inputs = pad_sequence(inputs, batch_first=True).numpy()\n        pulses_num = inputs.shape[1]\n        new_inputs = []\n        j=0\n        while j < pulses_num:\n            if j+max_num_pulses > pulses_num:\n                break\n            new_inputs.append(inputs[:,j:j+max_num_pulses])\n            j += max_num_pulses\n\n        left = pulses_num % max_num_pulses\n        if left > 0:\n            temp = torch.tensor(inputs[:,j:j+max_num_pulses], dtype=torch.float)\n            s = [d for d in temp.shape]\n            s[1] = max_num_pulses - s[1]\n            p = torch.zeros(s, dtype=torch.float)\n            temp = torch.cat([temp, p], dim=1)\n            new_inputs.append(temp.numpy())\n\n        new_inputs = torch.tensor(np.array(new_inputs), dtype=torch.float, device=device)\n        new_inputs = new_inputs.permute(1, 0, 2, 3)\n\n        return new_inputs\n\n    def train_collate(data):\n        inputs=[]\n        labels=[]\n        for x, y in data:\n            inputs.append(torch.tensor(x, dtype=torch.float))\n            labels.append(y)\n\n        labels = torch.tensor(labels, dtype=torch.float, device=device)\n\n        inputs = proccess_inputs(inputs)\n\n        return inputs, labels\n    \n    def test_collate(data):\n        inputs = proccess_inputs(data)\n\n        return inputs\n    \n    if train:\n        return train_collate\n    return test_collate","metadata":{"execution":{"iopub.status.busy":"2023-01-29T17:42:53.703616Z","iopub.execute_input":"2023-01-29T17:42:53.704056Z","iopub.status.idle":"2023-01-29T17:42:53.717052Z","shell.execute_reply.started":"2023-01-29T17:42:53.704017Z","shell.execute_reply":"2023-01-29T17:42:53.716118Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def exists(val):\n    return val is not None\n\ndef default(val, d):\n    return val if exists(val) else d","metadata":{"execution":{"iopub.status.busy":"2023-01-29T17:42:53.722634Z","iopub.execute_input":"2023-01-29T17:42:53.723322Z","iopub.status.idle":"2023-01-29T17:42:53.730559Z","shell.execute_reply.started":"2023-01-29T17:42:53.723289Z","shell.execute_reply":"2023-01-29T17:42:53.729518Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class LayerNorm(nn.Module):\n    def __init__(self, config):\n        super().__init__()\n        self.weight = nn.Parameter(torch.ones(config[\"d_model\"]))\n        self.variance_epsilon = config[\"norm_eps\"]\n\n    def forward(self, x):\n\n        variance = x.to(torch.float32).pow(2).mean(-1, keepdim=True)\n        x = x * torch.rsqrt(variance + self.variance_epsilon)\n\n        return self.weight * x","metadata":{"execution":{"iopub.status.busy":"2023-01-29T17:42:53.732040Z","iopub.execute_input":"2023-01-29T17:42:53.732647Z","iopub.status.idle":"2023-01-29T17:42:53.741546Z","shell.execute_reply.started":"2023-01-29T17:42:53.732613Z","shell.execute_reply":"2023-01-29T17:42:53.740649Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class FeedForward(nn.Module):\n    def __init__(self, config):\n        super().__init__()\n        self.wi_0 = nn.Linear(config[\"d_model\"], config[\"d_ff\"], bias=False)\n        self.wi_1 = nn.Linear(config[\"d_model\"], config[\"d_ff\"], bias=False)\n        self.wo = nn.Linear(config[\"d_ff\"], config[\"d_model\"], bias=False)\n        self.dropout = nn.Dropout(config[\"dropout_rate\"])\n        self.act = nn.SiLU()\n\n    def forward(self, x):\n        x_gelu = self.act(self.wi_0(x))\n        x_linear = self.wi_1(x)\n        x = x_gelu * x_linear\n        x = self.dropout(x)\n        x = self.wo(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2023-01-29T17:42:53.742974Z","iopub.execute_input":"2023-01-29T17:42:53.743330Z","iopub.status.idle":"2023-01-29T17:42:53.753004Z","shell.execute_reply.started":"2023-01-29T17:42:53.743296Z","shell.execute_reply":"2023-01-29T17:42:53.751890Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SelfAttention(nn.Module):\n    def __init__(self, config):\n        super().__init__()\n        inner_dim = config[\"dim_head\"] * config[\"heads\"]\n        self.heads = config[\"heads\"]\n        self.scale = config[\"dim_head\"] ** -0.5\n\n        self.to_q = nn.Linear(config[\"d_model\"], inner_dim, bias = False)\n        self.to_k = nn.Linear(config[\"d_model\"], inner_dim, bias = False)\n        self.to_v = nn.Linear(config[\"d_model\"], inner_dim, bias = False)\n        self.to_out = nn.Linear(inner_dim, config[\"d_model\"])\n\n        self.dropout = nn.Dropout(config[\"dropout_rate\"])\n\n    def forward(self, x, mask = None):\n        b, p, n, _, h = *x.shape, self.heads\n        q, k, v = self.to_q(x), self.to_k(x), self.to_v(x)\n\n        q, k, v = map(lambda t: rearrange(t, 'b p n (h d) -> b p h n d', h = h), (q, k, v))\n\n        q = q * self.scale\n\n        sim = torch.einsum('b p h i d, b p h j d -> b p h i j', q, k)\n\n        # mask\n\n        mask_value = -torch.finfo(sim.dtype).max\n\n        if mask is not None:\n            sim = sim.masked_fill_(~mask, mask_value)\n\n\n        # attention\n\n        attn = sim.softmax(dim = -1)\n        attn = self.dropout(attn)\n\n        # aggregate\n\n        out = torch.einsum('b p h i j, b p h j d -> b p h i d', attn, v)\n        \n        # merge heads\n\n        out = rearrange(out, 'b p h n d -> b p n (h d)')\n        \n        # combine heads and linear output\n\n        return self.to_out(out)","metadata":{"execution":{"iopub.status.busy":"2023-01-29T17:42:53.754664Z","iopub.execute_input":"2023-01-29T17:42:53.755072Z","iopub.status.idle":"2023-01-29T17:42:53.767936Z","shell.execute_reply.started":"2023-01-29T17:42:53.755035Z","shell.execute_reply":"2023-01-29T17:42:53.767057Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SublayerConnection(nn.Module):\n\n    def __init__(self, config):\n        super(SublayerConnection, self).__init__()\n        self.norm = LayerNorm(config)\n        self.dropout = nn.Dropout(config[\"dropout_rate\"])\n\n    def forward(self, sublayer, x, **kwargs):\n        return x + self.dropout(sublayer(self.norm(x), **kwargs))","metadata":{"execution":{"iopub.status.busy":"2023-01-29T17:42:53.769532Z","iopub.execute_input":"2023-01-29T17:42:53.769888Z","iopub.status.idle":"2023-01-29T17:42:53.778264Z","shell.execute_reply.started":"2023-01-29T17:42:53.769854Z","shell.execute_reply":"2023-01-29T17:42:53.777196Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class EncoderLayer(nn.Module):\n    def __init__(self, config):\n        super().__init__()\n        self.att = SelfAttention(config)\n        self.mlp = FeedForward(config)\n        \n        self.sublayer1 = SublayerConnection(config)\n        self.sublayer2 = SublayerConnection(config)\n\n    def forward(self, x, mask = None):\n        x = self.sublayer1(self.att, x, mask = mask)\n        \n        x = self.sublayer1(self.mlp, x)\n        \n        return x","metadata":{"execution":{"iopub.status.busy":"2023-01-29T17:42:53.779919Z","iopub.execute_input":"2023-01-29T17:42:53.780741Z","iopub.status.idle":"2023-01-29T17:42:53.790539Z","shell.execute_reply.started":"2023-01-29T17:42:53.780704Z","shell.execute_reply":"2023-01-29T17:42:53.789747Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class FirstLayer(nn.Module):\n    def __init__(self, input_dim, hidden_dim, output_dim=None, dropout=0.1):\n        super().__init__()\n        self.input_dim = input_dim\n        self.hidden_dim = hidden_dim\n        self.output_dim = output_dim if output_dim else input_dim\n        self.dropout = dropout\n        \n        self.wi_0 = nn.Linear(self.input_dim, self.hidden_dim, bias=False)\n        self.wi_1 = nn.Linear(self.input_dim, self.hidden_dim, bias=False)\n        self.wo = nn.Linear(self.hidden_dim, self.output_dim, bias=False)\n        self.dropout = nn.Dropout(self.dropout)\n        self.act = nn.SiLU()\n\n    def forward(self, x):\n        x_gelu = self.act(self.wi_0(x))\n        x_linear = self.wi_1(x)\n        x = x_gelu * x_linear\n        x = self.dropout(x)\n        x = self.wo(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2023-01-29T17:42:53.791867Z","iopub.execute_input":"2023-01-29T17:42:53.792524Z","iopub.status.idle":"2023-01-29T17:42:53.801397Z","shell.execute_reply.started":"2023-01-29T17:42:53.792489Z","shell.execute_reply":"2023-01-29T17:42:53.800527Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_targets(x, y):\n    t = x[:, :, 0]\n    c = y.squeeze(-1)#torch.mul(y.squeeze(-1), x[:, :, 1])\n    ox = x[:, :, 3]\n    oy = x[:, :, 4]\n    oz = x[:, :, 5]\n\n    ev_t_min = torch.min(t, 1).values\n    ev_t_max = torch.max(t, 1).values\n    \n    ev_t_min = ev_t_min.unsqueeze(-1)\n    ev_t_max = ev_t_max.unsqueeze(-1)\n    \n    ## Now we can just implement our formula! w0 and w1 are the time-weighted charge cases.\n    ## Gather the values we need\n    \n    w1 = torch.div(torch.mul(c, t - ev_t_min), (ev_t_max - ev_t_min))\n    w0 = c - w1\n    wx0 = torch.mul(ox, w0)\n    wy0 = torch.mul(oy, w0)\n    wz0 = torch.mul(oz, w0)\n    wx1 = torch.mul(ox, w1)\n    wy1 = torch.mul(oy, w1)\n    wz1 = torch.mul(oz, w1)\n    \n    w0 = torch.sum(w0, dim=1)\n    w1 = torch.sum(w1, dim=1)\n    wx0 = torch.sum(wx0, dim=1)\n    wy0 = torch.sum(wy0, dim=1)\n    wz0 = torch.sum(wz0, dim=1)\n    wx1 = torch.sum(wx1, dim=1)\n    wy1 = torch.sum(wy1, dim=1)\n    wz1 = torch.sum(wz1, dim=1)\n    \n    wx0 = torch.div(wx0, w0)\n    wy0 = torch.div(wy0, w0)\n    wz0 = torch.div(wz0, w0)\n    \n    wx1 = torch.div(wx1, w1)\n    wy1 = torch.div(wy1, w1)\n    wz1 = torch.div(wz1, w1)\n    \n    ox = wx0 - wx1\n    oy = wy0 - wy1\n    oz = wz0 - wz1\n    \n    zenith = torch.arccos(oz / torch.sqrt(ox**2 + oy**2 + oz**2))\n    azimuth = torch.arctan2(oy, ox)\n    azimuth = torch.add(azimuth, (azimuth<0).float() * 2 * torch.tensor(math.pi))\n\n    azimuth = azimuth.unsqueeze(1)\n    zenith = zenith.unsqueeze(1)\n    target = torch.cat([azimuth, zenith], dim=1)\n    \n    return target","metadata":{"execution":{"iopub.status.busy":"2023-01-29T17:42:53.802891Z","iopub.execute_input":"2023-01-29T17:42:53.803662Z","iopub.status.idle":"2023-01-29T17:42:53.817720Z","shell.execute_reply.started":"2023-01-29T17:42:53.803627Z","shell.execute_reply":"2023-01-29T17:42:53.816891Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Encoder(nn.Module):\n    def __init__(self, config):\n        super().__init__()\n        \n        self.firstLayer = FirstLayer(input_dim=config['feat_dim'],\n                                     hidden_dim=config['feat_dim']*3,\n                                     output_dim=config[\"d_model\"],\n                                     dropout=config[\"dropout_rate\"])\n        self.layers = nn.ModuleList([EncoderLayer(config) for _ in range(config[\"depth\"])])\n\n        self.final_norm = LayerNorm(config)\n        \n        self.net = nn.Sequential(\n            nn.Linear(in_features=config[\"d_model\"], out_features=config[\"head_dim\"]),\n            nn.Sigmoid(),\n        )\n\n    def forward(self, x, mask = None):\n        y = self.firstLayer(x)\n        \n        for layer in self.layers:\n            y = layer(y, mask)\n\n        y = self.final_norm(y)\n        y = self.net(y)\n        \n        y = y.reshape(tuple((y.shape[0], -1, *y.shape[3:])))\n        x = x.reshape(tuple((x.shape[0], -1, *x.shape[3:])))\n        # 'time', 'charge', 'auxiliary', 'x', 'y', 'z'\n#         t = x[:, :, 0]\n#         c = y.squeeze(-1)#torch.mul(y.squeeze(-1), x[:, :, 1])\n#         ox = x[:, :, 3]\n#         oy = x[:, :, 4]\n#         oz = x[:, :, 5]\n        \n#         xtc = torch.mul(torch.mul(ox, t), c)\n#         ytc = torch.mul(torch.mul(oy, t), c)\n#         ztc = torch.mul(torch.mul(oz, t), c)\n#         ttc = torch.mul(torch.mul(t, t), c)\n        \n#         tc = torch.mul(t, c)\n#         xc = torch.mul(ox, c)\n#         yc = torch.mul(oy, c)\n#         zc = torch.mul(oz, c)\n        \n#         t = tc\n#         ox = xc\n#         oy = yc\n#         oz = zc\n#         xt = xtc\n#         yt = ytc\n#         zt = ztc\n#         tt = ttc\n        \n#         t = torch.sum(t, dim=1)\n#         ox = torch.sum(ox, dim=1)\n#         oy = torch.sum(oy, dim=1)\n#         oz = torch.sum(oz, dim=1)\n#         tt = torch.sum(tt, dim=1)\n#         xt = torch.sum(xt, dim=1)\n#         yt = torch.sum(yt, dim=1)\n#         zt = torch.sum(zt, dim=1)\n        \n#         c = torch.sum(c, dim=1)\n        \n#         t = torch.div(t, c)\n#         ox = torch.div(ox, c)\n#         oy = torch.div(oy, c)\n#         oz = torch.div(oz, c)\n#         tt = torch.div(tt, c)\n#         xt = torch.div(xt, c)\n#         yt = torch.div(yt, c)\n#         zt = torch.div(zt, c)\n        \n#         ox = (xt - (ox * t)) / (tt - (t * t)) * -1\n#         oy = (xt - (oy * t)) / (tt - (t * t)) * -1\n#         oz = (xt - (oz * t)) / (tt - (t * t)) * -1\n        \n#         zenith = torch.arccos(oz / torch.sqrt(ox**2 + oy**2 + oz**2))\n#         azimuth = torch.arctan2(oy, ox)\n#         azimuth = torch.add(azimuth, (azimuth<0).float() * 2 * torch.tensor(math.pi))\n\n#         azimuth = azimuth.unsqueeze(1)\n#         zenith = zenith.unsqueeze(1)\n#         target = torch.cat([azimuth, zenith], dim=1)\n        target = get_targets(x, y)\n        return target","metadata":{"execution":{"iopub.status.busy":"2023-01-29T17:42:53.819108Z","iopub.execute_input":"2023-01-29T17:42:53.819760Z","iopub.status.idle":"2023-01-29T17:42:53.831926Z","shell.execute_reply.started":"2023-01-29T17:42:53.819726Z","shell.execute_reply":"2023-01-29T17:42:53.831199Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_lr_scheduler(optimizer, batch_size = 8, last_epoch=-1):\n    lr_start   = 0.00001\n    lr_max     = 0.001\n    lr_min     = 0.0001\n    lr_ramp_ep = 10\n    lr_sus_ep  = 0\n    lr_decay   = 0.98\n    def lrfn(epoch):\n        if epoch < lr_ramp_ep: lr = (lr_max - lr_start) / lr_ramp_ep * epoch + lr_start\n        elif epoch < lr_ramp_ep + lr_sus_ep: lr = lr_max\n        else: lr = (lr_max - lr_min) * lr_decay**(epoch - lr_ramp_ep - lr_sus_ep) + lr_min\n        print(\"Learning rate\", lr)\n        return lr\n    lr_scheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lrfn, last_epoch=last_epoch, verbose=False)\n    return lr_scheduler","metadata":{"execution":{"iopub.status.busy":"2023-01-29T17:42:53.833293Z","iopub.execute_input":"2023-01-29T17:42:53.833907Z","iopub.status.idle":"2023-01-29T17:42:53.845579Z","shell.execute_reply.started":"2023-01-29T17:42:53.833873Z","shell.execute_reply":"2023-01-29T17:42:53.844628Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def angular_dist_score(az_true, zen_true, az_pred, zen_pred):\n    \n    # pre-compute all sine and cosine values\n    sa1 = np.sin(az_true)\n    ca1 = np.cos(az_true)\n    sz1 = np.sin(zen_true)\n    cz1 = np.cos(zen_true)\n    \n    sa2 = np.sin(az_pred)\n    ca2 = np.cos(az_pred)\n    sz2 = np.sin(zen_pred)\n    cz2 = np.cos(zen_pred)\n    \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 =  np.clip(scalar_prod, -1, 1)\n    \n    # convert back to an angle (in radian)\n    return np.average(np.abs(np.arccos(scalar_prod)))","metadata":{"execution":{"iopub.status.busy":"2023-01-29T17:42:53.846916Z","iopub.execute_input":"2023-01-29T17:42:53.847531Z","iopub.status.idle":"2023-01-29T17:42:53.856519Z","shell.execute_reply.started":"2023-01-29T17:42:53.847495Z","shell.execute_reply":"2023-01-29T17:42:53.855540Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def validate(model, val_loader):\n    if not isinstance(model, nn.DataParallel):\n        model = nn.DataParallel(model)\n    \n    model = model.to(device)\n    model.eval()\n\n    loss_list = []\n    pred = []\n    loss_fn = nn.MSELoss()\n    \n    with torch.no_grad():\n        for i, (x, labels) in enumerate(val_loader):            \n            y = model(x)\n            loss = loss_fn(y, labels)\n            \n            y = y.detach().to('cpu').numpy()\n            labels = labels.to('cpu').numpy()\n            for j in range(len(y)):\n                az_pred, zen_pred = y[j]\n                az_true, zen_true = labels[j]\n                pred.append([angular_dist_score(az_true, zen_true, az_pred, zen_pred)])\n            loss_list.append(loss.item())\n\n\n    loss = np.mean(loss_list)\n    pred = np.mean([p for l in pred for p in l])\n    return loss, pred","metadata":{"execution":{"iopub.status.busy":"2023-01-29T17:42:53.857611Z","iopub.execute_input":"2023-01-29T17:42:53.858201Z","iopub.status.idle":"2023-01-29T17:42:53.871432Z","shell.execute_reply.started":"2023-01-29T17:42:53.858074Z","shell.execute_reply":"2023-01-29T17:42:53.870410Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_on_huge_dataset(model, args, checkpoint=None, with_epochs=False, epoch=0):\n    torch.set_grad_enabled(True)\n\n    start_time = time.time()\n\n    if not isinstance(model, nn.DataParallel):\n        model = nn.DataParallel(model)\n\n    model = model.to(device)\n    \n    # Set up the optimizer\n    trainables = [p for p in model.parameters() if p.requires_grad]\n    print('Total parameter number is : {:.3f} million'.format(sum(p.numel() for p in model.parameters()) / 1e6))\n    print('Total trainable parameter number is : {:.3f} million'.format(sum(p.numel() for p in trainables) / 1e6))\n\n    if args[\"optimizer\"] == 'adam':\n        optimizer = torch.optim.Adam(model.parameters(), lr=args[\"lr\"], weight_decay=5e-7, betas=(0.95, 0.999))\n    elif args[\"optimizer\"] == \"adamw\":\n        optimizer = torch.optim.AdamW(model.parameters(), lr=args[\"lr\"], weight_decay=5e-7, amsgrad=True)\n    else:\n        optimizer = torch.optim.SGD(model.parameters(), lr=args[\"lr\"], momentum=0.9, nesterov=True, weight_decay=5e-7)\n    \n    last_epoch = -1\n    last_chunk = 1\n    if checkpoint:\n        model.load_state_dict(checkpoint[\"model_state_dict\"])\n        optimizer.load_state_dict(checkpoint[\"optimizer_state_dict\"])\n        last_epoch = checkpoint[\"epoch\"]\n        last_chunk = checkpoint[\"chunk\"]\n        \n        if with_epochs:\n            last_epoch = last_chunk - 1\n    \n    scheduler = None\n    if args[\"scheduler\"] == \"LambdaLR\":\n        scheduler = get_lr_scheduler(optimizer, batch_size = args[\"batch_size\"] * args[\"NUM_ACCUMULATION_STEPS\"], last_epoch=last_epoch)\n    elif args[\"scheduler\"] == \"cosine\":\n        scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=args[\"cosin_T_max\"], last_epoch=last_epoch)\n    \n    loss_fn = nn.MSELoss()\n\n    model.train()\n    \n    i = 0\n    for batch_id in range(1, 661):\n        batch_id = 1\n        old = None\n        if i!=0:\n            break\n        parquet_file = pq.ParquetFile(f\"/kaggle/input/icecube-data/train_meta_batch_{batch_id}.parquet\")\n        for u, p_batch in enumerate(parquet_file.iter_batches(batch_size=5000)):\n            if old:\n                p_batch = old # test if the model can overfit if not then there is some problem\n            else:\n                old = p_batch\n            if i==10:\n                break\n            i += 1\n            if last_chunk >= i:\n                continue\n            chunk = p_batch.to_pandas()\n            chunk_size = len(chunk)\n            chunk_max_pulses = chunk['pulses'].max()\n            chunk_min_pulses = chunk['pulses'].min()\n            print(f\"chunk size: {chunk_size}, chunk max pulses: {chunk_max_pulses}, chunk min pulses: {chunk_min_pulses}\")\n            train_chunk, test_chunk = train_test_split(chunk, test_size=0.1, random_state=42)\n\n            train_dataset = EventDataset(df=train_chunk)\n            test_dataset = EventDataset(df=test_chunk)\n\n            max_num_pulses = min(chunk_max_pulses, 100)\n            train_collate = get_collate(max_num_pulses)\n            max_tensor_size = 200000 * 5 * 1 # 200000 pulses * 5 features * 1 events(batch_size)\n            batch_size = int(max_tensor_size // (chunk_max_pulses * 5))\n            batch_size = min(batch_size, args[\"batch_size\"])\n            print(f\"chunk batch size: {batch_size}\")\n            train_loader = DataLoader(train_dataset,\n                                      batch_size=batch_size,\n                                      collate_fn=train_collate,\n                                      shuffle=True)\n            test_loader = DataLoader(test_dataset,\n                                     batch_size=batch_size,\n                                     collate_fn=train_collate)\n\n            begin_time = time.time()\n            model.train()\n\n            loss_train = []\n\n            for k, (x, labels) in enumerate(tqdm(train_loader)):\n                y = model(x)\n\n                loss = loss_fn(y, labels)\n\n                loss = loss / args[\"NUM_ACCUMULATION_STEPS\"]\n\n                loss.backward()\n\n                if ((k + 1) % args[\"NUM_ACCUMULATION_STEPS\"] == 0) or (k + 1 == len(train_loader)):\n                    optimizer.step()\n                    optimizer.zero_grad()\n\n                loss_train.append(loss.item())\n\n\n            train_loss = np.mean(loss_train)\n            val_loss, pred = validate(model, test_loader)\n            lr = scheduler.get_last_lr()[0]\n\n            del train_loader, test_loader, train_chunk, test_chunk\n            gc.collect()\n\n            print(f\"chunk: {i}, lr: {lr:.8f}, train loss: {train_loss:.6f}, val loss: {val_loss:.6f}, val angular error : {pred:.6f}\")\n\n            if scheduler:\n                scheduler.step()\n\n            torch.save({\n                'epoch': epoch,\n                'model_state_dict': model.state_dict(),\n                'optimizer_state_dict': optimizer.state_dict(),\n                'chunk': i\n            }, 'model.pth')\n\n            if time.time() - start_time > 60*60*9:\n                break","metadata":{"execution":{"iopub.status.busy":"2023-01-29T17:42:53.872957Z","iopub.execute_input":"2023-01-29T17:42:53.873561Z","iopub.status.idle":"2023-01-29T17:42:53.896917Z","shell.execute_reply.started":"2023-01-29T17:42:53.873526Z","shell.execute_reply":"2023-01-29T17:42:53.895894Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"args = {\n    \"lr\": 0.001,\n    \"lrscheduler_start\": 15,\n    \"lrscheduler_step\": 10,\n    \"lrscheduler_decay\": 0.5,\n    \"warmup\": True,\n    \"optimizer\": [\"adam\", \"adamw\", \"sgd\"][1],\n    \"scheduler\": [\"LambdaLR\"][0],\n    \"batch_size\": 256,\n    \"NUM_ACCUMULATION_STEPS\": 1,\n    \"fold\": 0,\n    \"n_epochs\": 0\n}\n\nd_model = 256\nencoder_config = {\n    \"d_model\": d_model,\n    \"d_ff\": int(d_model*2.5),\n    \"dropout_rate\": 0.1,\n    \"heads\": 12,\n    \"dim_head\": 64,\n    \"depth\": 4,\n    \"norm_eps\": 1e-6,\n    \"feat_dim\": 6,\n    \"head_dim\": 1\n}\nmodel = Encoder(encoder_config)\ntrain_on_huge_dataset(model, args)","metadata":{"execution":{"iopub.status.busy":"2023-01-29T17:42:53.898262Z","iopub.execute_input":"2023-01-29T17:42:53.898803Z"},"trusted":true},"execution_count":null,"outputs":[]}]}