{"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":"import numpy as np\nimport pandas as pd\nimport polars as pl\nimport os\nimport time\nfrom tqdm.auto import tqdm\nimport numba as nb\n\nimport torch\nfrom torch import nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset\nfrom torch.utils.data import DataLoader\nimport math\n\nimport json","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-04-08T09:56:08.153963Z","iopub.execute_input":"2023-04-08T09:56:08.154571Z","iopub.status.idle":"2023-04-08T09:56:11.217470Z","shell.execute_reply.started":"2023-04-08T09:56:08.154537Z","shell.execute_reply":"2023-04-08T09:56:11.215951Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def angular_dist_score(az_true, zen_true, az_pred, zen_pred):\n    \"\"\" https://www.kaggle.com/code/sohier/mean-angular-error \"\"\"\n    if not (np.all(np.isfinite(az_true)) and\n            np.all(np.isfinite(zen_true)) and\n            np.all(np.isfinite(az_pred)) and\n            np.all(np.isfinite(zen_pred))):\n        raise ValueError(\"All arguments must be finite\")\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-04-08T09:56:11.224409Z","iopub.execute_input":"2023-04-08T09:56:11.225075Z","iopub.status.idle":"2023-04-08T09:56:11.239913Z","shell.execute_reply.started":"2023-04-08T09:56:11.225027Z","shell.execute_reply":"2023-04-08T09:56:11.238768Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DATA_DIR = \"/kaggle/input/icecube-neutrinos-in-deep-ice/\"\nPREP_DIR = \"/kaggle/input/icecube-preprocessed-data/\"\nTRAIN_META_DIR = \"/kaggle/input/train-meta-parquet/\"\nMODEL_DIR = \"/kaggle/input/icecube-models/\"\nWORK_DIR = \"/kaggle/working/\"","metadata":{"execution":{"iopub.status.busy":"2023-04-08T09:56:11.245096Z","iopub.execute_input":"2023-04-08T09:56:11.245529Z","iopub.status.idle":"2023-04-08T09:56:11.256119Z","shell.execute_reply.started":"2023-04-08T09:56:11.245488Z","shell.execute_reply":"2023-04-08T09:56:11.254785Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"VALIDATE = False # For checking on validation data\n\nif not VALIDATE:\n    PARQUETS_DIR = os.path.join(DATA_DIR + 'test')\n    BATCH_LIST = list(sorted(os.listdir(PARQUETS_DIR)))\n    metadata = pl.read_parquet(f'{DATA_DIR}/test_meta.parquet')\n    CHECK_PREDICTION = False\nelse:\n    PARQUETS_DIR = os.path.join(DATA_DIR + 'train')\n    vbatches = [655]\n    BATCH_LIST = [f'batch_{vb}.parquet' for vb in vbatches]\n    META_FILES = [f'train_meta_{vb}.parquet' for vb in vbatches]\n    def read_metadata():\n        meta = []\n        for mf, vb in zip(META_FILES, vbatches):\n            bmeta = pl.read_parquet(f'{TRAIN_META_DIR}/{mf}')\n            # Polars isn't going to get much adoption with stupid syntax like this :facepalm:\n            bmeta = bmeta.with_columns(pl.lit(vb).alias('batch_id'))\n            meta.append(bmeta)\n        return pl.concat(meta)\n    metadata = read_metadata()\n    CHECK_PREDICTION = True\n    \nGEOMETRY = os.path.join(PREP_DIR, \"sensor_geometry_with_transparency.csv\")\ngeometry = pl.scan_csv(GEOMETRY).with_columns(\n                [pl.col('sensor_id').cast(pl.Int16)]\n            )\n    \nNUM_BINS = 128\nFEATURE_NAMES = ['time', 'charge', 'auxiliary', 'x', 'y', 'z', 'qe', 'scatter', 'absorp']\nCHARGE_IDX = FEATURE_NAMES.index('charge')\nTIME_IDX = FEATURE_NAMES.index('time')\nAUX_IDX = FEATURE_NAMES.index('auxiliary')\nN_FEATURES = len(FEATURE_NAMES)\nMAX_SEQUENCE_LENGTH = 256\nBATCH_SIZE = 1000\n\nMAX_EVENTS = 200_000 if VALIDATE else 1000000000000 \n\ndevice = 'cuda' if torch.cuda.is_available() else 'cpu'\n\nwith open(os.path.join(PREP_DIR, f'angle_bins_{NUM_BINS}.json')) as fp:\n    bin_data = json.load(fp)\n    \nazimuth_bin_centers = torch.tensor(bin_data['azimuth_bin_centers']).type(torch.float32).to(device)\n# zenith_bin_centers = torch.tensor(bin_data['zenith_bin_centers']).type(torch.float32).to(device)\n\nzenith_centers = np.array(bin_data['zenith_bin_centers'])\nkernel_length = 15\nnum_zenith_padding = (kernel_length - 1)//2\npadded_zenith_bins = np.concatenate([\n        -zenith_centers[num_zenith_padding-1 : : -1],\n        zenith_centers,\n        2 * np.pi - zenith_centers[-1 : -num_zenith_padding-1 : -1],\n])\nzenith_bin_centers = torch.tensor(padded_zenith_bins).type(torch.float32).to(device)\nZENITH_NUM_BINS = len(zenith_bin_centers)\n\n# Model configs\n\n# class v11_m2_s256_ep669:\n#     name = 'v11_m2_s256_ep669.ckpt'\n#     checkpoint_path = 'models/v11_m2_s256_ep669.ckpt'\n#     n_embd = [256]*12\n#     n_heads = [4]*12\n#     bias = False\n#     dropout = 0.0\n#     neck_dropout = 0.0\n#     neck_features = 1536\n#     unwanted_prefix = 'model'\n\nclass v13_m3_s256_ep362:\n    name = 'v13_m3_s256_ep362.ckpt'\n    checkpoint_path = 'models/v13_m3_s256_ep362.ckpt'\n    n_embd = [512]*15\n    n_heads = [8]*15\n    bias = False\n    dropout = 0.0\n    neck_dropout = 0.0\n    neck_features = 3072\n    unwanted_prefix = 'model'\n\n\nMODEL_CONFIGS = [v13_m3_s256_ep362()]","metadata":{"execution":{"iopub.status.busy":"2023-04-08T09:56:23.312043Z","iopub.execute_input":"2023-04-08T09:56:23.312436Z","iopub.status.idle":"2023-04-08T09:56:23.338001Z","shell.execute_reply.started":"2023-04-08T09:56:23.312403Z","shell.execute_reply":"2023-04-08T09:56:23.336907Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@nb.njit\ndef set_seed(value):\n    np.random.seed(value)\n\n@nb.jit( nb.types.Tuple( (nb.float32[:,:,:], nb.int32[:]) )(nb.float64[:,:], nb.int64[:,:]) )\ndef sample_and_pad(data, pulse_indexes):\n    data[:, CHARGE_IDX] = np.log10(data[:, CHARGE_IDX]) / 3.0\n    data[:, AUX_IDX] = data[:, AUX_IDX] - 0.5\n    data_x = np.zeros((len(pulse_indexes), MAX_SEQUENCE_LENGTH, data.shape[-1]), dtype=np.float32)\n    sequence_lengths = np.zeros(len(pulse_indexes), dtype=np.int32)\n    for ii in range(len(pulse_indexes)):\n        event_data = data[pulse_indexes[ii, 0] : pulse_indexes[ii, 1] + 1]\n        if len(event_data) > MAX_SEQUENCE_LENGTH:\n            naux_idx = np.where(event_data[:, AUX_IDX] == -0.5)[0]\n            aux_idx = np.where(event_data[:, AUX_IDX] == 0.5)[0]\n            if len(naux_idx) < MAX_SEQUENCE_LENGTH:\n                max_length_possible = min(MAX_SEQUENCE_LENGTH, len(event_data))\n                num_to_sample = max_length_possible - len(naux_idx)\n                aux_idx_sample = np.random.choice(aux_idx, size=num_to_sample, replace=False)\n                selected_idx = np.concatenate((naux_idx, aux_idx_sample))\n            else:\n                selected_idx = np.random.choice(naux_idx, size=MAX_SEQUENCE_LENGTH, replace=False)\n            selected_idx = np.sort(selected_idx)\n            event_data = event_data[selected_idx]\n        event_data[:, TIME_IDX] = ( event_data[:, TIME_IDX] - event_data[:, TIME_IDX].min() ) / 3e4\n        assert np.all(np.isfinite(event_data))\n        data_x[ii, :len(event_data), :] = event_data\n        sequence_lengths[ii] = len(event_data)                       \n    return data_x, sequence_lengths","metadata":{"execution":{"iopub.status.busy":"2023-04-08T09:56:27.424494Z","iopub.execute_input":"2023-04-08T09:56:27.425106Z","iopub.status.idle":"2023-04-08T09:56:32.867968Z","shell.execute_reply.started":"2023-04-08T09:56:27.425068Z","shell.execute_reply":"2023-04-08T09:56:32.866878Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def preprocess_data(bfile):\n    set_seed(42)\n    print(\"Reading\", bfile)\n    start_time = time.perf_counter()\n    batch_id = int(bfile.split('.')[0].split('_')[-1])\n    batch = pl.scan_parquet(f'{PARQUETS_DIR}/{bfile}')\n    batch = batch.join(geometry, on='sensor_id', how='left')\n    batch_meta = metadata.filter(pl.col('batch_id') == batch_id)\n\n    data = batch.select(FEATURE_NAMES).collect().to_numpy()\n    pulse_indexes = batch_meta.select(['first_pulse_index', 'last_pulse_index']).to_numpy()\n#     pusle_indexes = pulse_indexes[:MAX_EVENTS]\n    print(\"Read and merge\", bfile, \"in\", time.perf_counter() - start_time, \"s\")    \n    \n    start_time = time.perf_counter()\n    data_x, seq_lens = sample_and_pad(data, pulse_indexes)\n    print(\"Processed\", bfile, \"in\", time.perf_counter() - start_time, \"s\")\n\n    if VALIDATE:\n        data_y = batch_meta.select(['azimuth', 'zenith']).to_numpy()\n#         np.save(f'{WORK_DIR}/y_{batch_id}', data_y)\n\n    return data_x, seq_lens","metadata":{"execution":{"iopub.status.busy":"2023-04-08T09:56:32.870279Z","iopub.execute_input":"2023-04-08T09:56:32.870641Z","iopub.status.idle":"2023-04-08T09:56:32.879835Z","shell.execute_reply.started":"2023-04-08T09:56:32.870598Z","shell.execute_reply":"2023-04-08T09:56:32.878461Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if VALIDATE:\n    def check_packing_fraction():\n        dx, dl = preprocess_data(BATCH_LIST[0]);\n        print(dx.shape)\n        # Can get quite low for higher sequence lengths - Around 30% for 256 sequence length\n        print(\"packing fraction without sorting - \", np.mean(dl/dl.max()))\n        lensplit = np.split(dl[np.argsort(dl)], len(dl)//BATCH_SIZE)\n        pfrac = np.mean([np.mean(ls/max(ls)) for ls in lensplit])\n        print(f\"packing fraction with sortting - batch size {BATCH_SIZE} - \", pfrac)\n    \n    check_packing_fraction()","metadata":{"execution":{"iopub.status.busy":"2023-04-08T09:56:32.882068Z","iopub.execute_input":"2023-04-08T09:56:32.882451Z","iopub.status.idle":"2023-04-08T09:56:32.894782Z","shell.execute_reply.started":"2023-04-08T09:56:32.882412Z","shell.execute_reply":"2023-04-08T09:56:32.893764Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class IceCubeDataset(Dataset):\n    def __init__(self, bfile):\n        super().__init__()\n        dx, sl = preprocess_data(bfile)\n        self.x = torch.Tensor(dx)\n        self.l = torch.Tensor(sl)\n        \n        self.sort_idx = np.argsort(sl)\n        self.reverse_sort_idx = np.argsort(self.sort_idx)\n\n    def __len__(self):\n        return len(self.x)\n    \n    def __getitem__(self, index):\n        si = self.sort_idx[index] # Lazy packing, works well for batch size ~1000\n        return self.x[si], self.l[si]","metadata":{"execution":{"iopub.status.busy":"2023-04-08T09:56:32.897716Z","iopub.execute_input":"2023-04-08T09:56:32.898277Z","iopub.status.idle":"2023-04-08T09:56:32.907538Z","shell.execute_reply.started":"2023-04-08T09:56:32.898240Z","shell.execute_reply":"2023-04-08T09:56:32.906535Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class LayerNorm(nn.Module):\n    \"\"\" LayerNorm but with an optional bias. PyTorch doesn't support simply bias=False \"\"\"\n\n    def __init__(self, ndim, bias=False):\n        super().__init__()\n        self.weight = nn.Parameter(torch.ones(ndim))\n        self.bias = nn.Parameter(torch.zeros(ndim)) if bias else None\n\n    def forward(self, input):\n        return F.layer_norm(input, self.weight.shape, self.weight, self.bias, 1e-5)\n\ndef mlp(n_embd, bias=False, dropout=0.0, out_embd=None):\n    out_embd = n_embd if out_embd is None else out_embd\n    return nn.Sequential(\n        nn.Linear(n_embd, 4 * n_embd, bias=bias),\n        nn.GELU(approximate='tanh'),\n        nn.Linear(4 * n_embd, out_embd, bias=bias),\n        nn.Dropout(dropout)\n    )\n\n\nclass SelfAttention(nn.Module):\n    def __init__(self, prev_emdb, n_embd, n_heads, bias=False, dropout=0.0):\n        super().__init__()\n        self.prev_embd = prev_emdb\n        self.n_embd = n_embd\n        self.n_heads = n_heads\n        \n        self.c_attn = nn.Linear(prev_emdb, 3 * n_embd, bias=bias)\n        self.c_proj = nn.Linear(n_embd, n_embd, bias=bias)\n        self.dropout = dropout\n        self.resid_dropout = nn.Dropout(dropout)\n\n    def forward(self, x, attn_mask, cross_features=None):\n        B, T, _ = x.shape\n        C = self.n_embd\n        q, k ,v  = self.c_attn(x).split(self.n_embd, dim=2)\n        k = k.view(B, T, self.n_heads, C // self.n_heads).transpose(1, 2) # (B, nh, T, hs)\n        q = q.view(B, T, self.n_heads, C // self.n_heads).transpose(1, 2) # (B, nh, T, hs)\n        v = v.view(B, T, self.n_heads, C // self.n_heads).transpose(1, 2) # (B, nh, T, hs)\n#         y = F.scaled_dot_product_attention(q, k, v, attn_mask=attn_mask, dropout_p=self.dropout, is_causal=False)\n        y = F._scaled_dot_product_attention(q, k, v, \n                                           attn_mask=attn_mask,\n                                           dropout_p=self.dropout, \n                                           is_causal=False)[0]\n        y = y.transpose(1, 2).contiguous().view(B, T, C)\n        y = self.c_proj(y)\n        y = self.resid_dropout(y)\n        return y\n    \nclass AttentionBlock(nn.Module):\n    def __init__(self, prev_embd, n_embd, n_heads, bias=False, dropout=0.0):\n        super().__init__()\n        self.ln_1 = LayerNorm(prev_embd, bias)\n        assert n_embd % prev_embd == 0, f\"{prev_embd} {n_embd} should be divisble\"\n        self.attn = SelfAttention(prev_embd, n_embd, n_heads, bias, dropout)\n        self.ln_2 = LayerNorm(n_embd, bias)\n        self.mlp = mlp(n_embd, bias, dropout)\n\n    def forward(self, x, attn_mask, cross_features=None):\n        x = x + self.attn(self.ln_1(x), attn_mask, cross_features)\n        x = x + self.mlp(self.ln_2(x))\n        return x\n\nclass AttentionEncoder(nn.Module):\n    def __init__(self, config):\n        super().__init__()\n        dropout = config.dropout\n        bias = config.bias\n        \n        attn_layers = []\n        prev_embd = config.n_embd[0]\n        for n_embd, n_heads in zip(config.n_embd, config.n_heads):\n            attn_layers.append( AttentionBlock(prev_embd, n_embd, n_heads, bias, dropout) )\n            prev_embd = n_embd\n        self.attn = nn.ModuleList(attn_layers)\n    \n    def forward(self, x, attn_mask):\n        out = x\n        for attn_layer in self.attn:\n            out = attn_layer(out, attn_mask)\n        \n        return out\n\n    \nclass SequencePool(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.num_pools = 1\n    \n    def forward(self, x, sequence_lengths, padding_mask):\n        sumf = torch.sum(x * padding_mask.unsqueeze(2), dim=1) # Mask padded tokens\n        meanf = sumf / sequence_lengths.view(-1, 1) # Normalize avg pool values by seq length\n        out = meanf\n        return out\n    \nclass Neck(nn.Module):\n    def __init__(self, in_features, out_features, bias, dropout):\n        super().__init__()\n        self.mlp = nn.Sequential(\n            LayerNorm(in_features, bias=bias),\n            nn.Linear(in_features, 4 * in_features, bias=bias),\n            nn.GELU(approximate='tanh'),\n            nn.Linear(4 * in_features, out_features, bias=bias),\n            nn.Dropout(dropout)\n        )\n        self.n_repeats = out_features // in_features\n\n    def forward(self, x):\n        return x.repeat(1, self.n_repeats) + self.mlp(x)\n\n    \nclass MultiLabelClassifier(nn.Module):    \n    def __init__(self, n_features, max_block_size, num_classes, zenith_num_classes, config):\n        super().__init__()\n\n        self.inp = nn.Linear(n_features, config.n_embd[0])\n        self.drop_inputs = nn.Dropout(config.dropout)\n\n        self.encoder = AttentionEncoder(config)\n        \n        self.pool = SequencePool()\n\n        num_out_features = config.n_embd[-1] * self.pool.num_pools\n\n        self.neck_az = Neck(num_out_features, config.neck_features, config.bias, config.neck_dropout)\n        self.neck_zn = Neck(num_out_features, config.neck_features, config.bias, config.neck_dropout)\n        \n        self.azimuth = nn.Linear(config.neck_features, num_classes)\n        self.zenith = nn.Linear(config.neck_features, zenith_num_classes)\n\n    def get_masks(self, x, l):\n        key_padding_mask = torch.arange(x.shape[1]).view(1, -1).to(l.device) < l.view(-1, 1)\n        attn_mask = (key_padding_mask.unsqueeze(1) == key_padding_mask.unsqueeze(2)).unsqueeze(1)  # (B, 1, T, T)\n        return key_padding_mask, attn_mask\n\n    def forward(self, x):\n        inputs, seq_lengths = x\n        out = self.inp(inputs)\n        out = self.drop_inputs(out)\n        key_padding_mask, attn_mask = self.get_masks(inputs, seq_lengths)\n        out = self.encoder(out, attn_mask)\n        pool = self.pool(out, seq_lengths, key_padding_mask)\n    \n        az_out = self.azimuth(self.neck_az(pool))\n        zn_out = self.zenith(self.neck_zn(pool))\n        return az_out, zn_out","metadata":{"execution":{"iopub.status.busy":"2023-04-08T09:56:32.909288Z","iopub.execute_input":"2023-04-08T09:56:32.909796Z","iopub.status.idle":"2023-04-08T09:56:32.941847Z","shell.execute_reply.started":"2023-04-08T09:56:32.909757Z","shell.execute_reply":"2023-04-08T09:56:32.940660Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def prepare_submission(event_ids, azimuth, zenith, validate_ids=False):\n    if validate_ids:\n        sample_submission = pd.read_parquet(os.path.join(DATA_DIR, 'sample_submission.parquet'))\n        assert np.array_equal(event_ids, sample_submission.event_id.values)\n    zenith_clipped = torch.clip(zenith, 0.0, np.pi)\n    submission_df = pd.DataFrame(\n        {\n            'event_id': event_ids,\n            'azimuth': azimuth.cpu().numpy(),\n            'zenith': zenith_clipped.cpu().numpy(),\n        }\n    ).set_index('event_id')\n    return submission_df","metadata":{"execution":{"iopub.status.busy":"2023-04-08T09:56:32.943523Z","iopub.execute_input":"2023-04-08T09:56:32.944523Z","iopub.status.idle":"2023-04-08T09:56:32.957896Z","shell.execute_reply.started":"2023-04-08T09:56:32.944487Z","shell.execute_reply":"2023-04-08T09:56:32.956946Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_model(model_config):\n    model = MultiLabelClassifier(n_features=N_FEATURES, \n                                max_block_size=MAX_SEQUENCE_LENGTH,\n                                num_classes=NUM_BINS,\n                                zenith_num_classes=ZENITH_NUM_BINS,\n                                config=model_config)\n    checkpoint_path = os.path.join(MODEL_DIR, model_config.name)\n    state_dict = torch.load(checkpoint_path)['state_dict']\n    old_keys = list(state_dict.keys())\n    # print(\"Removing prefix\", model_config.unwanted_prefix)\n    for key in old_keys:\n        if model_config.unwanted_prefix in key:\n            new_key = key.split(model_config.unwanted_prefix)[1][1:]\n            state_dict[new_key] = state_dict.pop(key)\n    model.load_state_dict(state_dict)\n    model.eval()\n    model.to(device)\n    return model","metadata":{"execution":{"iopub.status.busy":"2023-04-08T09:56:32.959361Z","iopub.execute_input":"2023-04-08T09:56:32.959765Z","iopub.status.idle":"2023-04-08T09:56:32.970392Z","shell.execute_reply.started":"2023-04-08T09:56:32.959728Z","shell.execute_reply":"2023-04-08T09:56:32.969161Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"### Classes to Angles\n\n# common\ndef simple_average(pred_mbs):\n    return torch.mean( torch.stack(pred_mbs, dim=2), dim=2)\n\ndef argmax_average(pred_mbs, centers):\n    pred_classes = []\n    for pred in pred_mbs:\n        pred_classes.append(centers[pred.argmax(axis=1)].unsqueeze(0))\n    angles = torch.mean(torch.cat(pred_classes), dim=0)\n    return angles\n\ndef simpleavg_argmax(pred_mbs, centers):\n    pred = simple_average(pred_mbs)\n    angle = centers[pred.argmax(dim=1)]\n    return angle\n\n# azimuth\nclass AzimuthXY:\n    def __init__(self):\n        self.azx = torch.cos(azimuth_bin_centers) \n        self.azy = torch.sin(azimuth_bin_centers)\n\n    def __call__(self, az_softmax):\n        self.azx = self.azx.to(az_softmax.device)\n        self.azy = self.azy.to(az_softmax.device)\n        return az_softmax * self.azx, az_softmax * self.azy\n\naz_xy = AzimuthXY()\n\n\ndef azimuth_vectorsum(azx, azy):\n    azmx, azmy = torch.sum(azx, dim=1), torch.sum(azy, dim=1)\n    azn = torch.sqrt(azmx**2 + azmy**2)\n    az_pred = ( torch.arccos(azmx / azn) * torch.sign(azmy) ) % (np.pi * 2)\n    return az_pred\n\ndef az_simpleavg_vectorsum(pred_mbs):\n    pred = simple_average(pred_mbs)\n    azsf = torch.softmax(pred, dim=1)\n    azx, azy = az_xy(azsf)\n    az = azimuth_vectorsum(azx, azy)\n    return az\n\n# zenith\ndef zn_argmax_average(pred_mbs):\n    return argmax_average(pred_mbs, zenith_bin_centers)\n\n# predictions to angles\ndef ensemble_predictions(az_pred_mbs, zn_pred_mbs):\n    with torch.no_grad():\n        az_pred = az_simpleavg_vectorsum(az_pred_mbs)\n        zn_pred = zn_argmax_average(zn_pred_mbs)\n    return az_pred, zn_pred\n","metadata":{"execution":{"iopub.status.busy":"2023-04-08T09:56:32.972086Z","iopub.execute_input":"2023-04-08T09:56:32.972657Z","iopub.status.idle":"2023-04-08T09:56:32.992378Z","shell.execute_reply.started":"2023-04-08T09:56:32.972621Z","shell.execute_reply":"2023-04-08T09:56:32.991342Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def predict_on_batch(model, dataset, logits_filename=None):\n    az_pred_batch, zn_pred_batch = [], []\n    dataloader = DataLoader(dataset, batch_size=BATCH_SIZE)\n    for x, l in tqdm(dataloader):\n        b_max_len = int(l.max()) # Lazy packing, works well for batch size ~1000\n        with torch.no_grad():\n            azp, znp = model((x[:, :b_max_len].to(device), l.to(device)))\n        az_pred_batch.append(azp)\n        zn_pred_batch.append(znp)\n    \n    az_pred_batch = torch.cat(az_pred_batch, dim=0)[dataset.reverse_sort_idx]\n    zn_pred_batch = torch.cat(zn_pred_batch, dim=0)[dataset.reverse_sort_idx]\n\n    if logits_filename is not None:\n        np.savez_compressed(logits_filename, \n                            azimuth=az_pred_batch.cpu().numpy(), \n                            zenith=zn_pred_batch.cpu().numpy())\n\n    return az_pred_batch, zn_pred_batch","metadata":{"execution":{"iopub.status.busy":"2023-04-08T09:56:32.994281Z","iopub.execute_input":"2023-04-08T09:56:32.994822Z","iopub.status.idle":"2023-04-08T09:56:33.003614Z","shell.execute_reply.started":"2023-04-08T09:56:32.994785Z","shell.execute_reply":"2023-04-08T09:56:33.002451Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def check_score(batch_id, az_pred, zn_pred, text, dataset, convert_to_angle=True):\n    if not VALIDATE:\n        return\n    if convert_to_angle:\n        az, zn = ensemble_predictions([az_pred], [zn_pred])\n    else:\n        az, zn = az_pred, zn_pred\n    az_gt = metadata.filter(pl.col('batch_id') == batch_id).select('azimuth').to_numpy().squeeze()\n    zen_gt = metadata.filter(pl.col('batch_id') == batch_id).select('zenith').to_numpy().squeeze()\n#     az_gt, zen_gt = az_gt[:MAX_EVENTS], zen_gt[:zen_gt]\n    angular_dist = angular_dist_score(az_gt, zen_gt, az.cpu().numpy(), zn.cpu().numpy())\n    print(f\"\\n\\n \\t ang_dist {text}\", angular_dist, \"\\n\")","metadata":{"execution":{"iopub.status.busy":"2023-04-08T09:56:33.006941Z","iopub.execute_input":"2023-04-08T09:56:33.007556Z","iopub.status.idle":"2023-04-08T09:56:33.015523Z","shell.execute_reply.started":"2023-04-08T09:56:33.007326Z","shell.execute_reply":"2023-04-08T09:56:33.014354Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def model_load_check():\n    for model_config in MODEL_CONFIGS:\n        model = load_model(model_config)\n        del model\n        /\nmodel_load_check()","metadata":{"execution":{"iopub.status.busy":"2023-04-08T09:56:15.461812Z","iopub.status.idle":"2023-04-08T09:56:15.462653Z","shell.execute_reply.started":"2023-04-08T09:56:15.462391Z","shell.execute_reply":"2023-04-08T09:56:15.462419Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"az_pred, zn_pred = [], []\nfor bfile in BATCH_LIST:\n    batch_id = int(bfile.split('.')[0].split('_')[-1])\n    dataset = IceCubeDataset(bfile)\n        \n    az_pred_mbs, zn_pred_mbs = [], []\n    for model_config in MODEL_CONFIGS:\n        model = load_model(model_config)\n\n        str_id = f'batch_{batch_id} model_{model_config.name}'\n#         logits_filename = os.path.join(WORK_DIR, f'logits_{str_id}.npz') if VALIDATE else None\n        logits_filename = None\n\n        az_pred_mb, zn_pred_mb = predict_on_batch(model, dataset, logits_filename)\n        az_pred_mbs.append(az_pred_mb)\n        zn_pred_mbs.append(zn_pred_mb)\n        \n        check_score(batch_id, az_pred_mb, zn_pred_mb, text=str_id, dataset=dataset)\n        del model\n\n    az_pred_batch, zn_pred_batch = ensemble_predictions(az_pred_mbs, zn_pred_mbs)\n    az_pred.append(az_pred_batch)\n    zn_pred.append(zn_pred_batch)\n\n    check_score(batch_id, az_pred_batch, zn_pred_batch,\n                text=f'batch_{batch_id}', dataset=dataset, convert_to_angle=False)","metadata":{"execution":{"iopub.status.busy":"2023-04-08T09:56:15.464169Z","iopub.status.idle":"2023-04-08T09:56:15.465037Z","shell.execute_reply.started":"2023-04-08T09:56:15.464779Z","shell.execute_reply":"2023-04-08T09:56:15.464806Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"azimuth = torch.cat(az_pred, axis=0)\nzenith = torch.cat(zn_pred, axis=0)\n\nevent_ids = metadata.select('event_id').to_numpy().squeeze()\nsubmission_df = prepare_submission(event_ids, azimuth, zenith, validate_ids=(not CHECK_PREDICTION))\nsubmission_df.to_csv('submission.csv')\nprint('Saved submission')","metadata":{"execution":{"iopub.status.busy":"2023-04-08T09:56:15.466488Z","iopub.status.idle":"2023-04-08T09:56:15.467428Z","shell.execute_reply.started":"2023-04-08T09:56:15.467244Z","shell.execute_reply":"2023-04-08T09:56:15.467264Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# azimuth, zenith = discrete_to_angle(az_pred, zn_pred, azimuth_bin_centers, zenith_bin_centers)\n# event_ids = metadata.select('event_id').to_numpy().squeeze()\n# submission_df = prepare_submission(event_ids, azimuth, zenith, validate_ids=(not CHECK_PREDICTION))\n# submission_df.to_csv('submission.csv')\n# print('Saved submission')","metadata":{"execution":{"iopub.status.busy":"2023-04-08T09:56:15.469512Z","iopub.status.idle":"2023-04-08T09:56:15.470374Z","shell.execute_reply.started":"2023-04-08T09:56:15.470075Z","shell.execute_reply":"2023-04-08T09:56:15.470103Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CHECK_PREDICTION:\n    submission = pd.read_csv('submission.csv')\n    az_pred, zen_pred = submission.azimuth.values, submission.zenith.values\n    az_gt = metadata['azimuth'].to_numpy()\n    zen_gt = metadata['zenith'].to_numpy()\n    angular_dist = angular_dist_score(az_gt, zen_gt, az_pred, zen_pred)\n    print(\"Angular Distance Score\", angular_dist)","metadata":{"execution":{"iopub.status.busy":"2023-04-08T09:56:15.471861Z","iopub.status.idle":"2023-04-08T09:56:15.472598Z","shell.execute_reply.started":"2023-04-08T09:56:15.472337Z","shell.execute_reply":"2023-04-08T09:56:15.472363Z"},"trusted":true},"execution_count":null,"outputs":[]}]}