{"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":"# HelloWorld: IceCube - Neutrinos in Deep Ice","metadata":{}},{"cell_type":"markdown","source":"## Splitting large metadata file in chunks\n\nThe input detector response files are split into batches, but metadata is stored in a single file of almost 4GB. We want to avoid holding that in memory all the time, so the code below splits this metadata into batches.\n\nIn Jupyter you can write files using the `%%file` magic, and then run them with python using a command `!python FILENAME ARGS` within a cell.","metadata":{}},{"cell_type":"code","source":"%%file partition_meta.py\n\"\"\"\nScript that partitions large metadata file into smaller files\n- separate file per each batch.\n\nUsage:\n    python partition_meta.py SPLIT\nwhere SPLIT is either `train` or `test`\n\"\"\"\n\nfrom pathlib import Path\nimport sys, gc\n\nimport polars as pl\nimport pyarrow.parquet as pq\nfrom tqdm import trange\n\n# base path with input data\nbp = Path(\"/kaggle/input/icecube-neutrinos-in-deep-ice\")\n\ndef iter_through_meta(folder, chunk_size=100):\n    \"\"\"\n    Read the large metadata file in chunks of batches.\n    \n    Parameters\n    ----------\n    folder : str\n        Which part of the data to process. Should be either \"train\" or \"test\".\n\n    Returns\n    -------\n    generator\n        A generator of chunks of batches in pyarrow.Table format.\n    \"\"\"\n    assert folder in [\"train\", \"test\"], \"Argument `folder` should either be 'train' or 'test'\"\n    src_file = bp / f\"{folder}_meta.parquet\"\n    all_batch_ids = pl.read_parquet(src_file, columns=[\"batch_id\"]).unique().sort(\"batch_id\")\n    for chunk_first in trange(0, len(all_batch_ids), chunk_size):\n        selection = all_batch_ids[chunk_first: chunk_first + chunk_size][\"batch_id\"]\n        first, last = selection.min(), selection.max()\n        yield pq.read_table(src_file, filters=[\n            (\"batch_id\", \">=\", first),\n            (\"batch_id\", \"<=\", last),\n        ])\n        gc.collect()\n\ndef write_meta_batches(meta, folder):\n    \"\"\"\n    Take a chunk of metadata info and write it to disk with separate file per batch.\n\n    Parameters\n    ----------\n    meta : pyarrow.Table\n        A table with a chunk of batches with metadata.\n    folder : str\n        Which to write into. Should be either \"train\" or \"test\".\n    \"\"\"\n    folder.mkdir(exist_ok=True)\n    pq.write_to_dataset(meta, root_path=folder, partition_cols=['batch_id'], flavor='spark')\n\ndef main(folder):\n    \"\"\"\n    Read metatdata in chunks and save separate files per batch.\n    \n    Parameters\n    ----------\n    folder : str\n        Which part of the data to process. Should be either \"train\" or \"test\".\n    \"\"\"\n    print(f\"Working on {folder}\")\n    for table in iter_through_meta(folder):\n        write_meta_batches(table, Path(folder))\n    print(f\"Done ({folder})\")\n    \nif __name__ == \"__main__\":\n    try:\n        _, folder = sys.argv\n    except ValueError:\n        print('Usage: \"python partition_meta.py SPLIT\", where SPLIT = \"train\" or \"test\"')\n        exit(0)\n    main(folder)","metadata":{"execution":{"iopub.status.busy":"2023-04-03T17:18:05.766174Z","iopub.execute_input":"2023-04-03T17:18:05.766626Z","iopub.status.idle":"2023-04-03T17:18:05.798505Z","shell.execute_reply.started":"2023-04-03T17:18:05.766595Z","shell.execute_reply":"2023-04-03T17:18:05.797428Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Run the partitioning code (only if needed, which we check by the existence of the resulting folders):","metadata":{}},{"cell_type":"code","source":"from pathlib import Path\n\nif not Path(\"train\").exists():\n    !python partition_meta.py train\nif not Path(\"test\").exists():\n    !python partition_meta.py test","metadata":{"execution":{"iopub.status.busy":"2023-04-03T17:18:05.800286Z","iopub.execute_input":"2023-04-03T17:18:05.800702Z","iopub.status.idle":"2023-04-03T17:20:48.089300Z","shell.execute_reply.started":"2023-04-03T17:18:05.800666Z","shell.execute_reply":"2023-04-03T17:20:48.087998Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Kaggle has weird indicators for used memory: they seem to partially be including the cache. A more detailed info may be obtained lie this:","metadata":{}},{"cell_type":"code","source":"# Check how many resources we actually take\n!top -bn1 -o '%MEM'\n\nimport os\nprint(f\"(our pid: {os.getpid()})\")","metadata":{"execution":{"iopub.status.busy":"2023-04-03T17:20:48.091428Z","iopub.execute_input":"2023-04-03T17:20:48.091763Z","iopub.status.idle":"2023-04-03T17:20:49.278820Z","shell.execute_reply.started":"2023-04-03T17:20:48.091729Z","shell.execute_reply":"2023-04-03T17:20:49.277518Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Finally, let's check that the number of metadata files we created is the same as the number of pulse files:","metadata":{}},{"cell_type":"code","source":"# base path\nbp = Path(\"/kaggle/input/icecube-neutrinos-in-deep-ice\")\n\ndef check_num_batches(split):\n    num_actual = len(list(Path(split).glob(\"batch_id=*\")))\n    num_expected = len(list((bp / split).glob(\"batch_*.parquet\")))\n\n    if num_actual != num_expected:\n        print(\n            f\"WARNING!!! Found {num_actual} batch files when expected {num_expected} for \"\n            f'split \"{split}\". Check that partitioning code ran ok.'\n        )\n    else:\n        print(f\"Check ok ({split})\")\n\nfor split in [\"train\", \"test\"]:\n    check_num_batches(split)","metadata":{"execution":{"iopub.status.busy":"2023-04-03T17:20:49.282525Z","iopub.execute_input":"2023-04-03T17:20:49.284550Z","iopub.status.idle":"2023-04-03T17:20:49.465684Z","shell.execute_reply.started":"2023-04-03T17:20:49.284509Z","shell.execute_reply":"2023-04-03T17:20:49.464719Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Imports\n\n(making all the necessary imports in a single cell)","metadata":{}},{"cell_type":"code","source":"from pathlib import Path\nimport gc\nimport os\nfrom functools import lru_cache\n\nimport polars as pl\nimport pandas as pd\nimport numpy as np\nfrom IPython.display import display, HTML\nimport matplotlib.pyplot as plt\nfrom plotly.offline import iplot, init_notebook_mode\ninit_notebook_mode(connected=True) # https://stackoverflow.com/questions/67419817/uncaught-error-script-error-for-plotly-http-requirejs-org-docs-errors-html\nimport plotly.graph_objects as go\nimport plotly.express as px\nimport torch\nimport pytorch_lightning as ptl\nfrom tensorboard.backend.event_processing import event_accumulator","metadata":{"execution":{"iopub.status.busy":"2023-04-03T17:20:49.468365Z","iopub.execute_input":"2023-04-03T17:20:49.469387Z","iopub.status.idle":"2023-04-03T17:21:05.557908Z","shell.execute_reply.started":"2023-04-03T17:20:49.469349Z","shell.execute_reply":"2023-04-03T17:21:05.556786Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Having a look at our data","metadata":{}},{"cell_type":"markdown","source":"There's a special file describing the geometry of the detector. Here we read it and convert its `sensor_id` column to the same dtype as we'll see in the pulse files.","metadata":{}},{"cell_type":"code","source":"sensors_df = pl.read_csv(bp / \"sensor_geometry.csv\").with_columns(pl.col(\"sensor_id\").cast(pl.Int16))\n\nprint(\"sesnsors shape\", sensors_df.shape)","metadata":{"execution":{"iopub.status.busy":"2023-04-03T17:21:05.559873Z","iopub.execute_input":"2023-04-03T17:21:05.560255Z","iopub.status.idle":"2023-04-03T17:21:05.607626Z","shell.execute_reply.started":"2023-04-03T17:21:05.560209Z","shell.execute_reply":"2023-04-03T17:21:05.606675Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Here's a function to read a single batch of pulses:","metadata":{}},{"cell_type":"code","source":"def load_batch(folder, b_id):\n    \"\"\"\n    Load a single batch of pulses.\n    \n    Parameters\n    ----------\n    folder : str\n        Which part of the data to process. Should be either \"train\" or \"test\".\n    b_id : int\n        Index of the batch (as in the `batch_id` column).\n\n    Returns\n    -------\n    polars.DataFrame\n        Data frame with the loaded pulses.\n    polars.DataFrame\n        Data frame with corresponding metadata.\n    \"\"\"\n    pulses = pl.read_parquet(bp / folder / f\"batch_{b_id}.parquet\")\n    pulses = pulses.join(sensors_df, on=\"sensor_id\").drop(\"sensor_id\")\n\n    meta = pl.read_parquet(f\"{folder}/batch_id={b_id}/*.parquet\")\n\n    return pulses, meta\n\n# Let's create an in-memory cache version of this function to avoid repeatedly\n# loading the same batch when we examine various events from it. We'll delete this function\n# later to free the memory taken by cache.\nload_batch_c = lru_cache(maxsize=1)(load_batch)","metadata":{"execution":{"iopub.status.busy":"2023-04-03T17:21:05.609519Z","iopub.execute_input":"2023-04-03T17:21:05.609866Z","iopub.status.idle":"2023-04-03T17:21:05.617130Z","shell.execute_reply.started":"2023-04-03T17:21:05.609825Z","shell.execute_reply":"2023-04-03T17:21:05.616094Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Let's collect all the train and test `batch_id`s for easier positional indexing later on:","metadata":{}},{"cell_type":"code","source":"def get_batch_ids_from_folder(folder):\n    fnames = [x.stem for x in (bp / folder).glob(\"*.parquet\")]\n    assert all(x.startswith(\"batch_\") for x in fnames)\n    return np.sort(np.array([int(x[6:]) for x in fnames]))\n\nTRAIN_BIDS = get_batch_ids_from_folder(\"train\")\nTEST_BIDS = get_batch_ids_from_folder(\"test\")","metadata":{"execution":{"iopub.status.busy":"2023-04-03T17:21:05.618986Z","iopub.execute_input":"2023-04-03T17:21:05.619734Z","iopub.status.idle":"2023-04-03T17:21:05.631902Z","shell.execute_reply.started":"2023-04-03T17:21:05.619696Z","shell.execute_reply":"2023-04-03T17:21:05.630823Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Now we want to see what our data tables look like. We're doing it inside of a function to avoid creating globals taking up memory later on.","metadata":{}},{"cell_type":"code","source":"def preview_data():\n    b_pulses_train, b_meta_train = load_batch_c(\"train\", TRAIN_BIDS[0])\n    b_pulses_test, b_meta_test = load_batch(\"test\", TEST_BIDS[0])\n\n    display(HTML(\"<H1>Train meta:\"))\n    display(b_meta_train.head())\n\n    display(HTML(\"<hr><H1>Test meta:\"))\n    display(b_meta_test.head())\n\n    display(HTML(\"<hr><H1>Sensors:\"))\n    display(sensors_df.head())\n\n    display(HTML(\"<hr><H1>Pulses batch train:\"))\n    display(b_pulses_train.head())\n\n    display(HTML(\"<hr><H1>Pulses batch test:\"))\n    display(b_pulses_test.head())\n\n    display(HTML(\"<hr><H1>Sample submission:\"))\n    display(pl.read_parquet(bp / \"sample_submission.parquet\").head())\n\npreview_data()","metadata":{"execution":{"iopub.status.busy":"2023-04-03T17:21:05.637900Z","iopub.execute_input":"2023-04-03T17:21:05.638571Z","iopub.status.idle":"2023-04-03T17:21:10.662015Z","shell.execute_reply.started":"2023-04-03T17:21:05.638543Z","shell.execute_reply":"2023-04-03T17:21:10.660879Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Here's a helper function to load data only for a single event.","metadata":{}},{"cell_type":"code","source":"def load_single_event(folder, i_batch, i_evt, ignore_aux=True):\n    \"\"\"\n    Load a single event (using positional indexing).\n\n    Parameters\n    ----------\n    folder : str\n        Which part of the data to process. Should be either \"train\" or \"test\".\n    i_batch : int\n        Positional index of the batch. Should be in a range from 0 (inclusive)\n        to the number of batches in `folder` (exclusive).\n    i_evt : int\n        Positional index of the event within the batch.\n    ignore_aux : bool\n        Whether to exclude auxiliary pulses.\n\n    Returns\n    -------\n    int\n        The id of the read event (as in `event_id` column).\n    polars.DataFrame\n        Table with pulses.\n    Tuple[float, float] | None\n        Azimuth and zenith (if `folder` is 'train') or `None` (if `folder` is 'test').\n    \"\"\"\n    bids = dict(train=TRAIN_BIDS, test=TEST_BIDS)[folder]\n    pulses, meta = load_batch_c(folder, bids[i_batch])\n\n    meta_event = meta[i_evt]\n    (event_id,) = meta_event[\"event_id\"]\n\n    event_pulses = pulses.filter(pl.col(\"event_id\") == event_id)\n    if ignore_aux:\n        event_pulses = event_pulses.filter(~pl.col(\"auxiliary\")).drop(\"auxiliary\")\n\n    target = None\n    if \"azimuth\" in meta.columns:\n        (azimuth,) = meta_event[\"azimuth\"]\n        (zenith,) = meta_event[\"zenith\"]\n        target = (azimuth, zenith)\n\n    return event_id, event_pulses, target","metadata":{"execution":{"iopub.status.busy":"2023-04-03T17:21:10.663954Z","iopub.execute_input":"2023-04-03T17:21:10.664361Z","iopub.status.idle":"2023-04-03T17:21:10.672347Z","shell.execute_reply.started":"2023-04-03T17:21:10.664320Z","shell.execute_reply":"2023-04-03T17:21:10.670923Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"And a function to make a 3D plot of an event. The input arguments are aligned with the outputs from the `load_single_event` function above.","metadata":{}},{"cell_type":"code","source":"def plot_event(event_id, event_pulses, y=None):\n    fig = px.scatter_3d(\n        x=event_pulses[\"x\"],\n        y=event_pulses[\"y\"],\n        z=event_pulses[\"z\"],\n        color=event_pulses[\"time\"],\n        size=event_pulses[\"charge\"],\n    )\n\n    if y is not None:\n        (azimuth, zenith) = y\n        xyz = event_pulses[['x', 'y', 'z']]\n        xyz_mean = xyz.mean(axis=0)\n        r = 1200\n\n        (x0,) = xyz_mean['x']\n        (y0,) = xyz_mean['y']\n        (z0,) = xyz_mean['z']\n        (x1,) = xyz_mean['x'] + r * np.cos(azimuth) * np.sin(zenith)\n        (y1,) = xyz_mean['y'] + r * np.sin(azimuth) * np.sin(zenith)\n        (z1,) = xyz_mean['z'] + r * np.cos(zenith)\n\n        fig.add_trace(\n            go.Scatter3d(\n                x=np.linspace(x0, x1, 100),\n                y=np.linspace(y0, y1, 100),\n                z=np.linspace(z0, z1, 100),\n                marker=go.scatter3d.Marker(size=0.001)\n            )\n        )\n    fig.update_layout(\n        scene=dict(\n            xaxis=dict(range=[-600,600],),\n            yaxis=dict(range=[-600,600],),\n            zaxis=dict(range=[-600,600],),\n        ),\n        scene_aspectmode='cube',\n        title=dict(text=f\"evt #{event_id}\")\n    )\n\n    return fig","metadata":{"execution":{"iopub.status.busy":"2023-04-03T17:21:10.673984Z","iopub.execute_input":"2023-04-03T17:21:10.674826Z","iopub.status.idle":"2023-04-03T17:21:10.695532Z","shell.execute_reply.started":"2023-04-03T17:21:10.674760Z","shell.execute_reply":"2023-04-03T17:21:10.694440Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"OK, now let's finally make some plots! Here's an event:","metadata":{}},{"cell_type":"code","source":"iplot(plot_event(*load_single_event(\"train\", 0, 18)))","metadata":{"execution":{"iopub.status.busy":"2023-04-03T17:21:10.697505Z","iopub.execute_input":"2023-04-03T17:21:10.698286Z","iopub.status.idle":"2023-04-03T17:21:12.306191Z","shell.execute_reply.started":"2023-04-03T17:21:10.698248Z","shell.execute_reply":"2023-04-03T17:21:12.305117Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Not all events have the direction this nicely aligned with the pulses trail. E.g., check out this one:","metadata":{}},{"cell_type":"code","source":"iplot(plot_event(*load_single_event(\"train\", 0, 21)))","metadata":{"execution":{"iopub.status.busy":"2023-04-03T17:21:12.307615Z","iopub.execute_input":"2023-04-03T17:21:12.308644Z","iopub.status.idle":"2023-04-03T17:21:12.434187Z","shell.execute_reply.started":"2023-04-03T17:21:12.308602Z","shell.execute_reply":"2023-04-03T17:21:12.433113Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"To the best of my understanding, here the trail is likely left by a background particle (like a cosmic muon). If so, the neutrino direction information (if any) for such events should be found in the auxiliary pulses. Here's how the organizers of the competition describe the `auxiliary` flag:\n> `auxiliary` (`bool`): If `True`, the pulse was not fully digitized, is of lower quality, and was more likely to originate from noise. If `False`, then this pulse was contributed to the trigger decision and the pulse was fully digitized.\n\nWe can check what the same event looks like if we plot both regular and aux pulses:","metadata":{}},{"cell_type":"code","source":"iplot(plot_event(*load_single_event(\"train\", 0, 21, ignore_aux=False)))","metadata":{"execution":{"iopub.status.busy":"2023-04-03T17:21:12.436046Z","iopub.execute_input":"2023-04-03T17:21:12.436753Z","iopub.status.idle":"2023-04-03T17:21:12.556136Z","shell.execute_reply.started":"2023-04-03T17:21:12.436714Z","shell.execute_reply":"2023-04-03T17:21:12.555115Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data preprocessing\n\nIn this example notebook, we are building a very simple convolutional neural network to predict the neutrino direction. What we'll actually do is we'll discretize our pulses into a fixed shape matrix. Moreover, to avoid large sparse 4D matrices (X, Y, Z, time), we'll calculate time-coordinate 2D projections (TX, TY, TZ). We'll also shift all the pulse timings by the median time in an event. **The particular binning chosen here is very likely not optimal (as well as this binning approach in general).**","metadata":{}},{"cell_type":"code","source":"def normalize_time(df):\n    g = df[[\"event_id\", \"time\"]].groupby(\"event_id\")\n    times = g.quantile(0.5).rename(dict(time=\"t_mid\")).with_columns(pl.col(\"t_mid\").cast(pl.Int64))\n    return df.join(times, on=\"event_id\").with_columns((pl.col(\"time\") - pl.col(\"t_mid\")).alias(\"time_norm\"))\n\nNBINS_T = 10\nNBINS_S = 10\n\ndef preprocess_batch(batch_df):\n    # Time bins:\n    tbins = np.linspace(-5000, 10000, NBINS_T + 1)\n    tbin_width = tbins[1] - tbins[0]\n    tbin_centers = (tbins[:-1] + tbins[1:]) / 2\n\n    # Spacial bins:\n    sbins = np.linspace(-600, 600, NBINS_S + 1)\n    sbin_width = sbins[1] - sbins[0]\n    sbin_centers = (sbins[:-1] + sbins[1:]) / 2\n\n    # Normalize time and calculate a gaussian-like contribution of each pulse to each bin.\n    # We group all the data by `event_id` and then aggregate each group with the code below.\n    return normalize_time(batch_df).groupby(\"event_id\").agg([\n        ((\n            np.exp(-(                                           #\n                ((pl.col(\"time_norm\") - tmid) / tbin_width)**2  #  <== Calculate the gaussian-smoothed\n                + ((pl.col(ax) - smid) / sbin_width)**2         #  <== contribution of each pulse to each bin\n            ))\n        # Then weight by charge, calculate per bin sum and rename (we'll get a separate column per each bin):\n        ) * pl.col(\"charge\")).sum().alias(f\"t_{int(tmid)}_{ax}_{int(smid)}\".replace('-', \"m\"))\n        for ax in [\"x\", \"y\", \"z\"]                          #  <== Iterate over axes\n        for tmid in tbin_centers for smid in sbin_centers  #  <== Iterate over bins\n    ])\n\n# Here's a function to check what a preprocessed event looks like.\n# Note that we are plotting log(1 + amplitude).\ndef plot_example(i_batch=0, i_event=4):\n    b_pulses, _ = load_batch_c(\"train\", TRAIN_BIDS[i_batch])\n    e_id = b_pulses[\"event_id\"].unique().sort()[i_event]\n    prep_batch = preprocess_batch(b_pulses.filter(pl.col(\"event_id\") == e_id))\n    print(\"Preprocessed batch shape:\", prep_batch.shape)\n    img = prep_batch.drop(\"event_id\").to_numpy().reshape(3, NBINS_T, NBINS_S)  # 3 images (TX, TY, TZ) of NBINS_T by NBINS_S\n\n    plt.figure(figsize=(14, 4))\n    plt.imshow(np.concatenate([np.pad(np.log1p(x), 1, constant_values=np.nan) for x in img], axis=-1))\n    plt.colorbar();\n\nplot_example()","metadata":{"execution":{"iopub.status.busy":"2023-04-03T17:21:12.557878Z","iopub.execute_input":"2023-04-03T17:21:12.558571Z","iopub.status.idle":"2023-04-03T17:21:13.375294Z","shell.execute_reply.started":"2023-04-03T17:21:12.558533Z","shell.execute_reply":"2023-04-03T17:21:13.374115Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"In the example above we ran preprocessing over just a single event. Running on an entire batch may take a considerable amount of time. In order to avoid repeatedly preprocessing our data, we'll make use of on-disk cache with `joblib.Memory`. This approach does not work very healthily when caching a function defined in an interactive Jupyter session (it may result in repeated calculations). Therefore, we'll use the `%%file` magic again to create and import such a function:","metadata":{}},{"cell_type":"code","source":"%%file preprocessing_cache.py\n\nimport numpy as np\nimport polars as pl\nfrom joblib import Memory\nmemory = Memory(\"cache/\")\n\n@memory.cache(ignore=[\"load_batch\", \"preprocess_batch\"])\ndef load_and_preprocess_batch(\n    folder, batch_id, load_batch, preprocess_batch, num_subsample=None\n):\n    \"\"\"\n    Load and preprocess a batch, possibly taking only a subsample of a batch.\n    \n    Parameters\n    ----------\n    folder : str\n        Which part of the data to process. Should be either \"train\" or \"test\".\n    batch_id : int\n        Index of the batch (as in the `batch_id` column).\n    load_batch : Callable\n        Function loading a batch (to pass the interactively defined `load_batch` function).\n    preprocess_batch : Callable\n        Function preprocessing a batch (to pass the interactively defined `preprocess_batch`\n        function).\n    num_subsample : int | None\n        If provided, only preprocess a random subsample of this size from the batch.\n\n    Returns\n    -------\n    polars.DataFrame\n        Preprocessed pulses.\n    polars.DataFrame\n        Corresponding meta info.\n    \"\"\"\n    batch_pulses, batch_meta = load_batch(folder, batch_id)\n    if num_subsample is not None:\n        ids = batch_pulses[\"event_id\"].unique()\n        ids = np.random.choice(ids, num_subsample, replace=False)\n        batch_pulses = batch_pulses.filter(pl.col(\"event_id\").is_in(ids.tolist()))\n        batch_meta = batch_meta.filter(pl.col(\"event_id\").is_in(ids.tolist()))\n    batch_pulses = preprocess_batch(batch_pulses).sort(by=\"event_id\")\n    batch_meta = batch_meta.sort(by=\"event_id\")\n\n    assert batch_pulses[\"event_id\"].series_equal(batch_meta[\"event_id\"])\n    return batch_pulses, batch_meta","metadata":{"execution":{"iopub.status.busy":"2023-04-03T17:21:13.376896Z","iopub.execute_input":"2023-04-03T17:21:13.377271Z","iopub.status.idle":"2023-04-03T17:21:13.385667Z","shell.execute_reply.started":"2023-04-03T17:21:13.377230Z","shell.execute_reply":"2023-04-03T17:21:13.384309Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Here's a simple wrapper around the above defined function:","metadata":{}},{"cell_type":"code","source":"import preprocessing_cache\n\ndef load_and_preprocess_batch(folder, i_batch, num_subsample=None):\n    bids = dict(train=TRAIN_BIDS, test=TEST_BIDS)[folder]\n    return preprocessing_cache.load_and_preprocess_batch(\n        folder, bids[i_batch], load_batch, preprocess_batch, num_subsample\n    )","metadata":{"execution":{"iopub.status.busy":"2023-04-03T17:21:13.387236Z","iopub.execute_input":"2023-04-03T17:21:13.388141Z","iopub.status.idle":"2023-04-03T17:21:13.434942Z","shell.execute_reply.started":"2023-04-03T17:21:13.388099Z","shell.execute_reply":"2023-04-03T17:21:13.434015Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Since we're done with examining single events, we can delete the in-memory cached function to free some RAM:","metadata":{}},{"cell_type":"code","source":"# Free the memory taken by the cached version of the loader function as we don't need it any more.\n\ndel load_batch_c\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-04-03T17:21:13.436373Z","iopub.execute_input":"2023-04-03T17:21:13.436697Z","iopub.status.idle":"2023-04-03T17:21:13.810507Z","shell.execute_reply.started":"2023-04-03T17:21:13.436662Z","shell.execute_reply":"2023-04-03T17:21:13.809172Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Defining the dataset, model and running a training loop","metadata":{}},{"cell_type":"markdown","source":"First, we define our dataset as a subclass of `torch.utils.data.IterableDataset`:","metadata":{}},{"cell_type":"code","source":"class Dataset(torch.utils.data.IterableDataset):\n    def __init__(self, folder, file_ids, batch_size, drop_each_last=True, num_subsample=None, with_evt_id=False):\n        \"\"\"\n        Parameters\n        ----------\n        folder : str\n            Which part of the data to process. Should be either \"train\" or \"test\".\n        file_ids : Sequence[int]\n            A collection of batch-file indices (positional) to use.\n        batch_size : int\n            Size of the output batches (not to be confused with the batches in which the pulses\n            are provided by the competition).\n        drop_each_last : bool\n            Whether to avoid batches of size smaller than `batch_size`\n        num_subsample : int | None\n            Take this many elements from each file (`None` = take all).\n        with_evt_id : bool\n            Output event ids (useful for making and submitting the prediction).\n        \"\"\"\n        super().__init__()\n        self.folder = folder\n        self.file_ids = file_ids\n        self.batch_size = batch_size\n        self.drop_each_last = drop_each_last\n        self.num_subsample = num_subsample\n        self.with_evt_id = with_evt_id\n\n    def __iter__(self):\n        for i in np.random.choice(len(self.file_ids), len(self.file_ids), replace=False):\n            gc.collect()\n            batch_x, batch_y = load_and_preprocess_batch(\n                self.folder, self.file_ids[i], num_subsample=self.num_subsample\n            )\n            if \"azimuth\" in batch_y.columns:\n                batch_y = batch_y[[\"azimuth\", \"zenith\"]]\n            else:\n                batch_y = None\n\n            evt_ids = batch_x[\"event_id\"]\n            batch_x = batch_x.drop(\"event_id\")\n\n            if self.num_subsample is not None:\n                assert len(batch_x) == self.num_subsample\n\n            for i_evt in range(0, len(batch_x), self.batch_size):\n                minibatch_x = batch_x[i_evt: i_evt + self.batch_size]\n                if self.drop_each_last and len(minibatch_x) < self.batch_size:\n                    continue\n\n                minibatch_x = torch.from_numpy(\n                    minibatch_x.to_numpy().astype(np.float32).reshape(-1, 3, NBINS_T, NBINS_S)\n                )\n                minibatch_y = None if batch_y is None else (\n                    torch.from_numpy(\n                        batch_y[i_evt: i_evt + self.batch_size].to_numpy().astype(np.float32)\n                    )\n                )\n                if self.with_evt_id:\n                    yield minibatch_x, minibatch_y, evt_ids[i_evt: i_evt + self.batch_size]\n                else:\n                    yield minibatch_x, minibatch_y","metadata":{"execution":{"iopub.status.busy":"2023-04-03T17:21:13.812453Z","iopub.execute_input":"2023-04-03T17:21:13.812847Z","iopub.status.idle":"2023-04-03T17:21:13.827669Z","shell.execute_reply.started":"2023-04-03T17:21:13.812804Z","shell.execute_reply":"2023-04-03T17:21:13.826504Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Now the model.\n\n**The architecture** is quite simple: process each of the 3 projections by a downsampling convolutional network, then proecss the result with a fully-connected network. The output will be a vector with 3 components from which we'll determine the angles later.\n\n**The loss** will be calculated in the output vector representation. I.e., we'll take the target angles, convert them to a directional unit vector and calculate the MSE between this vector and our prediction.\n\n**For monitoring**, we'll overwite a `on_epoch_end` lightning method, in which we'll create/update a plot of loss values.","metadata":{}},{"cell_type":"code","source":"assert NBINS_T == 10 and NBINS_S == 10, \"Our model expects 10x10 representation\"\n\nclass ConvPredictor(torch.nn.Module):\n    def __init__(self, activation=torch.nn.ELU()):\n        super().__init__()\n        (self.model_tx, self.model_ty, self.model_tz) = [\n            torch.nn.Sequential(\n                torch.nn.Conv2d(1, 32, 3), activation, # 1x10 -> 32x8\n                torch.nn.Conv2d(32, 64, 3), activation, # -> 64x6\n                torch.nn.Conv2d(64, 128, 3), activation, # -> 128x4\n                torch.nn.Conv2d(128, 256, 3), activation, # -> 256x2\n            ) for _ in range(3)\n        ]\n        self.head = torch.nn.Sequential(\n            torch.nn.Linear(3 * 256 * 2 * 2, 32), activation,\n            torch.nn.Linear(32, 3)\n        )\n\n    def forward(self, x):\n        # x shape is (BATCH, 3, 10, 10)\n        x = torch.log(1.0 + x)\n        tx, ty, tz = x[:, 0: 1], x[:, 1: 2], x[:, 2: 3]\n        tx = self.model_tx(tx).view(x.shape[0], 256 * 4)\n        ty = self.model_ty(ty).view(x.shape[0], 256 * 4)\n        tz = self.model_tz(tz).view(x.shape[0], 256 * 4)\n        pred = self.head(torch.cat([tx, ty, tz], axis=1))\n        return pred\n\nclass LitModel(ptl.LightningModule):\n    def __init__(self, model):\n        super().__init__()\n        self.model = model\n\n    def calculate_loss(self, batch):\n        (X,), (Y,) = batch\n        (azimuth, zenith) = Y.T\n        vx = torch.cos(azimuth) * torch.sin(zenith)\n        vy = torch.sin(azimuth) * torch.sin(zenith)\n        vz = torch.cos(zenith)\n        v = torch.stack([vx, vy, vz], axis=1)\n\n        pred_v = self.model(X)\n        return torch.nn.functional.mse_loss(pred_v, v)\n\n    def training_step(self, batch, batch_idx):\n        loss = self.calculate_loss(batch)\n        self.log(\"train_loss\", loss, prog_bar=True)\n        return loss\n\n    def validation_step(self, batch, batch_idx):\n        loss = self.calculate_loss(batch)\n        self.log(\"val_loss\", loss, prog_bar=True)\n\n    def configure_optimizers(self):\n        opt = torch.optim.Adam(self.parameters(), lr=1e-4)\n        return opt\n\n    def on_validation_end(self):\n        super().on_validation_end()\n\n        self.logger.experiment.flush()\n        if not hasattr(self, \"tb_reader\"):\n            self.tb_reader = event_accumulator.EventAccumulator(self.logger.log_dir)\n            self.hdisplay = display(\"\", display_id=True)\n        ea = self.tb_reader\n        ea.Reload()\n\n        try:\n            train_values = [(x.step, x.value) for x in ea.Scalars(\"train_loss\")]\n            val_values = [(x.step, x.value) for x in ea.Scalars(\"val_loss\")]\n\n            if hasattr(self, \"last_fig\"):\n                plt.close(self.last_fig)\n\n            self.last_fig = plt.figure()\n            plt.plot(*zip(*train_values), label=\"Train\")\n            plt.plot(*zip(*val_values), label=\"Validation\");\n            plt.xlabel(\"Step\")\n            plt.ylabel(\"Loss\")\n            plt.legend()\n            self.hdisplay.update(self.last_fig)\n        except KeyError:\n            pass\n\n\nmodel = ConvPredictor()\nlit_model = LitModel(model)","metadata":{"execution":{"iopub.status.busy":"2023-04-03T17:21:13.829334Z","iopub.execute_input":"2023-04-03T17:21:13.829867Z","iopub.status.idle":"2023-04-03T17:21:13.880079Z","shell.execute_reply.started":"2023-04-03T17:21:13.829831Z","shell.execute_reply":"2023-04-03T17:21:13.879118Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"OK, it's time to train our model. For this example, we'll use a small subset of data, and only run training for a few epochs.","metadata":{}},{"cell_type":"code","source":"gc.collect()\n\nnp.random.seed(42)\nuse_n_files = 10\nchosen_file_ids = np.random.choice(len(TRAIN_BIDS), use_n_files + 1, replace=False)\ntrain_file_ids = chosen_file_ids[:-1]\nval_file_ids = chosen_file_ids[-1:]\nprint(\"Train file ids:\", train_file_ids)\nprint(\"Validation file ids:\", val_file_ids)\n\ntrain_dataset = Dataset(\n    folder=\"train\",\n    file_ids=train_file_ids,\n    batch_size=32,\n    drop_each_last=True,\n    num_subsample=50000,\n)\nval_dataset = Dataset(\n    folder=\"train\",\n    file_ids=val_file_ids,\n    batch_size=500,\n    num_subsample=5000,\n)\ntrain_dataloader = torch.utils.data.DataLoader(train_dataset)\nval_dataloader = torch.utils.data.DataLoader(val_dataset)\n\ntrainer = ptl.Trainer(\n    max_epochs=25,\n    callbacks=ptl.callbacks.ModelCheckpoint(monitor=\"val_loss\"),\n    accelerator='gpu', devices=1,\n)\ntrainer.fit(model=lit_model, train_dataloaders=train_dataloader, val_dataloaders=val_dataloader)","metadata":{"execution":{"iopub.status.busy":"2023-04-03T17:21:13.881555Z","iopub.execute_input":"2023-04-03T17:21:13.881985Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(trainer.checkpoint_callback.best_model_score.detach().cpu().numpy())\nprint(trainer.checkpoint_callback.best_model_path)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Let's restore the checkpoint with best score:","metadata":{}},{"cell_type":"code","source":"checkpoint = torch.load(trainer.checkpoint_callback.best_model_path)\nlit_model.load_state_dict(checkpoint[\"state_dict\"])","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer.validate(lit_model, val_dataloader)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Predicting on test and saving submission file","metadata":{}},{"cell_type":"code","source":"del trainer, train_dataset, val_dataset, train_dataloader, val_dataloader\ngc.collect()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def angles2vec(azimuth, zenith):\n    x = torch.cos(azimuth) * torch.sin(zenith)\n    y = torch.sin(azimuth) * torch.sin(zenith)\n    z = torch.cos(zenith)\n    return torch.stack([x, y, z], axis=1)\n\ndef vec2angles(vec):\n    norm = ((vec**2).sum(axis=1)**0.5)[:, None]\n    vec = vec / norm\n    zenith = torch.acos(vec[:, 2])\n    sin_zenith = torch.sin(zenith)\n    cos_azimuth = vec[:, 0] / sin_zenith\n    sin_azimuth = vec[:, 1] / sin_zenith\n    azimuth = torch.atan2(sin_azimuth, cos_azimuth)\n    return azimuth, zenith","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_ds = Dataset(\"test\", list(range(len(TEST_BIDS))), batch_size=1000, drop_each_last=False, with_evt_id=True)\n\ndef torch2numpy(x):\n    if isinstance(x, torch.Tensor):\n        x = x.cpu().numpy()\n    return x\nwith torch.no_grad():\n    predictions_df = pd.concat([\n        pd.DataFrame({\n            name: torch2numpy(value) for name, value in zip(\n                [\"event_id\", \"azimuth\", \"zenith\"],\n                (eids.to_numpy(),) + vec2angles(model(X))\n            )\n        }) for X, _, eids in test_ds\n    ]).sort_values(by=\"event_id\")\n\npredictions_df.to_csv(\"submission.csv\", index=False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!head submission.csv","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}