{"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":"# IceCube: submit ensemble prediction\n\nThis notebook creates predictions for GNN and trasformer models, wich saves submission_<model_name>.csv files.\n\nEnsemble module aggregates this files and predict final results and save in to submission.csv","metadata":{}},{"cell_type":"code","source":"GLOBAL_TEST_MODE    = False # False for submission, True for validation\nGLOBAL_TEST_BATCHES = 5     # number of batches for validation","metadata":{"execution":{"iopub.status.busy":"2023-04-18T17:30:19.620224Z","iopub.execute_input":"2023-04-18T17:30:19.620787Z","iopub.status.idle":"2023-04-18T17:30:19.627003Z","shell.execute_reply.started":"2023-04-18T17:30:19.620728Z","shell.execute_reply":"2023-04-18T17:30:19.625782Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\n\nimport sys\n\n!rm -rf ./submission*\n!rm  -r nnet\n!rm  -r software\n!mkdir software/\n!mkdir software/graphnet\n\n!scp -r /kaggle/input/icecube-weights/graphnet-main-20230216/* software/graphnet\n!scp -r /kaggle/input/graphnet-and-dependencies/software .\n\n# Install dependencies\n!pip install -q /kaggle/working/software/dependencies/torch-1.11.0+cu115-cp37-cp37m-linux_x86_64.whl\n!pip install -q /kaggle/working/software/dependencies/torch_cluster-1.6.0-cp37-cp37m-linux_x86_64.whl\n!pip install -q /kaggle/working/software/dependencies/torch_scatter-2.0.9-cp37-cp37m-linux_x86_64.whl\n!pip install -q /kaggle/working/software/dependencies/torch_sparse-0.6.13-cp37-cp37m-linux_x86_64.whl\n!pip install -q /kaggle/working/software/dependencies/torch_geometric-2.0.4.tar.gz\n\n# Install GraphNeT\n!cd software/graphnet;pip install --no-index --find-links=\"/kaggle/working/software/dependencies\" -e .[torch]\n\n# Install polars\n!pip install ../input/polars01516/typing_extensions-4.4.0-py3-none-any.whl\n!pip install ../input/polars01516/polars-0.15.16-cp37-abi3-manylinux_2_17_x86_64.manylinux2014_x86_64.whl\n\n# Append to PATH\nsys.path.append('/kaggle/working/software/graphnet/src')","metadata":{"execution":{"iopub.status.busy":"2023-04-18T17:19:18.952017Z","iopub.execute_input":"2023-04-18T17:19:18.952505Z","iopub.status.idle":"2023-04-18T17:19:20.035751Z","shell.execute_reply.started":"2023-04-18T17:19:18.952404Z","shell.execute_reply":"2023-04-18T17:19:20.034137Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Att","metadata":{}},{"cell_type":"markdown","source":"## V2","metadata":{}},{"cell_type":"code","source":"CHECK_BEFORE_SUBMIT = GLOBAL_TEST_MODE  # check before submit\n\n#===============================================================================\n# if CHECK_BEFORE_SUBMIT = True:\nFIRST_BATCH_ID   = 1                     # number of first batch for train and validation\nNUM_BATCHES      = GLOBAL_TEST_BATCHES   # number of batches for train\nBATCHES_IN_PACK  = 5                     # number of batches in pack\n#===============================================================================\n\nTEST_MODE           = CHECK_BEFORE_SUBMIT  # testing (submition) or train and validate  !!!!!!!\nDROP_AUX            = False                # выкидывать aux==True\nDOMS_AGG            = False                # агрегировать по сенсорам (время первого)\nDATA_KIND           = f\"{DROP_AUX:d}{DOMS_AGG:d}\"\n\nimport os, gc, sys, time, datetime, math, random,  psutil, copy\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom   pathlib   import Path        \nfrom   tqdm.auto import tqdm\nimport pandas as pd\nimport pyarrow, pyarrow.parquet as pq     # read by chanks\n\nimport torch\nfrom   torch import nn\n\n#===============================================================================\n\nclass CFG:   \n    T_max =  1024        # max tokens (pulses in event)\n    AF    =  24*2         # 24 if DROP_AUX else 24*2   # event features (set in Dataset.next_load) ??        \n    F     =   3         # 2 if DROP_AUX else 3      # 3:(aux, q, t) or 2:(q, t) next add 4:(x,y,z,core,a)\n    SE    =   4         # 10 if DROP_AUX else 9      # if is_cat == 0 equal E \n    E     =  64         # embedding \n    \n    L     = 13          # transformer layers\n    Eh    = 256         # embedding after rnn (if not or mean\n    Hin   =  32         # скрытый слой в генераторе фич     (4+8=12) -> 32 -> 64\n    Hout  = 128         # neurons in hidden layer of output MLP  64 -> 128 -> 3    \n\n    data     = DATA_KIND\n    arch     = \"\"\n    nums     = [180, 60] # число разбиений азимутального и зенитного углоыв\n\n    frozen   = False\n    loss     = 'k2'\n    ka_reg   = 0\n\n    device     = 'cuda' \n    params     = 0\n    samples    = 0\n    steps      = 0\n    last       = 0\n    score      = 0\n\n    def plt(end=\"\\n\"):\n        return \"\".join([f\" {k:7s}:{v}{end}\" for k,v in CFG.__dict__.items() if not k.startswith(\"__\") and k not in [\"get\",\"plt\", \"V\",  \"params\", \"device\", \"lr\",  \"best\", \"is_sqr\",\"is_cat\",\"is_pos\",\"is_rnn\",\"is_agg\",\"is_emb\",\"is_abs\",\"is_reg\",\"last\"] ])\n\n    def get(end=\", \"):\n        return \"\".join([f\"{k}:{v}{end}\" for k,v in CFG.__dict__.items() if not k.startswith(\"__\") and k not in [\"get\", \"par\"] ])\n\nprint(CFG.get())","metadata":{"execution":{"iopub.status.busy":"2023-04-18T16:44:10.379371Z","iopub.execute_input":"2023-04-18T16:44:10.379702Z","iopub.status.idle":"2023-04-18T16:44:11.803894Z","shell.execute_reply.started":"2023-04-18T16:44:10.379670Z","shell.execute_reply":"2023-04-18T16:44:11.802926Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#===============================================================================\nPATH      = Path(\"/kaggle/input/icecube-neutrinos-in-deep-ice\")  # path to dataset\nPATH_PHYS = Path(\"/kaggle/input/icecube-phys\")\nfiles_trn = [item for item in (PATH  / \"train\").glob('*')]  # all train files\nprint(f\"{len(files_trn):3d} train files\")\n#===============================================================================\ndef info(text, pref=\"\", end=\"\\n\"):\n    \"\"\" Information about the progress of calculations (time and memory) \"\"\"\n    gc.collect()\n    ram, t = psutil.virtual_memory().used / 1024**3,  time.time()    \n    print(f\"{pref}{(t-info.beg)/60:5.1f}m[{t-info.last:+5.1f}s] {ram:6.3f}Gb > {text}\",end=end)\n    info.last = time.time(); \ninfo.beg = info.last = time.time()\n\n#-------------------------------------------------------------------------------\n\ndef get_sensors():\n    \"\"\" Get sensor positions \"\"\"            \n    df = pd.read_csv(PATH / \"sensor_geometry.csv\")      \n    df['line_id'] = df.sensor_id // 60 + 1                 # string id\n    df['core']    = (df.line_id > 78).astype(np.float32)   # sensor from DeepCore\n    df.x = ( df.x * 1e-3 ).astype(np.float32)              # distances in kilometers\n    df.y = ( df.y * 1e-3 ).astype(np.float32)\n    df.z = ( df.z * 1e-3 ).astype(np.float32)    \n    \n    from scipy.interpolate import interp1d                 # add absorption\n    phys = pd.read_csv(PATH_PHYS / \"scattering_and_absorption.csv\")\n    phys.z = (phys.z * 1e-3).astype(np.float32)\n    phys.a = (phys.a * 1e-2).astype(np.float32)\n    interp = interp1d(phys.z, phys.a)\n    df['a'] = interp(df.z)\n\n    df['r'] = np.sqrt(df.x**2 + df.y**2)\n\n    return df[['sensor_id', 'line_id', 'core', 'x', 'y', 'z', 'a', 'r']]\n\n#-------------------------------------------------------------------------------\n\ndef get_target_angles(batch_id=1):\n    \"\"\" Get target angles for batch with batch_id \"\"\"    \n    assert batch_id > 0 and  batch_id < 661, \"Wrong batch_id\"    \n    file = pq.ParquetFile(PATH / \"train_meta.parquet\")    \n    for b in file.iter_batches(batch_size=200_000, columns=['event_id','batch_id','azimuth','zenith']):    \n        batch_df = b.to_pandas()\n        if batch_df.batch_id[0] == batch_id:      \n            batch_df.event_id= batch_df.event_id.astype(np.int64)      \n            batch_df.azimuth = batch_df.azimuth.astype(np.float32)\n            batch_df.zenith  = batch_df.zenith.astype(np.float32)                        \n            return batch_df[ ['event_id','azimuth','zenith'] ]\n\n#-------------------------------------------------------------------------------\n\ndef prepare_batch(df, verbose=True, drop_aux = DROP_AUX, doms_agg = DOMS_AGG):\n    \"\"\" Preparing a loaded batch, shifting and normalizing times \"\"\"    \n    df['event_id'] = df.index.astype(np.int64)\n    df = df.reset_index(drop=True)  # sensor_id, t, charge, aux, event_id    \n    df.rename(columns={\"time\": \"t\", \"auxiliary\": \"aux\", 'charge': 'q'}, inplace=True)\n    df.q = df.q.astype(np.float32)\n\n    if drop_aux:\n        df = df[ ~df.aux ]\n    \n    if doms_agg:\n        df = df.groupby(['event_id', 'sensor_id']).agg(\n            aux = ( 'aux', \"mean\"),\n            q   = ( 'q',   \"sum\"),\n            t   = ( 't',   \"min\"),            \n        )\n        df = df.reset_index()\n    \n    if verbose: info(f\"load_batch: loaded  {df.shape}\")\n        \n    times = df.groupby('event_id').agg( t_min = ('t', 'min') )\n    df = df.merge(times, left_on='event_id', right_index=True, how='left')\n    df.t = (( df.t - df.t_min ) * 0.299792458e-3 ).astype(np.float32)             \n    \n    if verbose: info(\"load_batch: shift_times\")    \n\n    return df[['event_id', 'sensor_id', 'aux', 'q', 't' ]]\n        \n#-------------------------------------------------------------------------------\n\ndef cut_pulses(df, max_pulses = 128, verbose=True):\n    \"\"\" Выкидываем последние и ненадёжные пульсы в событии если их больше max_pulses \"\"\"\n    tot = len(df)\n    if DROP_AUX:\n        df = df.sort_values(['event_id','t'])      \n    else:\n        df = df.sort_values(['event_id','aux','t'])   # do you need aux???\n\n    df = df.reset_index(drop=True)    \n\n    df = df.groupby('event_id').head(max_pulses)  # cut pulses by event\n    df = df.reset_index()                         # sorted by time later!\n\n    if not DROP_AUX:\n        df = df.sort_values(['event_id','t'])        \n        df = df.reset_index(drop=True)\n\n    if verbose: info(f\"cut_pulses (max={max_pulses}): removed {100*(tot-len(df))/tot:.2f}%\")\n    return df\n\n#-------------------------------------------------------------------------------\n\ndef angles2vector(df):\n    \"\"\" Add unit vector components from (azimuth,zenith) to the DataFrame df \"\"\"\n    df['nx'] = np.sin(df.zenith) * np.cos(df.azimuth)\n    df['ny'] = np.sin(df.zenith) * np.sin(df.azimuth)\n    df['nz'] = np.cos(df.zenith) \n    return df\n\n#-------------------------------------------------------------------------------\n\ndef delta_angle(n1, n2, eps=1e-8):\n    \"\"\" Вычислить углы между двумя векторами: n1,n2: (B,3) return: (B,) \"\"\"\n    n1 = n1 / (np.linalg.norm(n1, axis=1, keepdims=True) + eps)\n    n2 = n2 / (np.linalg.norm(n2, axis=1, keepdims=True) + eps)\n    cos = (n1*n2).sum(axis=1).clip(-1,1)\n    return np.arccos( cos )\n\n#-------------------------------------------------------------------------------\n\ndef get_event_features(df, target_df, suf = \"\", aux=True):    \n    \"\"\" Aggregated features characterizing the entire event \"\"\"\n\n    df['xt'] = df.x*df.t;  df['yt'] = df.y*df.t; df['zt'] = df.z*df.t;  df['tt'] = df.t**2;       \n    if aux:\n        for col in df.columns:\n            if col not in ['event_id', 'sensor_id', 'line_id']:\n                df[col] = df[col] * (1-df.aux)\n\n    df = df.groupby('event_id').agg(          # это по всем пульсам с любым aux     \n        tot      = ('t',        'count'),            \n        t_med    = ('t',        'median'),  \n        t        = ('t',        'mean'),  \n        x        = ('x',        'mean'),  \n        y        = ('y',        'mean'),  \n        z        = ('z',        'mean'),  \n        stdT     = ('t',        'std'),\n        stdX     = ('x',        'std'),\n        stdY     = ('y',        'std'),\n        stdZ     = ('z',        'std'),\n        xt       = ('xt',       'mean'),\n        yt       = ('yt',       'mean'),\n        zt       = ('zt',       'mean'),\n        tt       = ('tt',       'mean'),\n        q        = ('q',        'mean' ),\n        q_min    = ('q',        'min' ),\n        q_max    = ('q',        'max' ),\n        q_med    = ('q',        'median' ),\n        aux      = ('aux',      'mean' ),       \n        core     = ('core',     'mean' ),        \n        lines    = ('line_id',  'nunique' ),\n        doms     = ('sensor_id','nunique' ),                \n    )\n    df = df.reset_index()    \n    \n    df.aux   = df.aux  .astype(np.float32)\n    df.lines = df.lines.astype(np.float32)\n    df.doms  = df.doms .astype(np.float32)    \n    df.stdT  = df.stdT .astype(np.float32)    \n    df.stdX  = df.stdX .astype(np.float32)    \n    df.stdY  = df.stdY .astype(np.float32)    \n    df.stdZ  = df.stdZ .astype(np.float32)    \n\n    df['p_lines'] = np.log10(df.tot / df.lines).astype(np.float32)\n    df['p_doms']  = np.log10(df.tot / df.doms ).astype(np.float32)    \n\n    df.q         = np.log(1+df.q)\n    df.q_med     = np.log(1+df.q_med)\n    df.q_min     = np.log(1+df.q_min)\n    df.q_max     = np.log(1+df.q_max)\n    df.lines     = np.log10(df.lines)  / 10\n    df.doms      = np.log10(df.doms)   / 10\n    df['pulses'] =(np.log10(df.tot)    / 10).astype(np.float32)\n    \n    df = df.fillna(0.0)   # if exclude aux is possible problems for std?\n\n    if len(suf):   # добавлям суффикс в название колонки\n        cols = [ col + suf for col in df.columns]\n        df.columns = cols\n\n    return df\n\n#-------------------------------------------------------------------------------\n\ndef get_pulse_features(df):\n    \"\"\" \"\"\"\n    df = df.drop(columns=['line_id'])   # !!!! (embedding ?)\n\n    df.q    = np.log(1+df.q)    \n    for col in df.columns:\n        if col not in ['sensor_id', 'event_id', 'line_id', 'tot']:\n            df[col] = df[col].astype(np.float32)\n\n    return df    \n\n#===============================================================================\n#                     Create dataset for train and validation\n#===============================================================================\n\ndef get_files(batch_ids):\n    if CHECK_BEFORE_SUBMIT:      # imitate submision\n        files = [PATH  / \"train\" / f\"batch_{batch_id}.parquet\"  for batch_id in batch_ids]            \n    else:\n        files     = [item for item in (PATH  / \"test\" ).glob('*')]  # all test  files\n        batch_ids = [ i for i in range(661, len(files)+661)]        # :)        \n\n    return files, batch_ids\n\n\n#-------------------------------------------------------------------------------\n\ndef append_dict(data, T, df, agg_df):\n    \"\"\" \n    df:     event_id sensor_id aux\tq\tt  tot\n    agg_df: event_id, nx, ny, nz, tot, t_aver, ...., ux, uy, uz, qx, qy, qz\n    \"\"\"\n    assert len(df) % T == 0,  f\"wait len(df) = T*B, got len={len(df)}, T={T}\"    \n    B, F = len(df) // T, df.shape[-1] - 3 # drop: event_id, sensor_id, tot\n    ID   = agg_df[['event_id']].to_numpy()\n    Y    = agg_df[['nx','ny','nz']].to_numpy()\n    AGG  = agg_df.iloc[:, 5:].to_numpy()\n    SENS = df.sensor_id.to_numpy().reshape(B,T)\n    # (B*T, F) -> (B, T, F) -> (B, F, T) -> (B, F*T)\n    FEAT= df.iloc[:, 2: -1].to_numpy().reshape(B,T,F)      # drop tot !\n    #FEAT = np.transpose(FEAT, axes=(0,2,1))  # (B,F,T)\n    #FEAT = FEAT.reshape(B, F*T)              \n    assert len(ID)==len(Y) and len(ID)==len(AGG) and len(ID)==len(SENS) and len(ID)==len(FEAT), \\\n           f\"{ID.shape}, {Y.shape}, {AGG.shape}, {SENS.shape} {FEAT.shape} from df={df.shape} agg_df={agg_df.shape} (T={T},F={F})\"\n\n    if T in data:    # ID, Y, AGG, SENS, FEAT \n        v = data[T]\n        v[0] = torch.vstack((v[0], torch.tensor(ID,   dtype=torch.long)    ))\n        v[1] = torch.vstack((v[1], torch.tensor(SENS, dtype=torch.long)    ))\n        v[2] = torch.cat   ((v[2], torch.tensor(FEAT, dtype=torch.float32) ), dim=0 )\n        v[3] = torch.vstack((v[3], torch.tensor(AGG,  dtype=torch.float32) ))\n        v[4] = torch.vstack((v[4], torch.tensor(Y,    dtype=torch.float32) ))\n        \n    else:       \n        data[T] = [torch.tensor(ID,   dtype=torch.long   ),\n                   torch.tensor(SENS, dtype=torch.long   ),\n                   torch.tensor(FEAT, dtype=torch.float32),\n                   torch.tensor(AGG,  dtype=torch.float32),\n                   torch.tensor(Y,    dtype=torch.float32) ]                  \n\n    CFG.F = F\n#-------------------------------------------------------------------------------\n\ndef create_dataset(batch_ids, sensors_df, verbose):\n    \"\"\" \"\"\"\n    files, batch_ids = get_files(batch_ids)\n    data, events_df  = {}, pd.DataFrame({'event_id': []})\n    for i, (batch_id, fname) in tqdm(enumerate(zip(batch_ids, files))):         \n        info(f\"******  batch_id: {batch_id:3d}\")\n        df = pd.read_parquet(fname)            \n        if TEST_MODE: # In order not to lose at the end of any events (the organizers do not merge them)            \n            events_df = events_df.append(  pd.DataFrame({'event_id': df.index.unique() }) )            \n\n        df = prepare_batch(df, doms_agg=DOMS_AGG)\n        df = cut_pulses(df, max_pulses=CFG.T_max)        \n\n        df = df.merge(sensors_df, left_on=\"sensor_id\", right_on=\"sensor_id\", how=\"left\")\n        df = df[['event_id', 'line_id', 'sensor_id', 'core', 'aux', 'q', 't', 'x', 'y', 'z']]\n        df = df.fillna(0.0)\n        \n        info(f\"merged batch with sensors {df.shape}\")    \n\n        if CHECK_BEFORE_SUBMIT:\n            target_df = get_target_angles(batch_id=batch_id)\n            target_df = angles2vector(target_df).drop(columns=['azimuth','zenith'])            \n        else:\n            target_df = pd.DataFrame({'event_id': df.event_id.unique() })\n            target_df['nx']=0; target_df['ny']=0;  target_df['nz']=1;\n        info(\"loaded target angles\")            \n\n        agg_df = get_event_features(df, target_df, suf=\"\", aux=False)\n        if not DROP_AUX:            \n            if DOMS_AGG:  # при агригации некоторые сенсоры имеют нецелый aux (умножаем на 1-него)!\n                agg2_df = get_event_features(df, target_df, suf=\"_aux\", aux=True)            \n            else:         \n                agg2_df = get_event_features(df[ ~df.aux ].copy(), target_df, suf=\"_aux\", aux=False)            \n            agg_df = agg_df.merge(agg2_df, left_on='event_id',  right_on='event_id_aux', how='left')\n            agg_df = agg_df.drop(columns = ['event_id_aux','tot_aux'] )\n\n        agg_df = target_df.merge(agg_df, left_on=\"event_id\", right_on=\"event_id\", how=\"left\")        \n        info('get_event_features done')\n\n        df = get_pulse_features(df)        \n\n        if DROP_AUX:\n            df = df[['event_id', 'sensor_id',        'q', 't']]  #   'core' , 'x', 'y', 'z'\n        else:\n            df = df[['event_id', 'sensor_id', 'aux', 'q', 't']]  #   'core', 'x', 'y', 'z'\n\n        info('get_pulse_features done')\n\n        df = df.merge(agg_df[['event_id', 'tot']], left_on='event_id', right_on='event_id', how='left')\n\n        if verbose and i == 0: show_stats(df, agg_df)\n\n        tots = df.tot.unique()\n        info(f\"count pulses:  {tots.mean():.0f} [{tots.min()} ... {tots.max()}]\")                                    \n        for n in tqdm(tots): \n            # first pulse will be last (for RNN)\n            d1 = df    [df.    tot == n].sort_values(['event_id','t'], ascending=[True,False])           \n            d2 = agg_df[agg_df.tot == n].sort_values(['event_id'])           \n            append_dict(data, n, d1, d2)\n        cols_df, cols_agg_df = df.columns, agg_df.columns\n        del df, agg_df\n    info(\"collected data for dataset\")                \n        \n    return data, events_df.reset_index(drop=True), cols_df, cols_agg_df\n\n#===============================================================================\n#                                Diagnostic\n#===============================================================================\n\ndef show_stats(df, agg_df):\n    \"\"\" \"\"\"\n    pd.set_option('display.float_format', lambda x: '%.2f' % x)\n    display(df.head(5))\n    display(df.describe(percentiles=[]).transpose())    \n    display(df.info())\n    display(agg_df.head(2))\n    display(agg_df.describe(percentiles=[]).transpose())                                        \n    display(agg_df.info())\n\n#-------------------------------------------------------------------------------\n\ndef plot_metric(err, prefix=\"\", bins = 200):    \n    \"\"\" Build a histogram of errors; calculate the statistics and the share w of 'bad examples' \"\"\"\n    plt.figure(figsize=(6,4), facecolor ='w') \n    plt.axes().set_facecolor(\"ivory\"); plt.autoscale(tight=True)\n    p,_,_ = plt.hist(err, bins=bins, range=(0,np.pi), fc=\"lightblue\", density=True, alpha=0.5)\n    w = 2*p[len(p)//2: ].sum()*np.pi/bins    \n    x = np.linspace(0,np.pi,bins)\n    plt.plot(x, w * 0.5*np.sin(x),   c=\"darkred\")\n    plt.plot(x, p-w * 0.5*np.sin(x), c=\"darkblue\")\n    plt.title(f\"{prefix}mean={np.mean(err):.4f}, median={np.median(err):.3f}, w={w:.3f}\")    \n    plt.ylabel(\"Density\"); plt.xlabel(r\"$\\Delta \\Psi$ (rad)\"); plt.grid()\n    plt.show()\n\n#-------------------------------------------------------------------------------        \n\ndef dataset_stat(data):    \n    X = np.array(sorted(list(data.keys())))    \n    Y = np.array([len(data[n][0]) for n in X])    \n    plt.title(f\"Samples by len (T): min={Y.min()} [{ X[np.argmin(Y)] }]  max={Y.max()} [{ X[np.argmax(Y)] }] tot={np.sum(Y)}\")\n    plt.scatter(X,Y, s=3); plt.grid(); plt.show()        \n\n    T = X[np.argmin(Y)]\n    F = CFG.n_feat\n    FEAT_df = pd.DataFrame(data[T][-1].numpy())        # E, Y, AGG, SENS, FEAT     \n    FEAT_df.columns = [ f\"f{f}_t{t}\"  for f in range(F) for t in range(T) ]    \n    display(FEAT_df.head(5))\n    del FEAT_df; gc.collect()\n\n    cnt, mem = 0, 0\n    for k,v in data.items():\n        mem += sum( [v[i].numel()*v[i].element_size() for i in range(5)] )\n        cnt += len(v[0])\n    info(f\"dataset samples: {cnt}, memory: {mem/1024**3:.3f} Gb\")","metadata":{"execution":{"iopub.status.busy":"2023-04-18T16:44:11.807501Z","iopub.execute_input":"2023-04-18T16:44:11.807993Z","iopub.status.idle":"2023-04-18T16:44:12.137891Z","shell.execute_reply.started":"2023-04-18T16:44:11.807966Z","shell.execute_reply":"2023-04-18T16:44:12.136823Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def create_dataset_on_disk():\n    sensors_df = get_sensors()\n    CFG.V = len(sensors_df)   \n    info(f\"loaded sensors pos: tot={len(sensors_df)}\")\n    display(sensors_df.head(3))\n    doms = torch.tensor(sensors_df[['x','y','z','core','a','r']].astype(np.float32).to_numpy())\n    torch.save( { 'cols': ['x','y','z','core','a','r'], 'data': doms }, f\"doms.pt\")\n    del doms\n    \n    FILES = []\n    if TEST_MODE:\n        for i,batch_id in tqdm(enumerate(range(FIRST_BATCH_ID, FIRST_BATCH_ID+NUM_BATCHES,  BATCHES_IN_PACK)), total=NUM_BATCHES//BATCHES_IN_PACK):\n            pack_id = batch_id // BATCHES_IN_PACK + 1\n            data, _, cols_df, cols_agg_df = create_dataset(range(batch_id, batch_id + BATCHES_IN_PACK), sensors_df, i==0)\n            torch.save({'cols_df':      cols_df, \n                        'cols_agg_df':  cols_agg_df, \n                        'data': data },   f\"pack_{pack_id:02d}.pt\")        \n            del data; gc.collect()\n            info(f\"created pack {pack_id:2d}\")           \n            FILES.append(f'pack_{pack_id:02d}.pt')\n    else:\n        data, _, cols_df, cols_agg_df = create_dataset(None, sensors_df, verbose=False)\n        torch.save({'cols_df':      cols_df, \n                    'cols_agg_df':  cols_agg_df, \n                    'data': data }, f\"pack_01.pt\")        \n        FILES.append(f'pack_01.pt')                 # TODO !!!!\n        \n        del data; gc.collect()\n        info(f\"created pack\")   \n    return FILES\n    \n","metadata":{"execution":{"iopub.status.busy":"2023-04-18T16:44:12.139250Z","iopub.execute_input":"2023-04-18T16:44:12.140324Z","iopub.status.idle":"2023-04-18T16:44:12.150549Z","shell.execute_reply.started":"2023-04-18T16:44:12.140285Z","shell.execute_reply":"2023-04-18T16:44:12.149468Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Create Dataset V1","metadata":{}},{"cell_type":"code","source":"class Dataset:\n    \"\"\"    \n    \"\"\"\n    def __init__(self, files, batch_size, shuffle, device, drop, aux=False):\n        \"\"\" \n        data - словарь, его ключи - длина последовательности в пульсах, а значения список:\n                * EVENT_ID:  (B,1)    - event id                                \n                * SENSOR_ID: (B,T)    - sensor_id for each pulse (token)\n                * FEAT:      (B,T,F)  - F features for each pulse (token)     \n                * AGG:       (B, FA)  - aggregate features\n                * Y:         (B,3)    - (nx, ny, nz) - target direction (unit vector)\n        \"\"\"        \n        self.files      = files\n        self.shuffle    = shuffle\n        self.batch_size = batch_size\n        self.batch_max  = batch_size\n        self.T_max      = CFG.T_max\n        self.file_id    = 0\n        self.device     = device\n        self.drop       = drop\n        self.epoch      = 0      # число эпох после загрузки нового пака\n        self.data       = {}\n        self.is_T1      = 0\n        self.is_T2     = 1e8\n        \n        self.aux        = aux\n        \n        info(f\"Create dataset with {len(files)} files\")\n        \n    def load_next(self, verbose = False):\n        if verbose: info(f\"load_next> started: file_id={self.file_id}\")\n        del self.data        \n        self.epoch = 0\n        state = torch.load(self.files[self.file_id])\n        self.cols_df     = state['cols_df']\n        self.cols_agg_df = state['cols_agg_df']        \n        self.data = state['data']\n        if self.drop:\n            for v in self.data:\n                v.append( torch.onese(len(v[0]))*1.5 )\n        \n        self.file_id += 1\n        self.file_id = self.file_id % len(self.files)\n        if verbose: info(f\"load_next loaded,  tokens: {len(self.data)}\")\n\n        if self.device != 'cpu':\n            for k in self.data:\n                for i in range(len(self.data[k])):\n                    self.data[k][i] = self.data[k][i].to(self.device)\n            if verbose: info(f\"load_next sended to GPU\")\n\n        self.create_batches()          # create batches\n        self.batch_id   = 0            # current batch        \n\n    def set(self, batch_size, batch_max, T_max, is_T1=0, is_T2=1e8):\n        self.batch_size = batch_size\n        self.batch_max  = batch_max\n        self.T_max      = T_max\n        self.is_T1      = is_T1\n        self.is_T2      = is_T2\n        self.create_batches()\n\n    def create_batches(self):      \n        \"\"\" создать ссылки на индексы начала и конца батча \"\"\"         \n        self.batches    = []           # list of pointers to batch  [ (T, idx1, idx2) ]\n        for T,v in self.data.items():  # create pointers to batches [ (T, idx1, idx2) ]\n            if self.is_T1 <= T and T <= self.is_T2:\n                batch_size = min(int(self.batch_size * (self.T_max/T)**2),  self.batch_max)                \n                for i in range(0, v[0].shape[0], batch_size):\n                    self.batches.append((T, i, i + batch_size))\n\n    def reset(self):        \n        self.batch_id   = 0            \n        if self.shuffle:\n            for k in self.data:         # shuffle all samples:                \n                idx = torch.randperm( len(self.data[k][0]), device=self.device )\n                for i in range(len(self.data[k])):\n                    self.data[k][i] = self.data[k][i][idx]            \n            random.shuffle(self.batches) \n\n        if self.drop:\n            for k in self.data:         # sort samples by errors (best first):\n                idx = torch.argsort(self.data[k][-1])                \n                for i in range(len(self.data[k])):\n                    self.data[k][i] = self.data[k][i][idx]\n\n    def __next__(self):        \n        if self.batch_id >= len(self.batches):\n            self.epoch += 1\n            self.batch_id = 0\n            raise StopIteration  \n\n        p = self.batches[self.batch_id]\n        data  = self.data[p[0]]\n        self.batch_id += 1                                                                #!!!!!!!!!!!!!!!!\n\n        # (B,)    (B,T)    (B,T,F)  (B,24*2)  (B,3) \n        EVENT_ID, SENSOR_ID, FEAT,  AGG,  Y = data[0][p[1]:p[2]], data[1][p[1]:p[2]], data[2][p[1]:p[2]], data[3][p[1]:p[2], 0: CFG.AF ], data[4][p[1]:p[2]]\n        if self.aux:            \n            B,T,F = FEAT.shape\n            Y = torch.zeros(B,T, device=self.device)\n            num = int(T*0.1)\n            if num > 0:         \n                idx = torch.randperm(B*T, device=self.device)[:num]\n                sens_add =  SENSOR_ID.view(B*T)   [idx]                                 \n                feat_add =  FEAT.     view(B*T,-1)[idx]\n                idx = torch.randperm(B*T, device=self.device)\n                sens_add = sens_add[idx].view(B,T)\n                feat_add = feat_add[idx].view(B,T,-1)\n                y_add    = torch.ones(B, num, device=self.device)\n                SENSOR_ID = torch.cat([SENSOR_ID, sens_add], dim=1)\n                FEAT      = torch.cat([FEAT,      feat_add], dim=1)\n                Y         = torch.cat([Y, y_add],            dim=1)\n        return EVENT_ID, SENSOR_ID, FEAT, AGG, Y, None # (last ERR)\n\n    def __iter__(self):\n        return self\n    \n    def __len__(self):\n        return len(self.batches)     \n#---------------------------------------------------------------------------\n\n","metadata":{"execution":{"iopub.status.busy":"2023-04-18T16:44:12.152552Z","iopub.execute_input":"2023-04-18T16:44:12.152964Z","iopub.status.idle":"2023-04-18T16:44:12.178933Z","shell.execute_reply.started":"2023-04-18T16:44:12.152929Z","shell.execute_reply":"2023-04-18T16:44:12.177977Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model","metadata":{}},{"cell_type":"code","source":"sys.path.insert(1,\"/kaggle/input/nnet-lib/\")\nfrom importlib import reload\nimport nnet_v1\n\nreload(nnet_v1)\nfrom nnet_v1 import MLP, TransformerBlock, PointsBlock\n\n#===============================================================================\n\nclass FeatureGenerator(nn.Module):\n    def __init__(self, cfg, doms, V=5160):\n        \"\"\" \n        Генератор фич. Получает исходные фичи пульсов, выдаёт тензор (B,T,E) \n        V - number of sensors (embedding vocab) from sensors_df    \n        \"\"\"\n        super(FeatureGenerator, self).__init__()      \n\n        dom_feat = 4                       # {x,y,z,core}\n        if cfg.is_abs and not cfg.is_rho:  # {x,y,z,core,a} \n            dom_feat = 5        \n        else:                              # {x,y,z,core,a,r}\n            dom_feat = 6          \n\n        self.dom = nn.Embedding(V, dom_feat,_weight=doms[:,:dom_feat]).requires_grad_(False)\n        self.emb = nn.Embedding(V, cfg.SE) if cfg.is_emb else None # sensor embedding                \n     \n        F = dom_feat+cfg.F   # {x,y,z,core [,a]} + {t,q,aux}\n        if cfg.is_emb:\n            F += cfg.SE      # добавить SE-эмбединг сенсоров        \n        self.mlp = MLP(dict(input=F, hidden=cfg.Hin, output=cfg.n_embd)) if cfg.Hin > 0 else nn.Linear(F, cfg.n_embd)\n\n    def forward(self, SENSOR_ID, FEAT):                       \n        s  = self.dom(SENSOR_ID)                   # (B,T,4/5) 'x','y','z','core'[,'a'],r  \n        if self.emb is None:\n            x  = torch.cat([s, FEAT], dim=2)       # (B,T,E0)  pulse feat  +  'aux','q','t'      \n            x  = self.mlp(x)                       # (B,T,E)   increase num of features            \n        else:\n            em = self.emb(SENSOR_ID)\n            x  = torch.cat([s, FEAT, em], dim=2)   # (B,T,F+6+SE) concat with pulse features       \n            x  = self.mlp(x)                       # (B,T,E)   increase num of features            \n        return x\n#-------------------------------------------------------------------------------\n\nclass Transformer(nn.Module):\n    def __init__(self, cfg, drop=0):\n        \"\"\" Трансформер (B,T,E) -> (B,T,E) \"\"\"\n        super(Transformer, self).__init__()     \n        if cfg.is_pnt:\n            self.blocks= nn.ModuleList([PointsBlock(     dict(emb=cfg.n_embd, res=cfg.res)) for _ in range(cfg.L)])         \n        else: \n            self.blocks= nn.ModuleList([TransformerBlock(dict(emb=cfg.n_embd, res=cfg.res, att={'heads': 8})) for _ in range(cfg.L)])         \n        self.drop = nn.Dropout(drop)\n        self.ln   = nn.LayerNorm(cfg.n_embd)        \n\n    def forward(self, x):\n        x   = self.drop(x)              \n        x   = self.ln(x)                         # layer normalization (? batch)\n\n        for i, block in enumerate(self.blocks):  \n            if CFG.frozen and self.training and i > CFG.L_frozen: torch.set_grad_enabled(True) \n            x = block(x)                         # (B,T,E)            \n        return x\n#-------------------------------------------------------------------------------\n\nclass Summator(nn.Module):\n    def __init__(self, cfg):\n        \"\"\" Собирает со всех пульсов фичи в один вектор (B,T,E) -> (B,E) \"\"\"\n        super(Summator, self).__init__()               \n        self.ln    = nn.LayerNorm(cfg.n_embd)        \n        self.rnn   = nn.GRU(cfg.n_embd, cfg.Eh, batch_first=True) if cfg.is_rnn else None\n        self.cfg   = cfg\n\n    def forward(self, x):    \n        x = self.ln(x)                           # !!??\n        if self.rnn is None:\n            if self.cfg.is_max:                  # (B,2*E)\n                x =torch.cat([x.mean(dim=1), x.amax(dim=1)], dim=1)\n            else:\n                x = torch.mean(x, dim=1)         # (B,E) \n        else: \n            _, x = self.rnn(x)                   # (1,B,Eh) integrate all pulses\n            x = torch.squeeze(x, dim=0)          # (B,Eh)            \n        return x\n#-------------------------------------------------------------------------------\n\nclass Conv1D(nn.Module):\n    def __init__(self,cfg):\n        # K = 2 (-1), 3 (-2), 4 (-3),  5 (-4)\n        E = cfg.n_embd\n        cfg.blocks = [\n            {'L': 10, 'K': 2,  'T': 10},  # T > 'T'\n            {'L': 10, 'K': 3,  'T': 20},\n            {'L': 10, 'K': 4,  'T': 30},\n            {'L': 10, 'K': 5,  'T': 40},\n            {'L': 12, 'K': 5,  'T': 48},\n        ]\n        self.block_id = [0] * cfg.blocks[0]['T']\n        for  i in range(1, len(cfg.blocks) ):\n            self.block_id = self.block_id + [ i ] * (cfg.blocks[i]['T'] - cfg.blocks[i-1]['T']) \n\n        blocks = [ nn.Identity() ]\n        for b in cfg.blocks:\n            blocks.append( nn.Conv1d(E,E,kernel_size=b['K']) )\n        self.blocks = nn.ModuleList(blocks)\n        self.ln   = nn.LayerNorm(cfg.n_embd)   \n\n    def forward(self, x):\n        T = x.shape[1]\n        x = self.ln(x)\n        i = self.block_id[T] if T < len(self.block_id) else self.block_id[-1]\n        block = self.blocks[i]            \n        x = block(x)                                                   # (B,T,E)            \n        return x\n#-------------------------------------------------------------------------------\n\nclass Model(nn.Module):\n    def __init__(self,cfg, doms):\n        \"\"\" Модель предсказания (общая для регрессора и классификатора) \"\"\"\n        super(Model, self).__init__()    \n        self.cfg = copy.deepcopy(cfg)           \n        self.gen = FeatureGenerator(cfg, doms)        \n        if self.cfg.tri == 0:\n            self.att = Transformer(cfg)\n        else:\n            self.atts= nn.ModuleList([Transformer(cfg) for _ in range(cfg.tri)])         \n        self.sum = Summator(cfg)\n         \n        self.is_agg = cfg.is_agg\n        if cfg.is_rnn:\n            Ein  = cfg.Eh\n        else: \n            Ein  = 2*cfg.n_embd  if cfg.is_max else cfg.n_embd\n            if self.cfg.tri > 0:\n                Ein *= self.cfg.tri\n        if cfg.is_agg:\n            Ein += cfg.AF\n        Eout = 3 if cfg.is_reg else cfg.nums[0]*cfg.nums[1]  \n\n        self.mlp = MLP(dict(input=Ein, hidden=cfg.Hout, output=Eout))        \n\n    def forward(self, SENSOR_ID, FEAT, AGG, Y):    \n        if CFG.frozen: torch.set_grad_enabled(False) \n\n        x = self.gen(SENSOR_ID, FEAT)            # (B,T,E)\n        if self.cfg.tri == 0:\n            x = self.att(x)                      # (B,T,E)\n            x = self.sum(x)                      # (B,E)\n        else:\n            res = []\n            for att in self.atts:\n                res.append( self.sum (att(x)) )\n            x = torch.cat(res, dim=-1)           # (B,E*tri)\n        \n\n        if CFG.frozen and self.training: torch.set_grad_enabled(True) \n\n        if self.is_agg:\n            AGG = torch.clip(AGG, -10, 10)       # на всякий случай\n            x = torch.cat([x, AGG], dim=1)       # (B, Eh+AF) or (B, E+AF)\n\n        x = self.mlp(x)                          # (B,...)\n        return x\n\n#-------------------------------------------------------------------------------\n\nclass Regression(nn.Module):\n    def __init__(self,cfg, doms):\n        \"\"\" Регрессионная модель - предсказываем три компоненты направления \"\"\"\n        super(Regression, self).__init__()     \n        self.model = Model(cfg, doms)          \n\n    def forward(self, SENSOR_ID, FEAT, AGG, Y, eps=1e-8):    \n        x = self.model(SENSOR_ID, FEAT, AGG, Y) # (B,3)\n\n        if   CFG.loss == 'cos':\n            kappa = torch.norm(x, dim=1, keepdim=True).clip(eps)\n            y = x / kappa                          \n            cos = (y*Y).sum(dim=1).mean()\n            loss = ( 1 - cos.mean() ) + CFG.ka_reg * (x*x).sum(dim=1).mean()  #  * B / CFG.batch_size  # !!!???\n        elif CFG.loss == 'prod':\n            loss = -((x*Y).sum(dim=1)).mean() + CFG.ka_reg * (x*x).sum(dim=1).mean()   \n        elif CFG.loss == 'vMF':\n            kappa = torch.norm(x, dim=1, keepdim=True).clip(eps)\n            logC  = -kappa + torch.log( ( kappa+eps )/( 1-torch.exp(-2*kappa)+2*eps ) )\n            loss =  -( (x*Y).sum(dim=1) + logC ).mean() \n        elif CFG.loss == 'k2':             \n            loss = -((x*Y).sum(dim=1)).mean() + 0.5 * (x*x).sum(dim=1).mean()\n        elif CFG.loss == 'azze':\n            kappa = torch.norm(x, dim=1, keepdim=True)\n            y = x / kappa.clip(eps)              \n            r2y   =  y[:,0]*y[:,0] + y[:,1]*y[:,1] \n            r2Y   =  Y[:,0]*Y[:,0] + Y[:,1]*Y[:,1]\n            ryY   =  torch.sqrt(r2y*r2Y)\n\n            cos  = (y*Y).sum(dim=1)\n            cosA = (y[:,0]*Y[:,0] + y[:,1]*Y[:,1]) / ryY.clip(eps)                \n            cosZ = y[:,2]*Y[:,2] + ryY      # cos(theta_y-theta_Y)\n\n            loss = 1 - cos.mean()               \\\n                 + CFG.az_reg * (1-cosA.mean()) \\\n                 + CFG.ze_reg * (1-cosZ.mean()) \\\n                 + CFG.ka_reg * (x*x).sum(dim=1).mean()\n\n        with torch.no_grad():                \n            kappa = torch.norm(x.detach(), dim=1, keepdim=True).clip(eps)                        \n            y = x.detach() / kappa\n            ang_err, az_err, ze_err = Phys.angle_errors(y, Y, eps=eps)            \n        return loss, y.detach(), ang_err.detach(), az_err.detach(), torch.abs(ze_err.detach()),  kappa.detach()\n\n#-------------------------------------------------------------------------------\n\nclass Classifier(nn.Module):\n    def __init__(self, cfg, doms):\n        \"\"\" \"\"\"\n        super(Classifier, self).__init__()               \n        self.model = Model(cfg, doms)      \n        self.cfg   = cfg    \n\n    def forward(self, SENSOR_ID, FEAT, AGG, Y, eps=1e-8):    \n        x = self.model(SENSOR_ID, FEAT, AGG, Y) # (B,3)\n\n        kappa = torch.square(x).mean()\n        az_true, ze_true = Phys.vector2angles(Y)\n        id_true = Phys.angles2index(az_true, ze_true, n_az=self.cfg.nums[0], n_ze=self.cfg.nums[1])\n        CE_loss = nn.CrossEntropyLoss()\n        loss = CE_loss(x, id_true) + CFG.ka_reg * kappa\n        pred = x.detach().argmax(axis=1)\n        az_pred, ze_pred = Phys.index2angles(pred, n_az=self.cfg.nums[0], n_ze=self.cfg.nums[1])\n        y = Phys.angles2vector(az_pred, ze_pred)        \n        ang_err, az_err, ze_err = Phys.angle_errors(y, Y, eps=eps)            \n        return loss, y.detach(), ang_err.detach(), az_err.detach(), torch.abs(ze_err.detach()), torch.sqrt(kappa.detach())\n\n#-------------------------------------------------------------------------------\n\nclass AuxModel(nn.Module):\n    def __init__(self, cfg, doms):\n        \"\"\" \"\"\"\n        super().__init__()               \n        self.cfg = copy.deepcopy(cfg) \n        self.model = None          \n        self.gen = FeatureGenerator(cfg, doms)\n        self.att = Transformer(cfg) \n        self.sum = Summator(cfg)\n        inp =  3*cfg.n_embd if self.cfg.is_max else 2*cfg.n_embd\n        self.mlp = MLP(dict(input=inp, stretch=4, output=1))\n\n    #                                       (B,T)\n    def forward(self, SENSOR_ID, FEAT, AGG, Y, eps=1e-8):            \n        x = self.gen(SENSOR_ID, FEAT)            # (B,T,E)\n        x = self.att(x)                          # (B,T,E)                \n        y = self.sum(x)[:,None,:]                # (B,1,E) or (B,1,2*E) \n        y = y.repeat(1,x.shape[1],1)             # (B,T,E)        \n        x = torch.cat([x,y], dim=-1)             \n        x = self.mlp(x).squeeze(dim=-1)          # (B,T)\n        x = torch.sigmoid(x)\n\n        kappa = x.pow(2).mean()        \n        loss = ((x-Y)**2).mean() + CFG.ka_reg * kappa\n        err  = 1 - ((x.detach() > 0.5) == Y).to(torch.float32).mean(dim=-1)\n        mean1, mean2 = Y.mean(dim=-1), x.detach().mean(dim=-1)\n\n        return loss, x.detach(), err, mean1, mean2, torch.zeros(1, device=x.device)\n\n#-------------------------------------------------------------------------------\n\nclass Phys:\n    def angle_errors(n1, n2, eps=1e-8):\n        \"\"\" Calculate angles between two unit (!!!) vectors:: n1,n2: (B,3) return: (B,) \"\"\"\n        with torch.no_grad():\n            cos = (n1*n2).sum(axis=1)                     # angles between vectors\n            angle_err = torch.arccos( cos.clip(-1,1) )    \n        \n            r1   =  n1[:,0]*n1[:,0] + n1[:,1]*n1[:,1]    # angles between vectors in (x,y)    \n            r2   =  n2[:,0]*n2[:,0] + n2[:,1]*n2[:,1]\n            norm = torch.sqrt(r1*r2)\n            cosX = (n1[:,0]*n2[:,0] + n1[:,1]*n2[:,1]) / norm.clip(eps)    \n            azimuth_err = torch.arccos( cosX.clip(-1,1) )\n                                \n            zerros = norm < eps                            # azimuth angle not defined\n            azimuth_err[zerros] = torch.rand((len(n1[zerros]),), device=n1.device)*torch.pi\n    \n            zenith1  = torch.arccos( n1[:,2].clip(-1,1) )\n            zenith2  = torch.arccos( n2[:,2].clip(-1,1) )\n            zenith_err = zenith2 - zenith1    \n        \n        return angle_err, azimuth_err, zenith_err\n\n    def vector2angles(n, eps=1e-8):\n        \"\"\"  Get spherical angles of vector n: (B,3) \"\"\"                \n        n = n / torch.norm(n, dim=1, keepdim=True).clip(eps)\n                                \n        azimuth = torch.arctan2( n[:,1],  n[:,0])    \n        azimuth[azimuth < 0] += 2*torch.pi\n                                \n        zenith = torch.arccos( n[:,2].clip(-1,1) )                                \n    \n        return azimuth, zenith\n\n    def angles2vector(azimuth, zenith):\n        \"\"\" Add unit vector components from (azimuth,zenith) to the DataFrame df \"\"\"\n        nx = (torch.sin(zenith) * torch.cos(azimuth)).view(-1,1)\n        ny = (torch.sin(zenith) * torch.sin(azimuth)).view(-1,1)\n        nz = torch.cos(zenith).view(-1,1)\n        return torch.cat([nx,ny,nz], dim=1)\n\n    def angles2index(azimuth, zenith, n_az, n_ze):\n        \"\"\" \"\"\"\n        az = torch.floor(n_az * azimuth / (2*np.pi)).clip(0,n_az-1)\n        ze = torch.floor(n_ze * (torch.cos(zenith)+1)/2).clip(0,n_az-1)\n        return (az*n_ze + ze).to(torch.long)\n\n    def index2angles(index, n_az, n_ze):\n        \"\"\" \"\"\"\n        az = (index // n_ze)  \n        ze = torch.arccos( (2*((index - az*n_ze)/n_ze) - 1).clip(-1,1) )\n        az = az * (2*torch.pi) / n_az\n        return az, ze        \n\n#===============================================================================\n\ndef load_model(fname, doms):\n    state = torch.load(fname)\n    if 'tri'       not in state['config']: state['config']['tri']       = 0\n    if 'res'       not in state['config']: state['config']['res']       = 1\n    if 'is_max'    not in state['config']: state['config']['is_max']    = 0\n    if 'is_rho'    not in state['config']: state['config']['is_rho']    = 0    \n    if 'is_pnt'    not in state['config']: state['config']['is_pnt']    = 0        \n    if 'n_embd'    not in state['config']: state['config']['n_embd']    =  state['config']['E']\n    if 'n_head'    not in state['config']: state['config']['n_head']    =  8\n    if 'causal'    not in state['config']: state['config']['causal']    =  False    \n    if 'drop_attn' not in state['config']: state['config']['drop_attn'] =  0    \n    if 'drop_mlp'  not in state['config']: state['config']['drop_mlp']  =  0    \n    if 'hidden'    not in state['config']: state['config']['hidden']    =  4    \n    if 'fun'       not in state['config']: state['config']['fun']       =  'gelu'     \n    if 'poly'      not in state['config']: state['config']['poly']      = False\n\n    print(state['config'])    \n    class ModelConfig(object): pass     \n    cfg = ModelConfig()                 \n    for key in state['config']:\n        setattr(cfg, key, state['config'][key])\n\n    if cfg.is_reg:\n        model = Regression(cfg, doms)\n    else:\n        model = Classifier(cfg, doms)\n    model.load_state_dict(state['model'])     \n\n    history     = state.get('history',[])\n    labels      = state.get('labels', [])\n    checks      = state.get('checks', [])\n\n    CFG.samples = history[-1][0] if len(history) else 0       \n    CFG.steps   = state.get('steps', 0)\n    CFG.last    = state.get('last',  0)\n    CFG.score   = state.get('score', 100)\n    \n    return model, history, labels, checks\n\n#===============================================================================\n\ndef model_state(model):\n    cfg = model.model.cfg if model.model is not None else model.cfg\n    state = {'info':      \"IceCube at\", \n             'date':      datetime.datetime.now(),   # дата и время\n             'config':    dict([(k,v) for k,v in cfg.__dict__.items() if not k.startswith(\"__\")]),\n             'model' :    model.state_dict(),        # параметры модели         \n             'optimizer': optimizer.state_dict(),    # состояние оптимизатора\n             'history':   history,\n             'labels':    labels,\n             'checks':    checks,\n             'score':     CFG.score,\n             'steps':     CFG.steps,\n             'samples':   CFG.samples,\n             }    \n    return state\n\n#-------------------------------------------------------------------------------\ndef save_model(model, fname):              \n    state = model_state(model)  \n    torch.save(state, fname)\n\ndef save_best_model(model, score, folder=\"/content/drive/MyDrive/IceCube/\"):              \n    cfg = model.model.cfg if model.model is not None else model.cfg\n    CFG.last = CFG.samples    \n    CFG.score= round(score, 5)        \n    CFG.best = f'att_{score:.4f}_D{DATA_KIND}_A{ModelCFG.arch}_F{ModelCFG.F}_AF{ModelCFG.AF}_E{ModelCFG.n_embd}_L{cfg.L}_Eh{ModelCFG.Eh}_Hout{ModelCFG.Hout}_T{CFG.T_max}'\n    state = model_state(model)  \n    torch.save(state, folder+CFG.best+\".pt\")\n\ndef save_model_checkpoint(model,  error_val, error_trn, folder=\"/content/drive/MyDrive/IceCube/checkpoints/\"):                \n    state = model_state(model)  \n    now = datetime.datetime.now().strftime(\"%m-%d %H:%M:%S\")\n    fname = f'{now} er_val_{error_val:.4f}_trn_{error_trn:.4f} az_{history[-1][3]:.3f} ze_{history[-1][4]:.3f} ka_{history[-1][-4]:.3f}.pt'    \n    torch.save(state, folder+fname)\n\n#===============================================================================\n\ndef load_model_old(fname, model):\n    CFG.best = fname  \n    state = torch.load('/content/drive/MyDrive/IceCube/'+CFG.best)       # /best\n    print(state['config'])\n    #for k in state['model'].keys(): print(k)\n    model.load_state_dict(state['model'])     \n    CFG.samples = state['history'][-1][0] if len(state['history']) else 0       \n    CFG.steps   = state['steps']          if 'steps' in state      else 0\n    CFG.last    = state['config']['last']\n    CFG.score   = state['score']\n    history     = state['history']\n    labels      = state['labels'] if 'labels' in state else []\n    checks      = state['checks'] if 'checks' in state else []    \n    print(state['labels'])\n    return history, labels, checks\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Submission","metadata":{}},{"cell_type":"code","source":"def angle_errors(n1, n2, eps=1e-8):\n    \"\"\" Calculate angles between two vectors:: n1,n2: (B,3) return: (B,) \"\"\"\n    n1 = n1 / (np.linalg.norm(n1, axis=1, keepdims=True) + eps)\n    n2 = n2 / (np.linalg.norm(n2, axis=1, keepdims=True) + eps)\n    \n    cos = (n1*n2).sum(axis=1)                     # angles between vectors\n    angle_err = np.arccos( cos.clip(-1,1) )    \n        \n    r1   =  n1[:,0]*n1[:,0] + n1[:,1]*n1[:,1]    # angles between vectors in (x,y)    \n    r2   =  n2[:,0]*n2[:,0] + n2[:,1]*n2[:,1]\n    cosX = (n1[:,0]*n2[:,0] + n1[:,1]*n2[:,1]) / (np.sqrt(r1*r2) + eps)    \n    azimuth_err = np.arccos( cosX.clip(-1,1) )\n                                \n    zerros = r1 < eps                            # azimuth angle not defined\n    azimuth_err[zerros] = np.random.random((len(n1[zerros]),))*np.pi\n    \n    zenith1  = np.arccos( n1[:,2].clip(-1,1) )\n    zenith2  = np.arccos( n2[:,2].clip(-1,1) )\n    zenith_err = zenith2 - zenith1    \n        \n    return angle_err, azimuth_err, zenith_err\n\n#-------------------------------------------------------------------------------\n\ndef stat(acc, bins = 200):    \n    plt.figure(figsize=(6,4), facecolor ='w') \n    plt.axes().set_facecolor(\"ivory\"); plt.autoscale(tight=True)\n    p,_,_ = plt.hist(acc, bins=bins, range=(0,np.pi), fc=\"lightblue\", density=True, alpha=0.5)\n    w = p[len(p)//2: ].sum()*np.pi/bins    \n    x = np.linspace(0,np.pi,bins)\n    plt.plot(x, w*np.sin(x), c=\"r\")\n    plt.plot(x, p-w*np.sin(x), c=\"b\")\n    plt.title(f\"Transformer:  mean={np.mean(acc):.4f}, median={np.median(acc):.3f}, w={w:.3f}\")    \n    plt.ylabel(\"Density\"); plt.xlabel(r\"$\\Delta \\Psi$ (rad)\")\n    plt.grid()\n    plt.show()\n    print(f\"Transformer:  mean={np.mean(acc):.4f}, median={np.median(acc):.3f}, w={w:.3f}\")\n#------------------------------------------------------------------------------\n\ndef submission(model, events_df, model_numb, del_all=False):\n    model = model.to(CFG.device)\n    model.train(False)     \n    \n    scores, counts,  = [], []    \n    pred_df = pd.DataFrame({\"event_id\":[], \"azimuth\":[], \"zenith\":[] } )\n    true_df = pd.DataFrame({\"event_id\":[], \"azimuth\":[], \"zenith\":[] } )    \n    \n    dataset_tst.file_id = 0 # !    \n    dataset_tst.load_next(True)\n    dataset_tst.set(batch_size=512, batch_max=2048, T_max=256)\n    \n    for file in FILES:\n        info(f\"run {file}\")\n        dataset_tst.reset()\n        for b, (EVENT_ID, SENSOR_ID, FEAT, AGG, Y, ERR) in tqdm(enumerate(dataset_tst), total=len(dataset_tst)):        \n            SENSOR_ID, FEAT, AGG, Y = SENSOR_ID.to(CFG.device), FEAT.to(CFG.device), AGG.to(CFG.device), Y.to(CFG.device)\n            with torch.no_grad():\n                loss, x,  ang_err, az_err, ze_err, kappa = model(SENSOR_ID, FEAT, AGG, Y)                        \n                x = x.data        \n            \n            azimuth = torch.arctan2(x[:,1], x[:,0])        \n            azimuth[azimuth < 0] += 2*np.pi\n            zenith = torch.arccos( x[:,2] )                        \n            pred_df = pred_df.append(pd.DataFrame({\"event_id\":EVENT_ID.view(-1,).numpy(), \"azimuth\":azimuth.cpu().numpy(), \"zenith\":zenith.cpu().numpy()} ))            \n\n            azimuth = torch.arctan2(Y[:,1], Y[:,0])        \n            azimuth[azimuth < 0] += 2*np.pi\n            zenith = torch.arccos( Y[:,2] )                        \n            true_df = true_df.append(pd.DataFrame({\"event_id\":EVENT_ID.view(-1,).numpy(), \"azimuth\":azimuth.cpu().numpy(), \"zenith\":zenith.cpu().numpy()} ))            \n        dataset_tst.load_next()\n        \n    pred_df.event_id = pred_df.event_id.astype(np.int64)\n    print(len(pred_df)- len(pred_df.event_id.unique()) )\n\n    info(\"CHECK_BEFORE_SUBMIT = False !!!!!!!\")\n    pred_df = pred_df.sort_values(['event_id']).reset_index(drop=True)               # !!!\n    #display(pred_df)\n    \n    true_df.event_id = true_df.event_id.astype(np.int64)\n    true_df = true_df.sort_values(['event_id']).reset_index(drop=True)               # !!!\n    #display(true_df)\n    \n    pred_df = angles2vector(pred_df)\n    true_df = angles2vector(true_df)\n    u = pred_df[['nx','ny','nz']].to_numpy()\n    n = true_df[['nx','ny','nz']].to_numpy()    \n    ang_err, az_err, ze_err = angle_errors(u,n)                     # inverse vector!\n    info(f\"mean={np.mean(ang_err):.3f}   medan={np.median(ang_err):.3f}   num={len(ang_err)}   min,max=[{np.min(ang_err):.3f},{np.max(ang_err):.3f}] {CFG.get()}\")    \n    pred_df['ang_err'] = ang_err\n    \n    if events_df is not None:\n        events_df.event_id = events_df.event_id.astype(np.int64)\n        display(events_df)\n        pred_df = events_df.merge(pred_df, left_on='event_id', right_on='event_id', how=\"left\")\n        if del_all: \n            del events_df\n        \n    print(\"Count rows with nan values: \", pred_df.isna().any(axis=1).sum() )\n    #display(pred_df[pred_df.isna().any(axis=1)])\n    pred_df = pred_df.sort_values(['event_id']).reset_index(drop=True)               # !!!\n    pred_df = pred_df.fillna(0)                                                      # !!!\n    #display(pred_df.info())\n    #display(pred_df)    \n        \n    pred_df[['event_id','azimuth','zenith']].to_csv(f'submission-att-{model_numb}.csv', index=False)    \n    !head submission.csv\n    \n    info(f\"the End: len(df)={len(pred_df)}\")    \n    return pred_df\n    \n#-------------------------------------------------------------------------------    \n#-------------------------------------------------------------------------------    \n\n\nmodels = [\n            ('att1', False, '/kaggle/input/icecube-models/att_0.9986_D01_A_F3_AF48_E128_L12_Eh256_Hout2048_T256.pt'),\n            ('att2', True,  '/kaggle/input/icecube-models/att_rnn_1.0015_L10.pt'),\n            ('att4', False, '/kaggle/input/icecube-models/att04_1.0003_D00.pt'),\n         ]\n\nfor i, each_model in enumerate(models):\n    \n    DOMS_AGG = each_model[1]\n    \n    info.beg = info.last = time.time()\n    info(\"begin\")\n    FILES = create_dataset_on_disk()        \n    info(\"the End\")\n\n    state = torch.load(\"doms.pt\")\n    DOMS  = state['data']\n    info(f\"load doms.pt: {state['cols']},  DOMS:{DOMS.shape}\")\n    dataset_tst = Dataset(FILES,  batch_size=256, shuffle=False,  device='cpu', drop=False)\n    info(CFG.get())\n\n    model, _, _, _ = load_model(each_model[2], DOMS)        \n\n    CFG.params = sum(p.numel() for p in model.parameters())\n    info(\"number of parameters: %.3fk  [%d]\" % (CFG.params/1e3, CFG.params))   \n\n    model = model.to(CFG.device)\n#     print(len(model.model.att.blocks))\n\n\n    df = submission(model=model, model_numb=i+1, events_df=None)\n    stat(df.ang_err.to_numpy())    \n\n    #del dataset_tst, DOMS, model, state, df; \n    gc.collect()\ninfo(\"the End\")","metadata":{"execution":{"iopub.status.busy":"2023-04-18T16:55:19.205752Z","iopub.execute_input":"2023-04-18T16:55:19.206126Z","iopub.status.idle":"2023-04-18T16:55:35.112396Z","shell.execute_reply.started":"2023-04-18T16:55:19.206091Z","shell.execute_reply":"2023-04-18T16:55:35.111237Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# GNN","metadata":{}},{"cell_type":"markdown","source":"## Config","metadata":{}},{"cell_type":"code","source":"import numpy as np\n\nclass CFG:\n    MODELS                = [\n        'gnn_id_26_val_0_9926_exp_classific_16n_e1633',\n        'gnn_id_27_val_0_9961_exp_readout_emb_e196',        \n        'gnn_id_28_val_0_9919_exp_classific_24n_e726',\n        #'gnn_id_20_val_0_9964_exp_mlp_2048_e1343'\n    ]\n    MODE                  = 'test' if GLOBAL_TEST_MODE else 'submit'      # train/test/submit\n    ENV                   = 'kaggle'     # kaggle/colab\n    BATCH_RANGE           = (1,GLOBAL_TEST_BATCHES+1)      # test batches only for test mode 1,5+1\n    MAX_PULSES_PER_EVENT  = 1000  \n    USE_WANDB             = False\n    USE_CUDA              = True\n    TTA_AZ_ANGLES         = [0.0] # [0.0, -2.0, 5.0, 37.0, -19.0] # [0.0]        # np.linspace(0.0,360.0,7)\n    LOADER                = 'pl'  # pl,pd,cudf\n    SAMPLE_FILTER         = False\n    ANGLES                = 'az,ze'  # az,ze/az/ze\n    ZENITH_RANGE          = None # (0.0, math.pi/2.0) # (math.pi/2.0, math.pi) # None\n    FROZEN                = False\n    UNFOROZEN_LAYERS      = []\n\nif CFG.ENV == 'kaggle':\n    CFG.REMOTE_DATASET_PATH    = '/kaggle/input/icecube-neutrinos-in-deep-ice'\n    CFG.DATASET_PATH           = CFG.REMOTE_DATASET_PATH\n    CFG.CACHE_DIR              = '/tmp/cache'\n    CFG.META_PATH              = f'{CFG.CACHE_DIR}/icecube-neutrinos-in-deep-ice/meta'\n    CFG.SCATTER_ABSORT_TABLE   = f'/kaggle/input/icecube-weights/scattering_and_absorption.csv'\n    CFG.DOMS_EFF_TABLE         = f'/kaggle/input/icecube-weights/doms_eff.csv'\n    \n    if CFG.MODE == 'test':\n        CFG.INPUT_DATA_PATH    = f'{CFG.DATASET_PATH}/train'\n        CFG.META_TABLE         = f'{CFG.DATASET_PATH}/train_meta.parquet'\n    if CFG.MODE == 'submit':\n        CFG.INPUT_DATA_PATH    = f'{CFG.DATASET_PATH}/test'\n        CFG.META_TABLE         = f'{CFG.DATASET_PATH}/test_meta.parquet'\n        \nif CFG.ENV == 'colab':\n    CFG.DATASET_PATH           = 'icecube-neutrinos-in-deep-ice'\n    CFG.REMOTE_DATASET_PATH    = '/content/drive/MyDrive/work/projects/icecube/datasets/icecube-neutrinos-in-deep-ice'\n    CFG.META_PATH              = '/content/content/icecube-neutrinos-in-deep-ice/train_meta'\n    CFG.CACHE_DIR              = 'cache'\n    \nCFG.GEOMETRY_TABLE         = f'{CFG.DATASET_PATH}/sensor_geometry.csv'\n\n","metadata":{"execution":{"iopub.status.busy":"2023-04-18T16:55:55.867690Z","iopub.execute_input":"2023-04-18T16:55:55.868201Z","iopub.status.idle":"2023-04-18T16:55:55.882163Z","shell.execute_reply.started":"2023-04-18T16:55:55.868161Z","shell.execute_reply":"2023-04-18T16:55:55.880962Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Lib","metadata":{}},{"cell_type":"code","source":"#@title helpers\n\n# Append to PATH\nimport sys\nimport gc\nsys.path.append('software/graphnet/src')\n\nimport random\nimport pyarrow.parquet as pq\nimport sqlite3\nimport pandas as pd\nimport sqlalchemy\nfrom tqdm import tqdm\nimport os\nfrom typing import Any, Dict, List, Optional, Union\nimport numpy as np\nimport math\nimport torch\nfrom torch.optim.adam import Adam\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom inspect import getfullargspec\nimport shutil\nfrom os import path\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch_geometric.data import Data\nfrom torch_geometric.data import Batch, Data\nfrom torch.utils.data import Subset\nfrom torch.optim.optimizer import Optimizer\nfrom torch.optim.lr_scheduler import _LRScheduler\nfrom scipy.interpolate import interp1d                \nimport polars as pl\n\ndevice = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n\nsensors_data = None\n\nif CFG.LOADER == 'cudf':\n    import cudf\n\nif CFG.MODE=='train':\n    from torchinfo import summary\n\nif CFG.USE_WANDB:\n    import wandb\n\n    def wandb_init():\n        wandb.login(key=CFG.WANDBAPIKEY)\n        wandb.init(\n            # set the wandb project where this run will be logged\n            project=\"icecube\",\n            name=f\"{CFG.EXP_ID}_{CFG.EXP_COMMENT}\",\n            config=CFG.TRAIN_CFG,\n            #resume=\"must\"\n        )\n\n\ndef get_sensors(sensor_features):\n    \"\"\" Get sensor positions \"\"\"            \n    df = pd.read_csv(CFG.GEOMETRY_TABLE)      \n    df['line_id'] = df.sensor_id // 60 + 1                 # string id\n    df['core']    = (df.line_id > 78).astype(np.float32)   # sensor from DeepCore\n    df.x = df.x.astype(np.float32)              # distances in kilometers\n    df.y = df.y.astype(np.float32)\n    df.z = df.z.astype(np.float32)    \n    \n    phys = pd.read_csv(CFG.SCATTER_ABSORT_TABLE)\n    phys.z = (phys.z).astype(np.float32)\n    phys.a = (phys.a * 1e-3).astype(np.float32)\n    phys.b = (phys.b * 1e-2).astype(np.float32)\n    interp_scatter = interp1d(phys.z, phys.a)\n    interp_absort  = interp1d(phys.z, phys.b)\n    df['sc']  = interp_scatter(df.z)\n    df['abs'] = interp_absort(df.z)\n    df['r']   = np.sqrt(df.x**2 + df.y**2)*1e-3\n\n    eff = pd.read_csv(CFG.DOMS_EFF_TABLE)\n    df['eff'] = eff.astype(np.float32)\n\n    sensors_tensor = torch.tensor(df[sensor_features].values, dtype=torch.float32, device=device)\n\n    sensors_data = torch.nn.Embedding(5160, len(sensor_features), device=device, _weight=sensors_tensor).requires_grad_(False)\n\n    return sensors_data\n\ndef split_meta_data():\n    if not path.isdir(CFG.META_PATH):\n        os.makedirs(CFG.META_PATH, exist_ok=True)\n    meta_data_iter = pq.ParquetFile(CFG.META_TABLE).iter_batches(batch_size = 200_000)\n    batch_ids = []\n    for meta_data_batch in tqdm(meta_data_iter):\n        meta_data_batch = meta_data_batch.to_pandas()\n        batch_id = pd.unique(meta_data_batch['batch_id'])[0]\n        if CFG.MODE == 'test':\n            if batch_id < CFG.BATCH_RANGE[0]:\n                continue            \n            if batch_id == CFG.BATCH_RANGE[1]:\n                break   \n        meta_data_batch.to_parquet(f'{CFG.META_PATH}/batch_{batch_id}_meta.parquet')\n        batch_ids.append(batch_id)        \n    batch_range = (min(batch_ids),max(batch_ids)+1)\n    return batch_range\n\n\ndef load_batch(batch_id, max_events, doms_agg=False, bin_num=None):\n    reindex = False\n\n    if CFG.MODE == 'submit':\n        batch_meta_df = (\n                      pl.read_parquet(f'{CFG.META_PATH}/batch_{batch_id}_meta.parquet')\n                    ).select(['event_id','first_pulse_index','last_pulse_index']\n                    ).with_columns([\n                      pl.lit(0.0).alias('azimuth').cast(pl.Float32), \n                      pl.lit(0.0).alias('zenith').cast(pl.Float32)])\n    else:\n        batch_meta_df = (\n                      pl.read_parquet(f'{CFG.META_PATH}/batch_{batch_id}_meta.parquet')\n                    ).select(['event_id','azimuth','zenith','first_pulse_index','last_pulse_index'])\n\n    if CFG.ZENITH_RANGE:\n        batch_meta_df = batch_meta_df.filter((pl.col(\"zenith\") > CFG.ZENITH_RANGE[0]) & (pl.col(\"zenith\") <= CFG.ZENITH_RANGE[1]))\n\n      \n    batch_df = (\n                      pl.read_parquet(f'{CFG.INPUT_DATA_PATH}/batch_{batch_id}.parquet')\n                    ).select(['event_id','time','charge','auxiliary','sensor_id']\n                    ).with_columns([\n                      pl.col(\"time\").cast(pl.Float32),\n                      pl.col(\"charge\").cast(pl.Float32),\n                      #pl.lit(0.0).alias('auxiliary').cast(pl.Float32), \n                      pl.col(\"auxiliary\").cast(pl.Float32),                      \n                      pl.col(\"sensor_id\").cast(pl.Float32)])\n                    \n    if max_events:\n        batch_meta_df = batch_meta_df[:max_events]\n        batch_df = batch_df[:batch_meta_df[-1]['last_pulse_index'][0]+1]\n\n    # CFG.FILTER_BY_ENERGY = None\n    # if CFG.FILTER_BY_ENERGY:\n    #     batch_df = batch_df.filter(pl.col(\"charge\")>CFG.FILTER_BY_ENERGY)\n    #     reindex = True\n\n    reindex = True # for GRU must be sorted\n    if doms_agg:\n        batch_df = batch_df.groupby(['event_id', 'sensor_id']).agg([\n                      pl.col(\"auxiliary\").mean(),\n                      pl.col(\"charge\").sum(),\n                      pl.col(\"time\").min()\n                  ])\n    else:\n        batch_indexes = torch.tensor(batch_meta_df.select(['event_id','first_pulse_index','last_pulse_index']).to_numpy(), device=device, dtype=torch.long)\n\n    if reindex:\n        batch_df = batch_df.sort(['event_id', 'time']).with_row_count()\n        batch_indexes_df = batch_df.groupby(['event_id']).agg([\n                      pl.col(\"row_nr\").min().alias('first_pulse_index'),\n                      pl.col(\"row_nr\").max().alias('last_pulse_index'),\n                  ]).sort(['event_id'])\n        batch_indexes = torch.tensor(batch_indexes_df.to_numpy(), device=device, dtype=torch.long)\n\n    if CFG.MAX_PULSES_PER_EVENT:\n        #clip indexes\n        batch_indexes[:,2] = batch_indexes[:,2] - batch_indexes[:,1] + 1                         # pulses len\n        batch_indexes[:,2] = torch.clip(batch_indexes[:,2], min=0, max=CFG.MAX_PULSES_PER_EVENT) # clip to max\n        batch_indexes[:,2] = batch_indexes[:,1] + batch_indexes[:,2]                             # new indexes\n\n    batch_angles    = torch.tensor(batch_meta_df.select(['azimuth','zenith']).to_numpy(), device=device, dtype=torch.float32)\n\n    if CFG.ANGLES == 'az,ze':\n      batch_direcions = angles_to_vectors(batch_angles)\n    if CFG.ANGLES == 'az':\n      batch_direcions = angles_to_vectors_2D(batch_angles[:,0])\n    if CFG.ANGLES == 'ze':\n      batch_direcions = angles_to_vectors_2D(batch_angles[:,1])\n\n    batch_classes = None\n\n    if bin_num:\n        azimuth_edges, zenith_edges = build_az_ze_edges(bin_num)\n        batch_classes = angles_to_code(batch_angles, bin_num, azimuth_edges, zenith_edges)\n\n    batch_features  = torch.tensor(batch_df.select(['time','charge','auxiliary','sensor_id']).to_numpy(), device=device, dtype=torch.float32)\n    \n    return batch_indexes, batch_direcions, batch_classes, batch_features\n\ndef build_dataloader(batch_id, shuffle, config, max_events=None, indexes=None):\n  print(f'build loader for batch {batch_id} with limit: {max_events}')\n  dataset = BatchDataset(batch_id, max_events, config['doms_agg'], config.get('bin_num'))\n  if not indexes is None:\n    dataset = Subset(dataset, indexes)\n  dataloader = DataLoader(dataset, batch_size=config['batch_size'], num_workers=0, shuffle=shuffle, collate_fn=collate_fn)\n  return dataloader\n\ndef collate_fn(graphs: List[Data]) -> Batch:\n    batch = Batch.from_data_list(graphs)\n    # map x,y,z\n    batch.sensor_id = batch.sensor_id.long()\n    batch.x = torch.cat([sensors_data(batch.sensor_id), batch.x], axis=1)\n    return batch\n\nclass BatchDataset(Dataset):\n \n  def __init__(self, batch_id, max_events=None, doms_agg=False, bin_num=None):\n    self.batch_indexes, self.batch_directions, self.batch_classes, self.batch_features = load_batch(batch_id, max_events, doms_agg, bin_num)\n    self.cache = {}\n    self.bin_num = bin_num\n \n  def __len__(self):\n    return len(self.batch_indexes)\n   \n  def __getitem__(self,idx):\n    cached = self.cache.get(idx)\n    if cached:\n      return cached\n    event_id, event_first_pulse, event_last_pulse = self.batch_indexes[idx] \n    event_direction = self.batch_directions[idx] \n    event_direction = event_direction.unsqueeze(0)\n    \n    x = self.batch_features[event_first_pulse:event_last_pulse,:-1]\n    sensor_id = self.batch_features[event_first_pulse:event_last_pulse,-1]\n\n    graph = Data(x=x, edge_index=None)\n    graph.n_pulses  = event_last_pulse - event_first_pulse + 1\n    graph.event_ids = event_id\n    graph.sensor_id = sensor_id\n    graph.direction = event_direction\n    if self.bin_num:\n        graph.class_id = self.batch_classes[idx].unsqueeze(0)\n    self.cache[idx] = graph\n    return graph\n\ndef angles_to_vectors(angles):\n    vectors = torch.empty((angles.shape[0], 3), device=angles.device, dtype=angles.dtype)\n    zen = angles[:,1]\n    az  = angles[:,0]\n    sz = torch.sin(zen)\n    vectors[:,0] = torch.cos(az)*sz # x\n    vectors[:,1] = torch.sin(az)*sz # y\n    vectors[:,2] = torch.cos(zen)   # z\n    return vectors\n\ndef vectors_to_angles(vectors):\n    vectors = vectors.clone()\n    v_squared = vectors.pow(2.0)\n        \n    ## Shortcut optimization for azimuth: calculate 2d unit vectors for x and y independent of z\n    xy_sq = torch.sum(v_squared[:, 0:2], axis=1)\n    xy_d = torch.sqrt(xy_sq)[:, None]\n        \n    vectors[:, 0:2] = torch.where(xy_d == 0, xy_d, vectors[:, 0:2]/xy_d)\n\n    ## For z, use full 3d unit vector\n    d = torch.sqrt(xy_sq + v_squared[:, 2])\n    vectors[:, 2] = torch.where(d == 0, d, vectors[:, 2]/d)\n\n    ## As mentioned by others, clip solely to avoid floating point errors, the unit vectors should already be within this range.\n    vectors =  torch.clip(vectors, -1, 1)\n\n    azimuth = torch.arccos(vectors[:, 0])\n    ## if y < 0, convert from quadrants 1 and 2 to quadrants 3 and 4\n    azimuth = torch.where(vectors[:, 1] >= 0, azimuth, 2*torch.pi - azimuth)\n    azimuth = torch.where(torch.isfinite(azimuth), azimuth, torch.tensor(0.0, dtype=azimuth.dtype, device=azimuth.device))\n\n    zenith = torch.arccos(vectors[:, 2])\n        \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 = torch.where(torch.isfinite(zenith), zenith, torch.tensor(math.pi/2, dtype=zenith.dtype, device=azimuth.device))\n\n    angles = torch.stack([azimuth, zenith], axis=1)\n    return angles\n\n\ndef angular_dist_score(all_true, all_pred):\n    az_true  = all_true[:,0]\n    zen_true = all_true[:,1]\n    az_pred  = all_pred[:,0]\n    zen_pred = all_pred[:,1]\n    sa1 = torch.sin(az_true)\n    ca1 = torch.cos(az_true)\n    sz1 = torch.sin(zen_true)\n    cz1 = torch.cos(zen_true)\n    sa2 = torch.sin(az_pred)\n    ca2 = torch.cos(az_pred)\n    sz2 = torch.sin(zen_pred)\n    cz2 = torch.cos(zen_pred)\n    scalar_prod = sz1*sz2*(ca1*ca2 + sa1*sa2) + (cz1*cz2)\n    scalar_prod = torch.clip(scalar_prod, -1, 1) \n    distanses = torch.abs(torch.arccos(scalar_prod))   \n    return torch.mean(distanses), distanses\n\ndef angle_errors(n1, n2, eps=1e-8):\n    \"\"\" Calculate angles between two vectors:: n1,n2: (B,3) return: (B,) \"\"\"\n    n1 = n1 / (torch.linalg.vector_norm(n1, dim=1, keepdims=True) + eps)\n    n2 = n2 / (torch.linalg.vector_norm(n2, dim=1, keepdims=True) + eps)\n    \n    cos = (n1*n2).sum(axis=1)                     # angles between vectors\n    angle_err = torch.arccos( cos.clip(-1,1) )\n        \n    r1   =  n1[:,0]*n1[:,0] + n1[:,1]*n1[:,1]    # angles between vectors in (x,y)    \n    r2   =  n2[:,0]*n2[:,0] + n2[:,1]*n2[:,1]\n    cosX = (n1[:,0]*n2[:,0] + n1[:,1]*n2[:,1]) / (torch.sqrt(r1*r2) + eps)    \n    azimuth_err = torch.arccos( cosX.clip(-1,1) )\n                                \n    zerros = r1 < eps                            # azimuth angle not defined\n\n    azimuth_err[zerros] = torch.rand((len(n1[zerros]),), dtype=n1.dtype, device=n1.device)*np.pi\n    \n    zenith1  = torch.arccos( n1[:,2].clip(-1,1) )\n    zenith2  = torch.arccos( n2[:,2].clip(-1,1) )\n    zenith_err = torch.abs(zenith2 - zenith1)\n        \n    return angle_err.mean(), azimuth_err.mean(), zenith_err.mean()\n\ndef angles_to_vectors_2D(angles):\n    vectors = torch.empty((angles.shape[0], 2), device=angles.device, dtype=angles.dtype)\n    vectors[:,0] = torch.cos(angles) # x\n    vectors[:,1] = torch.sin(angles) # y\n    return vectors\n\ndef vectors_to_angles_2D(vectors, azimuth=False):\n    vectors = vectors.clone()\n    v_squared = vectors.pow(2.0)\n        \n    xy_sq = torch.sum(v_squared, axis=1)\n    xy_d = torch.sqrt(xy_sq)[:, None]\n        \n    vectors = torch.where(xy_d == 0, xy_d, vectors/xy_d)\n\n    vectors =  torch.clip(vectors, -1, 1)\n\n    angles = torch.arccos(vectors[:, 0])\n    angles = torch.where(torch.isfinite(angles), angles, torch.tensor(0.0, dtype=angles.dtype, device=angles.device))\n    if azimuth:\n      angles = torch.where(vectors[:, 1] >= 0, angles, 2*torch.pi - angles) # [0,2pi] range\n    return angles\n\ndef angular_dist_score_2D(ang_true, ang_pred):\n    sa1 = torch.sin(ang_true)\n    ca1 = torch.cos(ang_true)\n    sa2 = torch.sin(ang_pred)\n    ca2 = torch.cos(ang_pred)\n    scalar_prod = ca1*ca2 + sa1*sa2\n    scalar_prod = torch.clip(scalar_prod, -1, 1) \n    distanses = torch.abs(torch.arccos(scalar_prod))   \n    return torch.mean(distanses), distanses\n\ndef get_rot_matrix(phi):\n    s = torch.sin(phi)\n    c = torch.cos(phi)\n    rot = torch.stack([torch.stack([c, -s]),\n                       torch.stack([s, c])])\n    rot = rot.squeeze(-1)\n    return rot\n\ndef build_az_ze_edges(bin_num):\n    # Create Azimuth Edges\n    azimuth_edges = torch.tensor(np.linspace(0, 2 * np.pi, bin_num + 1), dtype=torch.float32).to(device)\n    # Create Zenith Edges\n    zenith_edges = []\n    zenith_edges.append(0)\n    for bin_idx in range(1, bin_num):\n        zenith_edges.append(np.arccos(np.cos(zenith_edges[-1]) - 2 / (bin_num)))\n    zenith_edges.append(np.pi)\n    zenith_edges = torch.tensor(np.array(zenith_edges), dtype=torch.float32).to(device)\n    return azimuth_edges, zenith_edges\n\ndef build_angle_bin_vector(azimuth_edges, zenith_edges, bin_num):\n    angle_bin_zenith0 = np.tile(zenith_edges[:-1], bin_num)\n    angle_bin_zenith1 = np.tile(zenith_edges[1:], bin_num)\n    angle_bin_azimuth0 = np.repeat(azimuth_edges[:-1], bin_num)\n    angle_bin_azimuth1 = np.repeat(azimuth_edges[1:], bin_num)\n\n    angle_bin_area = (angle_bin_azimuth1 - angle_bin_azimuth0) * (np.cos(angle_bin_zenith0) - np.cos(angle_bin_zenith1))\n    angle_bin_vector_sum_x = (np.sin(angle_bin_azimuth1) - np.sin(angle_bin_azimuth0)) * ((angle_bin_zenith1 - angle_bin_zenith0) / 2 - (np.sin(2 * angle_bin_zenith1) - np.sin(2 * angle_bin_zenith0)) / 4)\n    angle_bin_vector_sum_y = (np.cos(angle_bin_azimuth0) - np.cos(angle_bin_azimuth1)) * ((angle_bin_zenith1 - angle_bin_zenith0) / 2 - (np.sin(2 * angle_bin_zenith1) - np.sin(2 * angle_bin_zenith0)) / 4)\n    angle_bin_vector_sum_z = (angle_bin_azimuth1 - angle_bin_azimuth0) * ((np.cos(2 * angle_bin_zenith0) - np.cos(2 * angle_bin_zenith1)) / 4)\n\n    angle_bin_vector_mean_x = angle_bin_vector_sum_x / angle_bin_area\n    angle_bin_vector_mean_y = angle_bin_vector_sum_y / angle_bin_area\n    angle_bin_vector_mean_z = angle_bin_vector_sum_z / angle_bin_area\n\n    angle_bin_vector = np.zeros((1, bin_num * bin_num, 3))\n    angle_bin_vector[:, :, 0] = angle_bin_vector_mean_x\n    angle_bin_vector[:, :, 1] = angle_bin_vector_mean_y\n    angle_bin_vector[:, :, 2] = angle_bin_vector_mean_z\n\n    angle_bin_vector = torch.tensor(angle_bin_vector, dtype=torch.float32).to(device)\n    return angle_bin_vector\n\ndef angles_to_code(angles,bin_num,azimuth_edges,zenith_edges):\n    azimuth_code = (angles[:, 0] > azimuth_edges[1:].reshape((-1, 1))).sum(axis=0)\n    zenith_code = (angles[:, 1] > zenith_edges[1:].reshape((-1, 1))).sum(axis=0)\n    angle_code = bin_num * azimuth_code + zenith_code\n    return angle_code\n\ndef code_to_vector(pred, bin_num, angle_bin_vector, max=False, epsilon=1e-8):\n    # convert prediction to vector\n    if max:\n        pred_vector = angle_bin_vector[0,pred.argmax(axis=1)]\n    else:\n        pred_vector = (pred.reshape((-1, bin_num * bin_num, 1)) * angle_bin_vector).sum(axis=1)\n            \n    # normalize\n    pred_vector_norm = torch.sqrt((pred_vector**2).sum(axis=1))\n    mask = pred_vector_norm < epsilon\n    pred_vector_norm[mask] = 1\n    \n    # assign <1, 0, 0> to very small vectors (badly predicted)\n    pred_vector /= pred_vector_norm.reshape((-1, 1))\n\n    pred_vector[mask] = torch.tensor([1., 0., 0.], dtype=pred.dtype, device=pred.device)\n\n    return pred_vector\n\ndef code_to_angle(pred, bin_num, angle_bin_vector, max=False, epsilon=1e-8):\n    # convert prediction to vector\n    pred_vector = code_to_vector(pred, bin_num, angle_bin_vector, max, epsilon)\n\n    # convert to angle\n    azimuth = torch.arctan2(pred_vector[:, 1], pred_vector[:, 0])\n    azimuth[azimuth < 0] += 2 * torch.pi\n    zenith = torch.arccos(pred_vector[:, 2])\n\n    angles = torch.cat([azimuth.view(-1,1), zenith.view(-1,1)], axis=1)\n    return angles\n\nclass Lion(Optimizer):\n  r\"\"\"Implements Lion algorithm.\"\"\"\n\n  def __init__(self, params, lr=1e-4, betas=(0.9, 0.99), weight_decay=0.0):\n    \"\"\"Initialize the hyperparameters.\n    Args:\n      params (iterable): iterable of parameters to optimize or dicts defining\n        parameter groups\n      lr (float, optional): learning rate (default: 1e-4)\n      betas (Tuple[float, float], optional): coefficients used for computing\n        running averages of gradient and its square (default: (0.9, 0.99))\n      weight_decay (float, optional): weight decay coefficient (default: 0)\n    \"\"\"\n\n    if not 0.0 <= lr:\n      raise ValueError('Invalid learning rate: {}'.format(lr))\n    if not 0.0 <= betas[0] < 1.0:\n      raise ValueError('Invalid beta parameter at index 0: {}'.format(betas[0]))\n    if not 0.0 <= betas[1] < 1.0:\n      raise ValueError('Invalid beta parameter at index 1: {}'.format(betas[1]))\n    defaults = dict(lr=lr, betas=betas, weight_decay=weight_decay)\n    super().__init__(params, defaults)\n\n  @torch.no_grad()\n  def step(self, closure=None):\n    \"\"\"Performs a single optimization step.\n    Args:\n      closure (callable, optional): A closure that reevaluates the model\n        and returns the loss.\n    Returns:\n      the loss.\n    \"\"\"\n    loss = None\n    if closure is not None:\n      with torch.enable_grad():\n        loss = closure()\n\n    for group in self.param_groups:\n      for p in group['params']:\n        if p.grad is None:\n          continue\n\n        # Perform stepweight decay\n        p.data.mul_(1 - group['lr'] * group['weight_decay'])\n\n        grad = p.grad\n        state = self.state[p]\n        # State initialization\n        if len(state) == 0:\n          # Exponential moving average of gradient values\n          state['exp_avg'] = torch.zeros_like(p)\n\n        exp_avg = state['exp_avg']\n        beta1, beta2 = group['betas']\n\n        # Weight update\n        update = exp_avg * beta1 + grad * (1 - beta1)\n        p.add_(torch.sign(update), alpha=-group['lr'])\n        # Decay the momentum running average coefficient\n        exp_avg.mul_(beta2).add_(grad, alpha=1 - beta2)\n\n    return loss\n\nclass PiecewiseLinearLR(_LRScheduler):\n    \"\"\"Interpolate learning rate linearly between milestones.\"\"\"\n\n    def __init__(\n        self,\n        optimizer: Optimizer,\n        milestones: List[int],\n        factors: List[float],\n        last_epoch: int = -1,\n        verbose: bool = False,\n    ):\n        \"\"\"Construct `PiecewiseLinearLR`.\n\n        For each milestone, denoting a specified number of steps, a factor\n        multiplying the base learning rate is specified. For steps between two\n        milestones, the learning rate is interpolated linearly between the two\n        closest milestones. For steps before the first milestone, the factor\n        for the first milestone is used; vice versa for steps after the last\n        milestone.\n\n        Args:\n            optimizer: Wrapped optimizer.\n            milestones: List of step indices. Must be increasing.\n            factors: List of multiplicative factors. Must be same length as\n                `milestones`.\n            last_epoch: The index of the last epoch.\n            verbose: If ``True``, prints a message to stdout for each update.\n        \"\"\"\n        # Check(s)\n        if milestones != sorted(milestones):\n            raise ValueError(\"Milestones must be increasing\")\n        if len(milestones) != len(factors):\n            raise ValueError(\n                \"Only multiplicative factor must be specified for each milestone.\"\n            )\n\n        self.milestones = milestones\n        self.factors = factors\n        super().__init__(optimizer, last_epoch, verbose)\n\n    def _get_factor(self) -> np.ndarray:\n        # Linearly interpolate multiplicative factor between milestones.\n        return np.interp(self.last_epoch, self.milestones, self.factors)\n\n    def get_lr(self) -> List[float]:\n        \"\"\"Get effective learning rate(s) for each optimizer.\"\"\"\n        if not self._get_lr_called_within_step:\n            warnings.warn(\n                \"To get the last learning rate computed by the scheduler, \"\n                \"please use `get_last_lr()`.\",\n                UserWarning,\n            )\n        return [base_lr * self._get_factor() for base_lr in self.base_lrs]\n\n#@title Conv\n\n\"\"\"Class(es) implementing layers to be used in `graphnet` models.\"\"\"\n\nfrom typing import Any, Callable, Optional, Sequence, Union\n\nfrom torch.functional import Tensor\nfrom torch_geometric.nn import EdgeConv, TransformerConv\nfrom torch_geometric.nn.pool import knn_graph\nfrom torch_geometric.typing import Adj\nfrom pytorch_lightning import LightningModule\n\nclass DynEdgeConv(EdgeConv, LightningModule):\n    \"\"\"Dynamical edge convolution layer.\"\"\"\n\n    def __init__(\n        self,\n        nn: Callable,\n        aggr: str = \"max\",\n        nb_neighbors: int = 8,\n        features_subset: Optional[Union[Sequence[int], slice]] = None,\n        last_layer = False,\n        **kwargs: Any,\n    ):\n        \"\"\"Construct `DynEdgeConv`.\n\n        Args:\n            nn: The MLP/torch.Module to be used within the `EdgeConv`.\n            aggr: Aggregation method to be used with `EdgeConv`.\n            nb_neighbors: Number of neighbours to be clustered after the\n                `EdgeConv` operation.\n            features_subset: Subset of features in `Data.x` that should be used\n                when dynamically performing the new graph clustering after the\n                `EdgeConv` operation. Defaults to all features.\n            **kwargs: Additional features to be passed to `EdgeConv`.\n        \"\"\"\n        # Check(s)\n        if features_subset is None:\n            features_subset = slice(None)  # Use all features\n        assert isinstance(features_subset, (list, slice))\n\n        # Base class constructor\n        super().__init__(nn=nn, aggr=aggr, **kwargs)\n\n        # Additional member variables\n        self.nb_neighbors = nb_neighbors\n        self.features_subset = features_subset\n        self.last_layer = last_layer\n\n    def forward(\n        self, x: Tensor, edge_index: Adj, batch: Optional[Tensor] = None\n    ) -> Tensor:\n        \"\"\"Forward pass.\"\"\"\n        # Standard EdgeConv forward pass\n        x = super().forward(x, edge_index)\n\n        if not self.last_layer:   # unnesessary last layer\n            # Recompute adjacency\n            edge_index = knn_graph(\n                x=x[:,self.features_subset],\n                k=self.nb_neighbors,\n                batch=batch,\n            ).to(self.device)\n\n        return x, edge_index\n\n\nclass DynTransformerConv(TransformerConv, LightningModule):\n    \"\"\"Dynamical edge convolution layer.\"\"\"\n\n    def __init__(\n        self, \n        in_channels, \n        out_channels: int, \n        nb_neighbors: int = 8,\n        features_subset: Optional[Union[Sequence[int], slice]] = None,\n        last_layer = False,\n        heads: int = 1, \n        concat: bool = True, \n        beta: bool = False, \n        dropout: float = 0.0, \n        edge_dim: Optional[int] = None, \n        bias: bool = True, \n        root_weight: bool = True, **kwargs,\n    ):\n\n        # Check(s)\n        if features_subset is None:\n            features_subset = slice(None)  # Use all features\n        assert isinstance(features_subset, (list, slice))\n\n        # Base class constructor\n        super().__init__(in_channels=in_channels, out_channels=out_channels, \n                         heads=heads, concat=concat, beta=beta, dropout=dropout, edge_dim=edge_dim, bias=bias, \n                         root_weight=root_weight, **kwargs)\n\n        # Additional member variables\n        self.nb_neighbors = nb_neighbors\n        self.features_subset = features_subset\n        self.last_layer = last_layer\n\n    def forward(\n        self, x: Tensor, edge_index: Adj, batch: Optional[Tensor] = None\n    ) -> Tensor:\n        \"\"\"Forward pass.\"\"\"\n        # Standard EdgeConv forward pass\n        x = super().forward(x, edge_index)\n\n        if not self.last_layer:   # unnesessary last layer\n            # Recompute adjacency\n            edge_index = knn_graph(\n                x=x[:,self.features_subset],\n                k=self.nb_neighbors,\n                batch=batch,\n            ).to(self.device)\n\n        return x, edge_index\n\n#@title DynEdge\n\n\"\"\"Implementation of the DynEdge GNN model architecture.\"\"\"\nfrom typing import List, Optional, Sequence, Tuple, Union\n\nimport torch\nfrom torch import Tensor, LongTensor\nfrom torch_geometric.data import Data\nfrom torch_scatter import scatter_max, scatter_mean, scatter_min, scatter_sum\n\n# from graphnet.models.components.layers import DynEdgeConv  # QuData: owerwrite\nfrom graphnet.utilities.config import save_model_config\nfrom graphnet.models.gnn.gnn import GNN\nfrom torch_geometric.utils.homophily import homophily\nfrom torch_geometric.nn.pool import TopKPooling\nfrom torch_geometric.utils import to_dense_batch\n\nGLOBAL_POOLINGS = {\n    \"min\": scatter_min,\n    \"max\": scatter_max,\n    \"sum\": scatter_sum,\n    \"mean\": scatter_mean,\n}\n\ndef calculate_xyzt_homophily(\n    x: Tensor, edge_index: LongTensor, batch: Batch\n) -> Tuple[Tensor, Tensor, Tensor, Tensor]:\n    \"\"\"Calculate xyzt-homophily from a batch of graphs.\n\n    Homophily is a graph scalar quantity that measures the likeness of\n    variables in nodes. Notice that this calculator assumes a special order of\n    input features in x.\n\n    Returns:\n        Tuple, each element with shape [batch_size,1].\n    \"\"\"\n    hx = homophily(edge_index, x[:, 0], batch).reshape(-1, 1)\n    hy = homophily(edge_index, x[:, 1], batch).reshape(-1, 1)\n    hz = homophily(edge_index, x[:, 2], batch).reshape(-1, 1)\n    ht = homophily(edge_index, x[:, -3], batch).reshape(-1, 1) # for dynamic reshape\n    return hx, hy, hz, ht\n\n\nclass DynEdge(GNN):\n    \"\"\"DynEdge (dynamical edge convolutional) model.\"\"\"\n\n    @save_model_config\n    def __init__(\n        self,\n        nb_inputs: int,\n        *,\n        nb_neighbours: int = 8,\n        features_subset: Optional[Union[List[int], slice]] = None,\n        dynedge_layers = None,\n        post_processing_layer_sizes: Optional[List[int]] = None,\n        post_processing_transformer: Optional[Dict] = None,\n        readout_layer_sizes: Optional[List[int]] = None,\n        global_pooling: Optional[Dict] = None,\n        add_global_variables_after_pooling: bool = False,\n        sensor_embedding = False,\n        local_pooling = None \n    ):\n        \"\"\"Construct `DynEdge`.\n\n        Args:\n            nb_inputs: Number of input features on each node.\n            nb_neighbours: Number of neighbours to used in the k-nearest\n                neighbour clustering which is performed after each (dynamical)\n                edge convolution.\n            features_subset: The subset of latent features on each node that\n                are used as metric dimensions when performing the k-nearest\n                neighbours clustering. Defaults to [0,1,2].\n            dynedge_layer_sizes: The layer sizes, or latent feature dimenions,\n                used in the `DynEdgeConv` layer. Each entry in\n                `dynedge_layer_sizes` corresponds to a single `DynEdgeConv`\n                layer; the integers in the corresponding tuple corresponds to\n                the layer sizes in the multi-layer perceptron (MLP) that is\n                applied within each `DynEdgeConv` layer. That is, a list of\n                size-two tuples means that all `DynEdgeConv` layers contain a\n                two-layer MLP.\n                Defaults to [(128, 256), (336, 256), (336, 256), (336, 256)].\n            post_processing_layer_sizes: Hidden layer sizes in the MLP\n                following the skip-concatenation of the outputs of each\n                `DynEdgeConv` layer. Defaults to [336, 256].\n            readout_layer_sizes: Hidden layer sizes in the MLP following the\n                post-processing _and_ optional global pooling. As this is the\n                last layer(s) in the model, the last layer in the read-out\n                yields the output of the `DynEdge` model. Defaults to [128,].\n            global_pooling_schemes: The list global pooling schemes to use.\n                Options are: \"min\", \"max\", \"mean\", and \"sum\".\n            add_global_variables_after_pooling: Whether to add global variables\n                after global pooling. The alternative is to  added (distribute)\n                them to the individual nodes before any convolutional\n                operations.\n        \"\"\"\n        # Latent feature subset for computing nearest neighbours in DynEdge.\n        if features_subset is None:\n            features_subset = slice(0, 3)\n\n\n        self._dynedge_layers = dynedge_layers\n        self._post_processing_layer_sizes = post_processing_layer_sizes\n        self._post_processing_transformer = post_processing_transformer\n        self._local_pooling_conf = local_pooling\n\n        # Read-out layer sizes\n        if readout_layer_sizes is None:\n            readout_layer_sizes = [\n                128,\n            ]\n\n        assert isinstance(readout_layer_sizes, list)\n        assert len(readout_layer_sizes)\n        assert all(size > 0 for size in readout_layer_sizes)\n\n        self._readout_layer_sizes = readout_layer_sizes\n\n        # Global pooling scheme(s)\n        if global_pooling is None:\n            global_pooling = {\"type\": \"simple\", \"schemes\": [\"min\",\"max\",\"mean\"], \"nb_out\": 768}\n\n        self._global_pooling_conf = global_pooling\n        self._global_pooling_model = None\n\n        self._add_global_variables_after_pooling = (\n            add_global_variables_after_pooling\n        )\n\n        # Base class constructor\n        super().__init__(nb_inputs, self._readout_layer_sizes[-1])\n\n        # Remaining member variables()\n        self._activation = torch.nn.LeakyReLU()\n        self._nb_inputs = nb_inputs\n        self._nb_global_variables = 5 + nb_inputs\n        self._nb_neighbours = nb_neighbours\n        self._features_subset = features_subset\n\n        self._sensor_embed = None\n        if sensor_embedding:            \n            sensor_embed_init_weights = torch.zeros((5160, 8), dtype=torch.float32, device=device)\n            sensor_embed_init_weights[:,:4] = 1.0\n            self._sensor_embed = torch.nn.Embedding(5160, 8, _weight=sensor_embed_init_weights)\n\n        self._construct_layers()\n        \n        self._local_pooling = None        \n        if self._local_pooling_conf:\n            if self._local_pooling_conf[\"type\"] == \"TopKPooling\":\n                self._local_pooling = TopKPooling(self._post_processing_layer_sizes[-1],self._local_pooling_conf[\"k\"])\n\n        if self._global_pooling_conf[\"type\"] == \"GRU\":\n            self._global_pooling_model = torch.nn.GRU(input_size=self._global_pooling_conf[\"nb_in\"],\n                                                      hidden_size=int(self._global_pooling_conf[\"nb_out\"]/2),\n                                                      batch_first=True,\n                                                      bidirectional=self._global_pooling_conf[\"bidirectional\"])\n\n    def build_dyn_edge_conv_layer(self, conf, nb_latent_features, last_layer):\n        sizes = conf['sizes']\n        layers = []\n        layer_sizes = [nb_latent_features] + list(sizes)\n        for ix, (nb_in, nb_out) in enumerate(\n            zip(layer_sizes[:-1], layer_sizes[1:])\n        ):\n            if ix == 0:\n                nb_in *= 2\n            layers.append(torch.nn.Linear(nb_in, nb_out))\n            layers.append(self._activation)\n        \n        conv_layer = DynEdgeConv(\n                torch.nn.Sequential(*layers),\n                aggr=\"add\",\n                nb_neighbors=self._nb_neighbours,\n                features_subset=self._features_subset,\n                last_layer=last_layer\n            )\n        return conv_layer, nb_out\n\n    def build_transformer_conv_layer(self, conf, nb_latent_features, last_layer):\n        conv_layer = DynTransformerConv(\n                in_channels=nb_latent_features,\n                out_channels=conf['nb_out'],\n                heads=conf['heads'],\n                nb_neighbors=self._nb_neighbours,\n                features_subset=self._features_subset,\n                last_layer=last_layer\n            )\n        return conv_layer, conf['nb_out']\n\n    def build_layer(self, layer_conf, nb_latent_features, last_layer):\n        if layer_conf['type'] == 'DynEdgeConv':\n            return self.build_dyn_edge_conv_layer(layer_conf, nb_latent_features, last_layer)\n        if layer_conf['type'] == 'TransformerConv':\n            return self.build_transformer_conv_layer(layer_conf, nb_latent_features, last_layer)\n\n    def _construct_layers(self) -> None:\n        \"\"\"Construct layers (torch.nn.Modules).\"\"\"\n        # Convolutional operations\n        nb_input_features = self._nb_inputs\n        if not self._add_global_variables_after_pooling:\n            nb_input_features += self._nb_global_variables\n\n        self._conv_layers = torch.nn.ModuleList()\n        nb_latent_features = nb_input_features\n        for layer_conf in self._dynedge_layers:\n            last_layer = len(self._conv_layers) == (len(self._dynedge_layers) - 1) # qudata: unnessesary last layer\n            conv_layer, nb_out = self.build_layer(layer_conf, nb_latent_features, last_layer)\n            self._conv_layers.append(conv_layer)\n            nb_latent_features = nb_out\n\n        # Post-processing operations\n        nb_latent_features = (\n            sum(layer_conf['nb_out'] for layer_conf in self._dynedge_layers)\n            + nb_input_features\n        )\n\n        post_processing_layers = []\n        if self._post_processing_layer_sizes:\n            \n            layer_sizes = [nb_latent_features] + list(\n                self._post_processing_layer_sizes\n            )\n            for nb_in, nb_out in zip(layer_sizes[:-1], layer_sizes[1:]):\n                post_processing_layers.append(torch.nn.Linear(nb_in, nb_out))\n                post_processing_layers.append(self._activation)            \n\n        if self._post_processing_transformer:\n            encoder_layer = torch.nn.TransformerEncoderLayer(\n                                                             d_model=self._post_processing_transformer[\"d_model\"], \n                                                             nhead=self._post_processing_transformer[\"nhead\"],\n                                                             dim_feedforward=self._post_processing_transformer[\"dim_feedforward\"]\n                                                            )\n            post_processing_layers.append(torch.nn.TransformerEncoder(encoder_layer, self._post_processing_transformer[\"num_layers\"]))\n\n        self._post_processing = torch.nn.Sequential(*post_processing_layers)\n\n        # Read-out operations\n        nb_latent_features = self._global_pooling_conf[\"nb_out\"]\n\n        if self._add_global_variables_after_pooling:\n            nb_latent_features += self._nb_global_variables\n\n        readout_layers = []\n        layer_sizes = [nb_latent_features] + list(self._readout_layer_sizes)\n        for nb_in, nb_out in zip(layer_sizes[:-1], layer_sizes[1:]):\n            readout_layers.append(torch.nn.Linear(nb_in, nb_out))\n            readout_layers.append(self._activation)\n\n        self._readout = torch.nn.Sequential(*readout_layers)\n\n    def _global_pooling_simple(self, x: Tensor, batch: LongTensor) -> Tensor:\n        \"\"\"Perform global pooling.\"\"\"\n        pooled = []\n        for pooling_scheme in self._global_pooling_conf[\"schemes\"]:\n            pooling_fn = GLOBAL_POOLINGS[pooling_scheme]\n            pooled_x = pooling_fn(x, index=batch, dim=0)\n            if isinstance(pooled_x, tuple) and len(pooled_x) == 2:\n                # `scatter_{min,max}`, which return also an argument, vs.\n                # `scatter_{mean,sum}`\n                pooled_x, _ = pooled_x\n            pooled.append(pooled_x)\n        pooled = torch.cat(pooled, dim=1)\n        return pooled\n\n    def _global_pooling_gru(self, x: Tensor, batch: LongTensor) -> Tensor:\n        x, mask = to_dense_batch(x, batch)\n        pooled = self._global_pooling_model(x)[0][:, -1]\n        return pooled\n\n    def _global_pooling(self, x: Tensor, batch: LongTensor) -> Tensor:\n        if self._global_pooling_conf[\"type\"] == \"simple\":\n            return self._global_pooling_simple(x, batch)\n        if self._global_pooling_conf[\"type\"] == \"GRU\":\n            return self._global_pooling_gru(x, batch)\n\n    def _calculate_global_variables(\n        self,\n        x: Tensor,\n        edge_index: LongTensor,\n        batch: LongTensor,\n        *additional_attributes: Tensor,\n    ) -> Tensor:\n        \"\"\"Calculate global variables.\"\"\"\n        # Calculate homophily (scalar variables)\n        h_x, h_y, h_z, h_t = calculate_xyzt_homophily(x, edge_index, batch)\n\n        # Calculate mean features\n        global_means = scatter_mean(x, batch, dim=0)\n\n        # Add global variables\n        global_variables = torch.cat(\n            [\n                global_means,\n                h_x,\n                h_y,\n                h_z,\n                h_t,\n            ]\n            + [attr.unsqueeze(dim=1) for attr in additional_attributes],\n            dim=1,\n        )\n\n        return global_variables\n\n    def forward(self, data: Data) -> Tensor:\n        \"\"\"Apply learnable forward pass.\"\"\"\n        if CFG.FROZEN: torch.set_grad_enabled(False) \n\n        # Convenience variables\n        x, edge_index, batch = data.x, data.edge_index, data.batch\n\n        # Sensor embeddings\n        if self._sensor_embed:    \n            #with torch.no_grad():         \n            s_emb = self._sensor_embed(data.sensor_id)\n            x[:,:4] = x[:,:4] * s_emb[:,:4] + s_emb[:,4:] # x,y,z,t * wx,wy,wz,wt + ax,ay,az,at\n\n        global_variables = self._calculate_global_variables(\n            x,\n            edge_index,\n            batch,\n            torch.log10(data.n_pulses),\n        )\n\n        # Distribute global variables out to each node\n        if not self._add_global_variables_after_pooling:\n            distribute = (\n                batch.unsqueeze(dim=1) == torch.unique(batch).unsqueeze(dim=0)\n            ).type(torch.float)\n\n            global_variables_distributed = torch.sum(\n                distribute.unsqueeze(dim=2)\n                * global_variables.unsqueeze(dim=0),\n                dim=1,\n            )\n\n            x = torch.cat((x, global_variables_distributed), dim=1)\n\n        # DynEdge-convolutions\n        skip_connections = [x]\n        for conv_layer_idx, conv_layer in enumerate(self._conv_layers): \n            if CFG.FROZEN and self.training and conv_layer_idx in CFG.UNFOROZEN_LAYERS: torch.set_grad_enabled(True)\n            x, edge_index = conv_layer(x, edge_index, batch)\n            skip_connections.append(x)\n\n        # Skip-cat\n        x = torch.cat(skip_connections, dim=1)\n\n        # Post-processing        \n        x = self._post_processing(x)\n\n        if self._local_pooling:\n            x, edge_index, _, batch, _, _ = self._local_pooling(x, edge_index, None, batch)\n\n        # (Optional) Global pooling\n        if self._global_pooling_conf:\n            x = self._global_pooling(x, batch=batch)\n            if self._add_global_variables_after_pooling:\n                x = torch.cat(\n                    [\n                        x,\n                        global_variables,\n                    ],\n                    dim=1,\n                )\n\n        if CFG.FROZEN and self.training: torch.set_grad_enabled(True) \n\n        # Read-out\n        x = self._readout(x)\n\n        return x\n\n#@title GraphNet\n\nfrom graphnet.data.constants import FEATURES, TRUTH\nfrom graphnet.models import StandardModel\nfrom graphnet.models.detector.icecube import IceCubeKaggle\n#from graphnet.models.gnn import DynEdge # QuData: owerwrite\nfrom graphnet.models.graph_builders import KNNGraphBuilder\nfrom graphnet.models.task.reconstruction import DirectionReconstructionWithKappa, ZenithReconstructionWithKappa, AzimuthReconstructionWithKappa\nfrom graphnet.training.callbacks import ProgressBar\nfrom graphnet.training.loss_functions import VonMisesFisher3DLoss, VonMisesFisher2DLoss\nfrom graphnet.training.labels import Direction\nfrom graphnet.utilities.logging import get_logger\nfrom graphnet.models.graph_builders import GraphBuilder\nfrom graphnet.models.detector.detector import Detector\n\nfrom typing import Any, Dict, List, Optional, Union\n\nimport torch\nfrom torch import Tensor\nfrom torch.nn import ModuleList\nfrom torch.optim import Adam\nfrom torch.utils.data import DataLoader\nfrom torch_geometric.data import Data\n\nfrom graphnet.models.coarsening import Coarsening\nfrom graphnet.utilities.config import save_model_config\nfrom graphnet.models.detector.detector import Detector\nfrom graphnet.models.gnn.gnn import GNN\nfrom graphnet.models.model import Model\nfrom graphnet.models.task import Task\nfrom torch.nn.functional import cosine_similarity, normalize\nfrom graphnet.training.loss_functions import LossFunction\nfrom graphnet.utilities.maths import eps_like\n\n\nclass CosineLoss(LossFunction):\n    def _forward(self, preds, target):\n        target = target.reshape(-1, 3)\n        return 1 - cosine_similarity(preds[:,:3], target, dim=1, eps=1e-8)\n\nclass CosineLoss2D(LossFunction):\n    def _forward(self, preds, target):\n        target = target.reshape(-1, 2)\n        return 1 - cosine_similarity(preds[:,:2], target, dim=1, eps=1e-8)\n\nclass VMFCustomLoss(LossFunction):\n    def _forward(self, preds, target, eps = 1e-8, kappa0=10):\n        target = target.reshape(-1, 3)          \n        kappa = preds[:,3] \n        preds = preds[:,:3]*kappa.unsqueeze(1)\n        logC  = -kappa + torch.log(kappa+eps) \n        mask  = kappa < kappa0\n        ka    = kappa[mask]\n        logC[mask] = torch.log( ka / ( torch.exp(ka) - torch.exp(-ka) + eps) )        \n        return -(preds* target).sum(axis=1) - logC\n\nclass VMFCustomLoss2(LossFunction):\n    def _forward(self, preds, target, eps = 1e-8, kappa0=10):    \n        target = target.reshape(-1, 3)      \n        kappa = preds[:,3] \n        preds = preds[:,:3]*kappa.unsqueeze(1)\n        logC  = torch.log(kappa/(1-torch.exp(-2*kappa))+eps) - kappa\n        return -( (n_true*n_pred).sum(axis=1) + logC)\n\nclass IceCubeCustom(Detector):\n    \"\"\"`Detector` class for Kaggle Competition.\"\"\"\n\n    def __init__(\n        self, graph_builder: GraphBuilder, scalers: List[dict] = None, features = None\n    ):\n        self._features = features\n\n        super().__init__(graph_builder, scalers)\n\n    @property\n    def features(self) -> List[str]:\n        return self._features\n\n    def _forward(self, data: Data) -> Data:\n        \"\"\"Ingest data, build graph, and preprocess features.\n\n        Args:\n            data: Input graph data.\n\n        Returns:\n            Connected and preprocessed graph data.\n        \"\"\"\n        # Check(s)\n        #self._validate_features(data)\n\n        # Preprocessing\n        data.x[:, 0] /= 500.0  # x\n        data.x[:, 1] /= 500.0  # y\n        data.x[:, 2] /= 500.0  # z\n        data.x[:, -3] = (data.x[:, -3] - 1.0e04) / 3.0e4  # time\n        data.x[:, -2] = torch.log10(data.x[:, -2]) / 3.0  # charge\n\n        return data\n\nclass DirectionReconstructionWithKappa2D(Task):\n    \"\"\"Reconstructs direction with kappa from the 3D-vMF distribution.\"\"\"\n\n    # Requires three features: untransformed points in (x,y,z)-space.\n    nb_inputs = 2\n\n    def _forward(self, x: Tensor) -> Tensor:\n        # Transform outputs to angle and prepare prediction\n        kappa = torch.linalg.vector_norm(x, dim=1) + eps_like(x)\n        vec_x = x[:, 0] / kappa\n        vec_y = x[:, 1] / kappa\n        return torch.stack((vec_x, vec_y, kappa), dim=1)\n\nclass DirectionReconstructionWithBins(Task):\n    \"\"\"Reconstructs direction with kappa from the 3D-vMF distribution.\"\"\"\n\n    # Requires three features: untransformed points in (x,y,z)-space.\n    nb_inputs = 576\n\n    def _forward(self, x: Tensor) -> Tensor:\n        return x\n\n\"\"\"Standard model class(es).\"\"\"\nclass StandardCustomModel(Model):\n    \"\"\"Main class for standard models in graphnet.\n\n    This class chains together the different elements of a complete GNN-based\n    model (detector read-in, GNN architecture, and task-specific read-outs).\n    \"\"\"\n\n    @save_model_config\n    def __init__(\n        self,\n        *,\n        detector: Detector,\n        gnn: GNN,\n        tasks: Union[Task, List[Task]],\n        coarsening: Optional[Coarsening] = None,\n        optimizer_class: type = Adam,\n        optimizer_kwargs: Optional[Dict] = None,\n        scheduler_class: Optional[type] = None,\n        scheduler_kwargs: Optional[Dict] = None,\n        scheduler_config: Optional[Dict] = None     \n    ) -> None:\n        \"\"\"Construct `StandardModel`.\"\"\"\n        # Base class constructor\n        super().__init__()\n\n        # Check(s)\n        if isinstance(tasks, Task):\n            tasks = [tasks]\n        assert isinstance(tasks, (list, tuple))\n        assert all(isinstance(task, Task) for task in tasks)\n        assert isinstance(detector, Detector)\n        assert isinstance(gnn, GNN)\n        assert coarsening is None or isinstance(coarsening, Coarsening)\n\n        # Member variable(s)\n        self._detector = detector\n        self._gnn = gnn\n        self._tasks = ModuleList(tasks)\n        self._coarsening = coarsening\n\n    def forward(self, data: Data) -> List[Union[Tensor, Data]]:\n        \"\"\"Forward pass, chaining model components.\"\"\"\n        if CFG.FROZEN: torch.set_grad_enabled(False) \n        if self._coarsening:\n            data = self._coarsening(data)\n        data = self._detector(data)\n        x = self._gnn(data)\n        preds = [task(x) for task in self._tasks]\n        return preds\n\n    def compute_loss(\n        self, preds: Tensor, data: Data, verbose: bool = False\n    ) -> Tensor:\n        \"\"\"Compute and sum losses across tasks.\"\"\"\n        losses = [\n            task.compute_loss(pred, data)\n            for task, pred in zip(self._tasks, preds)\n        ]\n        if verbose:\n            self.info(f\"{losses}\")\n        assert all(\n            loss.dim() == 0 for loss in losses\n        ), \"Please reduce loss for each task separately\"\n        return torch.sum(torch.stack(losses))\n\n    def _get_batch_size(self, data: Data) -> int:\n        return torch.numel(torch.unique(data.batch))\n\ndef build_model(config: Dict[str,Any]) -> StandardModel:\n    \"\"\"Builds GNN from config\"\"\"\n    # Building model\n\n    if not \"mode\" in config:\n        config[\"mode\"] = \"regression\"\n\n    len_train_dataloader = 1000 # len(train_dataloader)\n\n    detector = IceCubeCustom(\n        graph_builder=KNNGraphBuilder(\n              nb_nearest_neighbours=config[\"neighbours\"][0],\n              columns=config[\"features_subset\"]\n            ),\n        features=config[\"features\"]\n    )    \n    gnn = DynEdge(\n        #nb_inputs=detector.nb_outputs,\n        nb_inputs=len(config[\"features\"]),\n        nb_neighbours=config[\"neighbours\"][1],\n        dynedge_layers=config[\"dynedge_layers\"],\n        post_processing_layer_sizes=config[\"post_processing_layer_sizes\"],\n        post_processing_transformer=config[\"post_processing_transformer\"],\n        readout_layer_sizes=config[\"readout_layer_sizes\"],\n        global_pooling=config.get(\"global_pooling\"),\n        features_subset=config[\"features_subset\"],\n        sensor_embedding=config.get(\"sensor_embedding\", False),\n        local_pooling=config.get(\"local_pooling\"),\n    )\n\n    loss_function = None\n    task = None\n    \n    if config[\"loss_function\"] == \"VonMisesFisher3DLoss\":\n        loss_function = VonMisesFisher3DLoss()\n\n    if config[\"loss_function\"] == \"CosineLoss\":\n        loss_function = CosineLoss()\n\n    if config[\"loss_function\"] == \"VMFCustomLoss\":\n        loss_function = VMFCustomLoss()\n\n    if config[\"loss_function\"] == \"CosineLoss2D\":\n        loss_function = CosineLoss2D()\n\n    if config[\"loss_function\"] == \"CrossEntropyLoss\":\n        loss_function = torch.nn.CrossEntropyLoss()        \n\n    if config[\"target\"] == 'direction':\n        if CFG.ANGLES == 'az,ze':\n            task = DirectionReconstructionWithKappa(\n                    hidden_size=gnn.nb_outputs,\n                    target_labels=config[\"target\"],\n                    loss_function=loss_function,       \n                )\n            prediction_columns = [config[\"target\"] + \"_x\", \n                                  config[\"target\"] + \"_y\", \n                                  config[\"target\"] + \"_z\", \n                                  config[\"target\"] + \"_kappa\" ]\n            additional_attributes = ['zenith', 'azimuth', 'event_id', 'sensor_id']\n        else:\n            task = DirectionReconstructionWithKappa2D(\n                hidden_size=gnn.nb_outputs,\n                target_labels=config[\"target\"],\n                loss_function=loss_function,       \n            )\n            prediction_columns = [config[\"target\"] + \"_x\", \n                                  config[\"target\"] + \"_y\", \n                                  config[\"target\"] + \"_kappa\" ]\n            additional_attributes = ['zenith', 'azimuth', 'event_id', 'sensor_id']\n\n    if config[\"target\"] == 'class_id':\n        DirectionReconstructionWithBins.nb_inputs = config[\"bin_num\"]**2\n        task = DirectionReconstructionWithBins(\n                                hidden_size=gnn.nb_outputs,\n                                target_labels=config[\"target\"],\n                                loss_function=loss_function)\n        \n        prediction_columns = [] \n        additional_attributes = ['zenith', 'azimuth', 'event_id', 'sensor_id']\n\n\n    model = StandardCustomModel(\n        detector=detector,\n        gnn=gnn,\n        tasks=[task],\n    )\n    model.prediction_columns = prediction_columns\n    model.additional_attributes = additional_attributes\n    \n    return model","metadata":{"execution":{"iopub.status.busy":"2023-04-18T16:55:55.886373Z","iopub.execute_input":"2023-04-18T16:55:55.886663Z","iopub.status.idle":"2023-04-18T16:56:00.574098Z","shell.execute_reply.started":"2023-04-18T16:55:55.886635Z","shell.execute_reply":"2023-04-18T16:56:00.572814Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Preprocess","metadata":{}},{"cell_type":"code","source":"%%time\n\nif CFG.MODE in ['test','submit']:\n    CFG.BATCH_RANGE = split_meta_data()\n","metadata":{"execution":{"iopub.status.busy":"2023-04-18T16:56:00.579488Z","iopub.execute_input":"2023-04-18T16:56:00.580624Z","iopub.status.idle":"2023-04-18T16:56:00.625107Z","shell.execute_reply.started":"2023-04-18T16:56:00.580579Z","shell.execute_reply":"2023-04-18T16:56:00.624044Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"# Constants\n\nCFG.FEATURES = ['x', 'y', 'z', 'time', 'charge', 'auxiliary'] # FEATURES.KAGGLE\nCFG.NUM_EPOCHS = 10000\nCFG.NUM_STEPS = -1\n\nbin_num = None\nsensors_data = None\nazimuth_edges = None\nzenith_edges = None\nangle_bin_vector = None\n\ndef preds_to_vectors(preds):\n    preds = preds[0].detach()\n\n    if CFG.TRAIN_CFG['mode'] == 'regression':\n        vectors_pred = preds[:,:3]\n        kappa_preds =  preds[:,3]\n        return vectors_pred, kappa_preds\n\n    if CFG.TRAIN_CFG['mode'] == 'classification':\n        preds = torch.nn.functional.softmax(preds, dim=1)\n        vectors_pred_avg = code_to_vector(preds, bin_num, angle_bin_vector, max=False) \n        #vectors_pred_max = code_to_vector(preds, bin_num, angle_bin_vector, max=True) \n        return vectors_pred_avg, None\n    \ndef fix_cfg(cfg):\n    if not 'mode' in cfg:\n        cfg['mode'] = 'regression'\n    return cfg\n\ndef load_model(model_name):\n    global sensors_data\n    print('load model: ', model_name)\n    model_weights = f\"/kaggle/input/icecube-weights/{model_name}.pt\"\n    resume_state = torch.load(model_weights)\n    CFG.TRAIN_CFG = fix_cfg(resume_state[\"cfg\"][\"TRAIN_CFG\"])\n    sensors_data = get_sensors(CFG.TRAIN_CFG['features'][:-3])\n    model = build_model(config=CFG.TRAIN_CFG)\n    model = model.to(device)\n    model.load_state_dict(resume_state['model_state_dict'])\n    return model\n\ndef infer_model(model_name):\n    global azimuth_edges, zenith_edges, angle_bin_vector, bin_num\n    model = load_model(model_name)\n    model.eval()\n    \n    if CFG.TRAIN_CFG['mode'] == 'classification':\n        bin_num = CFG.TRAIN_CFG['bin_num']\n        azimuth_edges, zenith_edges = build_az_ze_edges(bin_num)\n        angle_bin_vector = build_angle_bin_vector(azimuth_edges.cpu().numpy(), zenith_edges.cpu().numpy(), bin_num)\n\n    ens_event_ids   = []\n    ens_angles_true = []\n    ens_angles_pred = []\n    ens_kappa_pred  = []\n\n    for az_angle in CFG.TTA_AZ_ANGLES:\n        print('tta azimuth angle:', az_angle)\n        rot_angle = torch.tensor(az_angle, device=device, dtype=torch.float32)\n        rot_phi = rot_angle * torch.pi / 180.0\n        rot_matrix_direct  = get_rot_matrix(rot_phi)\n        rot_matrix_reverse = get_rot_matrix(-rot_phi)\n\n        all_event_ids       = torch.tensor([], dtype=torch.long, device=device)\n        all_vectors_pred    = torch.tensor([], dtype=torch.float32, device=device)\n        all_vectors_target  = torch.tensor([], dtype=torch.float32, device=device)\n        all_kappa_pred      = torch.tensor([], dtype=torch.float32, device=device)\n\n        for batch_id in range(CFG.BATCH_RANGE[0], CFG.BATCH_RANGE[1]):\n            test_loader = build_dataloader(batch_id, False, CFG.TRAIN_CFG)\n            pbar = tqdm(enumerate(test_loader), total=len(test_loader))\n            for steps, batch in pbar: \n                pbar.set_description(f\"batch: {batch_id}\")\n                batch = batch.to(device)\n                batch.x[:,:2] = batch.x[:,:2] @ rot_matrix_direct\n\n                # Forward propagation\n                with torch.no_grad():\n                    preds = model(batch)\n                    \n                vectors_pred, kappa_pred = preds_to_vectors(preds)\n\n                vectors_target = batch.direction\n\n                all_vectors_pred   = torch.cat([all_vectors_pred, vectors_pred])\n                all_vectors_target = torch.cat([all_vectors_target, vectors_target])\n                all_event_ids      = torch.cat([all_event_ids, batch.event_ids])\n                if not kappa_pred is None:\n                    all_kappa_pred     = torch.cat([all_kappa_pred, kappa_pred])\n\n        all_vectors_pred[:,:2] = all_vectors_pred[:,:2] @ rot_matrix_reverse\n        all_angles_pred  = vectors_to_angles(all_vectors_pred)\n        all_angles_true  = vectors_to_angles(all_vectors_target)\n\n        ens_event_ids.append(all_event_ids)\n        ens_angles_true.append(all_angles_true)\n        ens_angles_pred.append(all_angles_pred)\n        ens_kappa_pred.append(all_kappa_pred)\n\n    # ensemble TTA\n    vectors_ens = None\n\n    for i in range(len(CFG.TTA_AZ_ANGLES)):\n        vectors_preds   = angles_to_vectors(ens_angles_pred[i])\n        if vectors_ens is None:\n            vectors_ens = vectors_preds\n        else:\n            vectors_ens += vectors_preds\n\n    gnn_event_ids = ens_event_ids[0]\n    gnn_angles_pred = vectors_to_angles(vectors_ens)\n    gnn_kappa_pred = ens_kappa_pred[0]\n\n    gnn_df = pd.DataFrame({\n                        'event_id': gnn_event_ids.cpu().numpy(),\n                        'azimuth':  gnn_angles_pred[:,0].cpu().numpy(),\n                        'zenith':   gnn_angles_pred[:,1].cpu().numpy(),\n                      })   \n    gnn_df.to_csv(f'submission_{model_name}.csv', index=False)\n\n    if CFG.MODE == 'test':\n        single_score, _ = angular_dist_score(ens_angles_true[0], ens_angles_pred[0])\n        ens_score, _    = angular_dist_score(ens_angles_true[0], gnn_angles_pred)\n        #gnn_kappa_pred = gnn_kappa_pred.cpu().numpy()\n        #if len(gnn_kappa_pred)>0:\n        #    gnn_df['kappa'] = gnn_kappa_pred        \n        #gnn_df.to_csv(f'submission_{model_name}_batch_{CFG.BATCH_RANGE[0]}_{(CFG.BATCH_RANGE[1]-1)}.csv', index=False)        \n        print(f'GNN({model_name}) single score:', single_score.cpu().item())\n        print(f'GNN({model_name}) ensemble score:', ens_score.cpu().item())\n                        \n    return gnn_event_ids, gnn_angles_pred\n","metadata":{"execution":{"iopub.status.busy":"2023-04-18T16:56:00.630344Z","iopub.execute_input":"2023-04-18T16:56:00.630713Z","iopub.status.idle":"2023-04-18T16:56:00.665996Z","shell.execute_reply.started":"2023-04-18T16:56:00.630675Z","shell.execute_reply":"2023-04-18T16:56:00.664658Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Infer","metadata":{}},{"cell_type":"code","source":"%%time\nfor model_name in CFG.MODELS:\n    gnn_event_ids, gnn_angles_pred = infer_model(model_name)\n    ","metadata":{"execution":{"iopub.status.busy":"2023-04-18T16:56:00.671512Z","iopub.execute_input":"2023-04-18T16:56:00.671900Z","iopub.status.idle":"2023-04-18T16:56:04.034190Z","shell.execute_reply.started":"2023-04-18T16:56:00.671864Z","shell.execute_reply":"2023-04-18T16:56:04.032425Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"- att-1: 0.9976\n- att-2: 1.0016\n- gnn_id_26_val_0_9926_exp_classific_16n_e1633: 0.99468\n- gnn_id_27_val_0_9961_exp_readout_emb_e196:    0.99710\n- gnn_id_28_val_0_9919_exp_classific_24n_e726:  0.99309\n","metadata":{}},{"cell_type":"code","source":"#!cp validate/submission* .","metadata":{"execution":{"iopub.status.busy":"2023-04-18T16:56:04.035577Z","iopub.execute_input":"2023-04-18T16:56:04.036042Z","iopub.status.idle":"2023-04-18T16:56:04.042288Z","shell.execute_reply.started":"2023-04-18T16:56:04.036003Z","shell.execute_reply":"2023-04-18T16:56:04.040957Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Ensemble","metadata":{}},{"cell_type":"markdown","source":"## Config","metadata":{}},{"cell_type":"code","source":"import torch\n\nclass CFG:\n    TEST_MODE = GLOBAL_TEST_MODE # GLOBAL_TEST_MODE\n    WEIGHTS   = '/kaggle/input/icecube-models/ens_err_0.9792_mlp512_aggF_gnn_1_2_3_att_1_2_4.pt'\n    MODELS    = [\n                # GNN new best:\n                'submission_gnn_id_26_val_0_9926_exp_classific_16n_e1633.csv',\n                'submission_gnn_id_28_val_0_9919_exp_classific_24n_e726.csv',\n                'submission_gnn_id_27_val_0_9961_exp_readout_emb_e196.csv',\n                #'submission_gnn_id_20_val_0_9964_exp_mlp_2048_e1343.csv',\n                # ATT: \n                'submission-att-1.csv',  # submission_att_all_0.9984_L12_batch_1_40.csv\n                'submission-att-2.csv',  # submission_att_rnn_1.0015_L10_batch_1_40.csv\n                'submission-att-3.csv',  # submission_att04_1.0003_D00_batch_1_40.csv\n              ]  \n    TRUE_TARGET = \"/kaggle/input/icecube-events/true_batch_1_40.parquet\"\n    \n    loss   = 'cos'\n    ka_reg = 0    \n    device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")","metadata":{"execution":{"iopub.status.busy":"2023-04-18T17:04:35.959011Z","iopub.execute_input":"2023-04-18T17:04:35.959408Z","iopub.status.idle":"2023-04-18T17:04:35.966536Z","shell.execute_reply.started":"2023-04-18T17:04:35.959376Z","shell.execute_reply":"2023-04-18T17:04:35.965463Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Lib","metadata":{}},{"cell_type":"code","source":"!cp -r /kaggle/input/nnet-lib/nnet /kaggle/working\n\nfrom nnet.models   import Transformer, TransformerBlock, MLP\nfrom nnet.data     import Data\nfrom nnet.optim    import CosScheduler, ExpScheduler\nfrom nnet.trainer  import Trainer\n\nimport torch, torch.nn as nn\nimport pandas as pd\nimport numpy as np\nfrom   tqdm.auto import tqdm\nimport os, gc, sys, time, datetime, math, random, copy, psutil\n\ndef metrics(x, Y, eps=1e-8):\n    kappa = torch.norm(x, dim=1, keepdim=True).clip(eps)\n    y = x / kappa \n    cos  = (y*Y).sum(dim=1) \n    if   CFG.loss == 'cos':\n        loss = 1 - cos.mean() + CFG.ka_reg * (kappa**2).mean()\n    elif CFG.loss == 'prod':\n        loss = -((x*Y).sum(dim=1)).mean() + CFG.ka_reg * kappa.mean()\n    elif CFG.loss == 'vMF':\n        logC = -kappa + torch.log( ( kappa+eps )/( 1-torch.exp(-2*kappa)+2*eps ) )\n        loss = -( (x*Y).sum(dim=1) + logC ).mean() \n    elif CFG.loss == 'k2':             \n        loss = -((x*Y).sum(dim=1)).mean() + 0.5 * (kappa**2).mean()\n    elif CFG.loss == 'k2ze':\n        loss = -((x*Y).sum(dim=1)).mean() + 0.5 * (kappa**2).mean() + torch.square(y[:,2]-Y[:,2]).mean()\n\n    err  = torch.abs( torch.arccos(  torch.clip(cos.detach() ,-1,1) ) )\n    return loss,  y.detach(),  torch.cat([err.view(-1,1), kappa.detach().view(-1,1)], dim=1 )\n\ndef angles2vector(df):\n    \"\"\" Add unit vector components from (azimuth,zenith) to the DataFrame df \"\"\"\n    df['nx'] = np.sin(df.zenith) * np.cos(df.azimuth)\n    df['ny'] = np.sin(df.zenith) * np.sin(df.azimuth)\n    df['nz'] = np.cos(df.zenith) \n    return df\n\n#-------------------------------------------------------------------------------\n\ndef vector2angles(n, eps=1e-8):\n    \"\"\"  Get spherical angles of vector n: (B,3) \"\"\"                \n    n = n / (np.linalg.norm(n, axis=1, keepdims=True) + eps)    \n                                \n    azimuth = np.arctan2( n[:,1],  n[:,0])    \n    azimuth[azimuth < 0] += 2*np.pi\n                                \n    zenith = np.arccos( n[:,2].clip(-1,1) )                                \n    \n    return azimuth, zenith\n#-------------------------------------------------------------------------------    \n\nclass MLP_Model(nn.Module):\n    \"\"\"\n    Все выходы моделей, возможно, дополненные агрегированными фичами события,\n    поступают на вход MLP\n    \"\"\"\n\n    def __init__(self, cfg: dict):        \n        super().__init__() \n        self.cfg = {\n            'name':      'mlp',\n            'n_models':  2,\n            'is_agg':    True,        \n            'AF':        4,                        # число агрегированных фич\n            'hidden':    256,\n            'drop':      0.01,\n        }   \n        if type(cfg) is dict:                       # добавляем, меняем свойства\n            self.cfg.update(copy.deepcopy(cfg))\n        cfg = self.cfg\n\n        F = cfg['n_models']*3\n        if cfg['is_agg']:  F += cfg['AF']\n        self.mlp = MLP( dict(input=F, hidden=cfg['hidden'], output=3, drop=cfg['drop']) )        \n\n    def forward(self, batch, eps=1e-8):      \n        x,  agg, Y = batch\n        B,T,E = x.shape\n        x = x.view(B, T*E)     \n\n        if self.cfg['is_agg']:\n            x = torch.cat([x, agg], dim=-1)\n    \n        x = self.mlp(x)\n\n        return metrics(x,Y)","metadata":{"execution":{"iopub.status.busy":"2023-04-18T17:31:20.393429Z","iopub.execute_input":"2023-04-18T17:31:20.394189Z","iopub.status.idle":"2023-04-18T17:31:21.552195Z","shell.execute_reply.started":"2023-04-18T17:31:20.394149Z","shell.execute_reply":"2023-04-18T17:31:21.550608Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"event_ids = pd.read_csv(CFG.MODELS[0])['event_id'].values\n\nif CFG.TEST_MODE:\n    true_df = pd.read_parquet(CFG.TRUE_TARGET)[:len(event_ids)]\n    models_df = true_df[['event_id','nx','ny','nz']].copy()\n    models_df.rename(columns={\"nx\": \"nx_true\", \"ny\": \"ny_true\", \"nz\": \"nz_true\"}, inplace=True)\nelse:\n    models_df = pd.DataFrame({\n                          'event_id': event_ids,\n                          'nx_true': np.zeros((len(event_ids))), \n                          'ny_true': np.zeros((len(event_ids))), \n                          'nz_true': np.ones((len(event_ids)))})\n\nfor i, m in tqdm(enumerate(CFG.MODELS), total=len(CFG.MODELS)):\n    df = pd.read_csv(m)    \n    df = angles2vector(df)[['event_id', 'nx','ny','nz']].copy()    \n    df.rename(columns={\"nx\": f\"nx_{i+1:0d}\", \"ny\": f\"ny_{i+1:0d}\", \"nz\": f\"nz_{i+1:0d}\"}, inplace=True)\n    models_df = models_df.merge(df, left_on='event_id', right_on='event_id', how='left')\ndel df\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Infer","metadata":{}},{"cell_type":"code","source":"def load_ensemble_model(fname):\n    state = torch.load(fname)  \n    #state['config']['n_models'] = 6\n    print(state['config'])    \n    model = MLP_Model(state['config'])    \n    model.load_state_dict(state['model'])     \n    return model, state['data']\n\n# Загружаем:\nmodel, data = load_ensemble_model(CFG.WEIGHTS)\nmodel.to(CFG.device)\n\ntrainer = Trainer(model, None, None)\ntrainer.plotter.plot(trainer.cfg, model, data, w=12, h=4)\n\n# Создаём датасет:\nB, T, E = len(models_df), len(CFG.MODELS), 3\nX = torch.Tensor(models_df.iloc[:, 4:].to_numpy()).view(B,T,E)   # (B,T,E)\nY = torch.Tensor(models_df.iloc[:,1:4].to_numpy())               # (B,E)\ndata_tst = Data( (X, Y,  Y), batch_size=1024,  device=CFG.device, shuffle=False, whole_batch=False)\n\noutput, losses, score = trainer.predict(model, data_tst, verbose=True)\n\n#print()\n#print(f\"output: {output.shape}  err:{score.mean(0)[0]:.4f}   kappa:{score.mean(0)[1]:.4f}\")","metadata":{"execution":{"iopub.status.busy":"2023-04-18T17:07:59.083683Z","iopub.execute_input":"2023-04-18T17:07:59.084191Z","iopub.status.idle":"2023-04-18T17:08:01.108967Z","shell.execute_reply.started":"2023-04-18T17:07:59.084147Z","shell.execute_reply":"2023-04-18T17:08:01.107846Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Submission","metadata":{}},{"cell_type":"code","source":"azimuth, zenith = vector2angles(output.cpu().numpy())\ndf = pd.DataFrame({\n                        'event_id': event_ids,\n                        'azimuth':  azimuth,\n                        'zenith': zenith\n                  })\n\ndf.to_csv('submission.csv', index=False)\n\n!head submission.csv","metadata":{"execution":{"iopub.status.busy":"2023-04-18T17:08:11.509258Z","iopub.execute_input":"2023-04-18T17:08:11.510337Z","iopub.status.idle":"2023-04-18T17:08:12.793397Z","shell.execute_reply.started":"2023-04-18T17:08:11.510287Z","shell.execute_reply":"2023-04-18T17:08:12.792149Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}