{"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":"Submission merging various models.  The models are pytorch models stored in the attached dataset\nthey use 3 or 4 letter model types as follows:\n1. dyno - Dynamic GNN with DynEdge convolutions (thanks @rasmusrse, @anjum48)\n2. dog - LSTM with W width and k layers, described as LSTM\\<W\\>x\\<k\\>\n3. egg - LSTM with an attention layer\n4. fir - LSTM with attention, outputs of LSTMs are summed rather than concatentated\n5. goo - LSTM using one hot encoding of sensor types (REG, VETO, DEEP) and added pulse counts\n    as a feature, not used in final solution\n\nThe variables use_models and weights contain the actual models and weights that are used for ensembling.\n\nEach model file contains the parameters for the model.  Because I had trouble getting jit storage to work with all models, I need the model code (in dog_net.py, egg_net.py etc) to be loaded first, then the model parameters are loaded.\n\nThis generates a cache of data on the first model for each batch_id,\nthe reuses on each model after the first one.\n\nSee Discussion post (12th place solution) for details on how the models work and were trained.\n","metadata":{}},{"cell_type":"code","source":"KAGGLE=True","metadata":{"execution":{"iopub.status.busy":"2023-04-19T16:04:46.248525Z","iopub.execute_input":"2023-04-19T16:04:46.249101Z","iopub.status.idle":"2023-04-19T16:04:46.254602Z","shell.execute_reply.started":"2023-04-19T16:04:46.249048Z","shell.execute_reply":"2023-04-19T16:04:46.253440Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Move software to working disk and install the dependencies needed for pytorch-geometric\nimport time\nstart=time.time()\n!rm  -r software\n!scp -r /kaggle/input/graphnet-and-dependencies/software .\nprint(f'{time.time()-start:8.3f} copy')\n# Install dependencies\n!pip install /kaggle/working/software/dependencies/torch-1.11.0+cu115-cp37-cp37m-linux_x86_64.whl\n!pip install /kaggle/working/software/dependencies/torch_cluster-1.6.0-cp37-cp37m-linux_x86_64.whl\n!pip install /kaggle/working/software/dependencies/torch_scatter-2.0.9-cp37-cp37m-linux_x86_64.whl\n!pip install /kaggle/working/software/dependencies/torch_sparse-0.6.13-cp37-cp37m-linux_x86_64.whl\n!pip install /kaggle/working/software/dependencies/torch_geometric-2.0.4.tar.gz\nprint(f'{time.time()-start:8.3f} install')\n","metadata":{"execution":{"iopub.status.busy":"2023-04-19T16:04:46.261535Z","iopub.execute_input":"2023-04-19T16:04:46.262243Z","iopub.status.idle":"2023-04-19T16:09:32.421935Z","shell.execute_reply.started":"2023-04-19T16:04:46.262204Z","shell.execute_reply":"2023-04-19T16:09:32.420817Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\n#Find the other modules we need\nsys.path.append('/kaggle/input/neutrinofiles')\nsys.path","metadata":{"execution":{"iopub.status.busy":"2023-04-19T16:09:32.423420Z","iopub.execute_input":"2023-04-19T16:09:32.425003Z","iopub.status.idle":"2023-04-19T16:09:32.442708Z","shell.execute_reply.started":"2023-04-19T16:09:32.424957Z","shell.execute_reply":"2023-04-19T16:09:32.441670Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport sys\nimport os\nimport pandas as pd\nimport numpy as np\nimport glob\nimport time\nimport math\nimport gc\nimport random\nimport multiprocessing\nimport copy\n#torch\nimport torch\nimport torch_geometric as geometric\nimport torch.nn as nn\nprint(torch.__version__)\n#from typing import List, Optional, Sequence, Tuple, Union, Callable, Any\n#from torch import Tensor, LongTensor\n#from torch_scatter import scatter_max, scatter_mean, scatter_min, scatter_sum\n#from torch_geometric.typing import Adj\n#give these names to remember\nimport dog_net as dog\nimport dyno_net as dyno\nimport egg_net as egg\nimport fir_net as fir\nimport dynedge  #for Net_dyno\nimport utils\nthe_net_modules={'dog':dog, 'dyno':dyno, 'egg':egg, 'fir': fir}\n","metadata":{"execution":{"iopub.status.busy":"2023-04-19T16:09:32.445594Z","iopub.execute_input":"2023-04-19T16:09:32.447720Z","iopub.status.idle":"2023-04-19T16:09:44.984779Z","shell.execute_reply.started":"2023-04-19T16:09:32.447681Z","shell.execute_reply":"2023-04-19T16:09:44.983676Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_names = [\n                 #('dog-css-state-528-s997-obj.pth',  'LSTM384x3',    256),   #\n                 #('dog-css-state-528-s997-obj.pth',  'LSTM384x3',    320),   #\n                # ('dog-css-state-528-s997-obj.pth',  'LSTM384x3',    384),   #\n                #('dog-fey-state-768-s996-obj.pth',  'LSTM 512',      1),    #\n                #('dog-fey-state-820-s996-obj.pth',   'LSTM 512',     128),  #\n                #('dog-gpn-state-636-s996-obj.pth',  'LSTM512',      128),   #\n                # ('dog-gpn-state-636-s996-obj.pth',  'LSTM512',      300),   #\n                # ('dog-ihv-state-664-s1015-obj.pth',  'GRU512x3',    128),   #\n                # ('dog-ihv-state-664-s1015-obj.pth',  'GRU512x3',    200),   #\n                # ('dog-ihv-state-664-s1015-obj.pth',  'GRU512x3',    256),   #\n                #('dog-kax-state-592-s1025-obj.pth',  'GRUx3',        128),   #\n                #('dog-ksv-state-804-s1005-obj.pth',  'LSTM 384',     128),   #\n                #('dog-ktd-state-442-s1009-obj.pth',  'LSTM 256',     128),   #\n                #('dog-nrw-state-248-s1037-obj.pth',  'dog ???',      128),\n                #('dog-nrw-state-500-s1008-obj.pth',  'LSTM 256',     128),   #\n                #('dog-nrw-state-596-s1008-obj.pth',  'LSTM 256',     128),   #\n                #('dog-osy-state-532-s999-obj.pth',   'LSTM 512',     128),   #\n                #('dog-osy-state-724-s996-obj.pth',  'LSTM 512',      1),     #\n                # ('dog-qnr-state-664-s982-obj.pth',  'LSTM512x3-384',256),    #\n                 ('dog-qnr-state-664-s982-obj.pth',  'LSTM512x3-384',384),    #\n                #('dog-rhh-state-340-obj.pth',        'dog_net ',    0),\n                #('dog-rrt-state-420-s1027-obj.pth',  'GRUx3 better', 128),\n                #('dog-sta-state-156-s1033-obj.pth', 'GRUx3',        0),\n                # ('dog-svp-state-544-s993-obj.pth',  'LSTM768x3-200', 200),  #\n                # ('dog-ust-state-640-s991-obj.pth',  'LSTM512x3',    256),   #\n                 ('dog-ust-state-752-s990-obj.pth',  'LSTM512x3',    128),   #\n                 ('dog-ust-state-752-s990-obj.pth',  'LSTM512x3',    200),   #\n                 ('dog-ust-state-752-s990-obj.pth',  'LSTM512x3',    256),   #\n                 ('dog-ust-state-752-s990-obj.pth',  'LSTM512x3',    320),   #\n                 ('dog-ust-state-752-s990-obj.pth',  'LSTM512x3',    384),   #\n    \n                 #('dog-uva-state-648-s986-obj.pth',   'LSTM512x3-200',200),  #\n                 ('dog-uva-state-648-s986-obj.pth',   'LSTM512x3-200',320),  #\n    \n                 #('dog-uva-state-768-s984-obj.pth',   'LSTM512x3-200',128),  #\n                 ('dog-uva-state-768-s984-obj.pth',   'LSTM512x3-200',200),  #\n                 ('dog-uva-state-768-s984-obj.pth',   'LSTM512x3-200',256),  #\n                 ('dog-uva-state-768-s984-obj.pth',   'LSTM512x3-200',320),  #\n                 ('dog-uva-state-768-s984-obj.pth',   'LSTM512x3-200',384),  #\n                #('dog-vgy-state-654-s1007-obj.pth',  'dog WD1e-6',   128),   #\n                 ('dyno-dxu-state-456-s998-obj.pth', 'GNN5-RO',     200),   #\n                 ('dyno-dxu-state-456-s998-obj.pth', 'GNN5-RO',     256),   #\n                 ('dyno-dxu-state-456-s998-obj.pth', 'GNN5-RO',     350),  #\n                #('dyno-igl-state-340-obj.pth',       'GNN4',        200),   #\n                 ('dyno-kxc-state-524-s997-obj.pth', 'GNN5-RO-350', 350),  #\n                 ('dyno-kxc-state-580-s996-obj.pth', 'GNN5-RO-350', 350),  #\n                # ('dyno-nyc-state-344-s1000-obj.pth', 'GNN5',       200),   #\n                 ('dyno-rjl-state-708-s999-obj.pth', 'GNN4-RO-350',  200),   #\n                 ('dyno-rjl-state-708-s999-obj.pth', 'GNN4-RO-350',  256),   #\n                 ('dyno-rjl-state-708-s999-obj.pth', 'GNN4-RO-350',  350),   #\n                 ('dyno-wnq-state-476-s1002-obj.pth', 'GNN4',       200),   #\n                 #('dyno-wpk-state-436-s1003-obj.pth','GNN4-RO',     256),   #\n                 ('egg-jhn-state-604-s994-obj.pth',  'LSTM384x3-ATT',     128),  #\n                 ('egg-jhn-state-604-s994-obj.pth',  'LSTM384x3-ATT',     200),  #\n                 ('egg-jhn-state-604-s994-obj.pth',  'LSTM384x3-ATT',     350),\n    \n                 #('egg-jqn-state-572-s984-obj.pth',  'LSTM512x3-ATT200',  128),\n                 ('egg-jqn-state-572-s984-obj.pth',  'LSTM512x3-ATT200',  200),\n                 #('egg-jqn-state-572-s984-obj.pth',  'LSTM512x3-ATT200',  320),\n                 ('egg-jqn-state-572-s984-obj.pth',  'LSTM512x3-ATT200',  384),\n    \n                 ('fir-irl-state-456-s987-obj.pth',  'LSTM512x3-200-mod', 200),\n                 #('fir-irl-state-456-s987-obj.pth',  'LSTM512x3-200-mod', 256),\n                 ('fir-irl-state-456-s987-obj.pth',  'LSTM512x3-200-mod', 300),\n                 ('fir-irl-state-456-s987-obj.pth',  'LSTM512x3-200-mod', 350),\n                 #('egg-jqn-state-348-s1005-obj.pth', 'LSTM512x3-ATT200',  256), #doesn't help\n                 ('egg-kvu-state-516-s981-obj.pth',  'LSTM512x3-ATT384',  384),\n                 ('egg-kvu-state-548-s980-obj.pth',  'LSTM512x3-ATT384',  384),\n    \n                 ('goo-uel-state-652-s988-obj.pth',  'GOO-200',           200),\n                 ('goo-uel-state-652-s988-obj.pth',  'GOO-200',           350),\n                 ('goo-wqz-state-664-s985-obj.pth',  'GOO-384',           384),\n    \n                 ('dyno-fdk-state-396-s1005-obj.pth', 'GNN6-200',         200),\n                 ('dyno-rqy-state-408-s1005-obj.pth', 'GNN6-350',         350),\n                 ('egg-fsh-state-620-s982-obj.pth',   'LSTM512x3-ATT384', 384),\n                 ('egg-fsh-state-640-s981-obj.pth',   'LSTM512x3-ATT384', 384)\n\n               ]\n\nprint('loaded')","metadata":{"execution":{"iopub.status.busy":"2023-04-19T16:09:44.988504Z","iopub.execute_input":"2023-04-19T16:09:44.988874Z","iopub.status.idle":"2023-04-19T16:09:45.003847Z","shell.execute_reply.started":"2023-04-19T16:09:44.988838Z","shell.execute_reply":"2023-04-19T16:09:45.001529Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MODE='train'   # or train\n#MODE='test'\nPATH_DATASET = \"/kaggle/input/icecube-neutrinos-in-deep-ice\"\nFULL=True      #do all batches or not\nBATCH_SIZE=256\nSAMPLES=4    # samples per event with large number of pulses\nMODEL_DIR='/kaggle/input/threemodels-v1'\ncache = {}\n#model_files = ['dog-rhh-state-340-obj.pth', 'dyno-igl-state-340-obj.pth', 'dyno-xyg-state-290-obj.pth']\n#These are all the models we might use, but we select a subset of them\n\n#use_models =[0, 7, 10, 11, 13, 14, 18] # Use these ones - we need a dog model first for caching!\n#use_models = [36, 28, 40, 43, 47, 48]\n#use_models = [3, 5, 9, 15, 18, 20, 25, 29]\nuse_models = [5, 9, 18, 20, 25, 36]  \n\nuse_model_names = [model_names[i] for i in use_models]\n#weights = [.28677, .46420, .25810]\n#These are the weights for the use_models indexes\n#weights = [0.31882195, 0.16722909, 0.12814019, 0.14714945, 0.24440681]\n#weights = [0.28649896, 0.06939899, 0.09794364, 0.17935315, 0.08443963,0.18615417, 0.12425975]\n#weights = [0.33073873, 0.12516523, 0.16632225, 0.1105064 , 0.09551677, 0.13892529, 0.25718829]\n#weights = [0.56121643, 0.62025675, 0.16434306, 0.13951461, 0.158012, 0.13922811]\n#weights = [0.19673503, 0.16515622, 0.24888259, 0.09604258, 0.07863544,\n#      0.09493275, 0.17866872, 0.28162951]\nweights = [0.30260045, 0.22815514, 0.09216829, 0.07076734, 0.18211481, 0.40845271]\n\nassert len(weights)==len(use_model_names)\nwhile 'dyno' in use_model_names[0][0]:\n    #we need to shift model names and weights so dog is first\n    use_model_names = use_model_names[1:] + use_model_names[:1]\n    weights = weights[1:] + weights[:1]\nassert 'dog' in use_model_names[0][0] or 'egg' in use_model_names[0][0]\n#Make sure all model files are present\nmx_pulses = 0  # Find the max number we will ever use\nmn_pulses = 9999\nfor i, (model_file, expl, pulses) in enumerate(use_model_names):\n    print(f'{model_file}')\n    assert os.path.exists(os.path.join(MODEL_DIR,model_file))\n    if pulses > mx_pulses:\n        mx_pulses = pulses\n    if pulses < mn_pulses:\n        mn_pulses = pulses\nprint(f'all model files present, max pulses {mx_pulses} min {mn_pulses}')\ngeometry = utils.read_geometry_file(PATH_DATASET)\ndevice = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\nprint(torch.cuda.get_device_name(0))\nprint('DEVICE', device)\nssub = pd.read_parquet(os.path.join(PATH_DATASET, \"sample_submission.parquet\"))\nssub.set_index(\"event_id\", inplace=True)\nprint(ssub.head())\nif len(ssub) > 5:\n    #Running in submission mode, make sure not TRAIN to avoid embarrasing mistake on submitting\n    MODE='test'\n    FULL=True\n    SAMPLES=4\nssub = ssub.apply(np.float32)\nprint(ssub.info())\nclass Param:\n    def print(self):\n        print('=== PARAMS ===')\n        for v in self.__dict__:\n            print(f'{v} = {self.__dict__[v]}')\n        print('==============')\n#These params will be used as default for models that do not have params stored\nparam = Param()\nparam.nb_inputs = 6\nparam.nearest_neighbors=8\nparam.add_global_variables_after_pooling=True\nparam.take_log_charge=True\nparam.MAX_PULSES = 2000\nparam.add_vest = False\nparam.feature_sort = None\nparam.num_lstm = 2\nparam.use_gru = False\nparam.test_batch = 999\nparam.take_log_charge=True\nparam.dynedge_layer_sizes = None\nparam.knn_use_time = False\nparam.sort_features=False\nparam.limit_pulses='aaa'  # not using\n#Note - pulses now determined by the entry in use_models array\nparam.x_size = 'aaa'    # num pulses data for X values  NOT USING\nparam.data_size = 'aaa' # num pulses for Data() values  NOT USING\nparam.dynedge_layer_sizes = None\nparam.knn_use_time = False\nparam.sort_features=False\nparam.use_qe = False\n\n","metadata":{"execution":{"iopub.status.busy":"2023-04-19T16:12:33.128036Z","iopub.execute_input":"2023-04-19T16:12:33.128503Z","iopub.status.idle":"2023-04-19T16:12:33.176039Z","shell.execute_reply.started":"2023-04-19T16:12:33.128465Z","shell.execute_reply":"2023-04-19T16:12:33.174599Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def collapse(eid_ary, vec, az_zen_true):\n    #groupby and take mean of results by eid_ary\n    start=time.time()\n    #try sorting everything together\n    args = eid_ary.argsort()\n    vec_sorted = vec[args]\n    eid_sorted = eid_ary[args]\n    eids = np.unique(eid_ary)\n    vec_mean=[]\n    for v in np.split(vec_sorted, np.unique(eid_sorted, return_index=True)[1][1:], axis=0):\n        #print(f'v is {v}')\n        vec_mean.append(v.mean(axis=0))\n    vec_mean = np.stack(vec_mean, axis=0)\n    print(f'vec_mean shape {vec_mean.shape} time for collapse {time.time()-start:8.4f}')\n    return eids, vec_mean, az_zen_true[:len(eids)]\n","metadata":{"execution":{"iopub.status.busy":"2023-04-19T16:12:33.179209Z","iopub.execute_input":"2023-04-19T16:12:33.179732Z","iopub.status.idle":"2023-04-19T16:12:33.189995Z","shell.execute_reply.started":"2023-04-19T16:12:33.179684Z","shell.execute_reply":"2023-04-19T16:12:33.187355Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def cartesian_to_polar_ary(x,y,z):\n    \"\"\"x,y,z are N,1 vectors\"\"\"\n    x2y2 = x**2 + y**2\n    r = np.sqrt(x2y2 + z**2)\n    mask = x2y2 < 1e-6\n    x2y2[mask] = 1e-6\n    azimuth = np.arccos(x / np.sqrt(x2y2)) * np.sign(y)\n    zenith = np.arccos(z / r)  # returns 0-2pi, cannot be negative\n    #Adjust built-in here\n    mask = azimuth<0\n    azimuth[mask] = azimuth[mask] + 2 * math.pi\n    #zenith can't be less than zero if it came from arccos\n    return azimuth, zenith\n\n\ndef process_one(model, dataloader, is_dyno):\n    #take a model, process through dataloader (iterable) and return stacked lists\n    eid_list=[]\n    az_list=[]\n    zen_list=[]\n    vec_list=[]\n    start=time.time()\n    tot=0\n    az_zen_true_list=[]\n    for i, Q in enumerate(dataloader):\n        if is_dyno:\n            data=Q\n            eid=data.eid.cpu().detach().numpy()\n            #print('process_one az_zen', data.az_zen.shape)\n            az_zen_true_list.append(data.az_zen.cpu().detach().numpy())\n            data = data.to(device)\n            pred = model(data)\n        else:\n            X, eid, az_zen = Q\n            X = X.to(device)\n            az_zen_true_list.append(az_zen)\n            #print('XX', X.shape)\n            valid_len = torch.zeros(1)  # not used\n            pred = model(X, valid_len)\n        pred_local = pred.cpu().detach().numpy()\n        #pred_local = np.ones((len(data.eid),3))\n        vec_list.append(pred_local)\n        az, zen = cartesian_to_polar_ary(pred_local[:,0],pred_local[:,1],pred_local[:,2])\n        if i % 50==0:\n            print(f'{time.time()-start:8.2f} minibatch {i} eids {len(eid)} pred shape {pred_local.shape}')\n        eid_list.append(eid)\n        az_list.append(az)\n        zen_list.append(zen)     \n        del pred\n        tot += BATCH_SIZE\n    return (np.concatenate(eid_list,axis=0),\n            np.concatenate(vec_list,axis=0),\n            np.concatenate(az_zen_true_list,axis=0))\n\ndef normalize_dog_batch(B):\n    B=B.copy()\n    for i in range(len(B)):\n        B[i,:,:] = normalize_dog(B[i,:,:])\n    return B\n    \ndef normalize_dog(X):\n    #Columns are before reordering \n    X=X.astype(np.float32)  # does a copy\n    N = len(X)\n    #dyno version:\n    #X[:,4] = np.log10(X[:,4]+.0001)/3\n    #dog version\n    #print('len', len(X), X[:,4])\n    X[:,4] = np.log10(X[:,4]+.01)/3\n    #x,y,z\n    X[:,:3] = X[:,:3] / 500\n    #Time - should it be subtract min?\n    #dyno\n    #X[:,3] = (X[:,3] - 1e4)/ 3e4\n    #dog version\n    mn = np.min(X[:,3])\n    X[:,3] = (X[:,3] - mn) / 1e4  # time\n    return X\n\ndef normalize_dyno(X):\n    \"\"\"normalize features - make each column of positions, zero mean and max of 1 or -1\"\"\"\n    #Columns are before reordering \n    X=X.astype(np.float64)\n    N = len(X)\n    #dyno version:\n    X[:,4] = np.log10(X[:,4]+.0001)/3\n    #x,y,z\n    X[:,:3] = X[:,:3] / 500\n    #Time - should it be subtract min?\n    #dyno\n    X[:,3] = (X[:,3] - 1e4)/ 3e4\n    return X\n\nclass Dataset_idx():\n    \"\"\"returns df by index for all those pulses, merged with geometry.  200,000 events per batch_id\n    mx_pulses is the most we will need (including later, as it is cached)\n    caches result\n    anything over sample_pulses gets resampled to a total of samples times\n    \"\"\"\n    \n    def __init__(self, file, mx_pulses, sample_pulses, samples=1):\n        self.batch_df = pd.read_parquet(file)\n        if not FULL:\n            self.batch_df = self.batch_df.iloc[:20000]  #just some of first pulses\n        self.eids = self.batch_df.index.unique().to_numpy()   #unique makes them sorted!\n        #get the pulse counts for each event\n        self.counts = self.batch_df.groupby('event_id')['charge'].count().values\n        self.num_over = np.sum(self.counts > sample_pulses)\n        self.mx_pulses = mx_pulses\n        self.samples = samples\n        #NOTE: we can probably save some effort here if do less samples on events where\n        #only over mx_pulses a little bit\n        print(f'for {file} eids {len(self.eids)} num_over {self.num_over}')\n        print('number of eids', len(self.eids))\n        print(f\"loaded {file} {len(self.batch_df)}\")\n        #repeat the eids over mx_pulses for each sample\n        if self.samples > 1:\n            big_eid_mask = self.counts > sample_pulses\n            big_eids = self.eids[big_eid_mask]\n            dups = np.tile(big_eids, self.samples-1)\n            self.eids = np.concatenate((self.eids, dups))\n            print('number of eids after resamples added', len(self.eids))\n        cache.clear()   # clear the cache\n\n    def __len__(self):\n        total_len = len(self.eids)\n        return total_len\n\n    def __getitem__(self, idx):\n        #Find the sensor type for each sensor\n        #print(idx)\n        #return X, eid  for that idx.  idx is 0-num_events, not the event_id\n        try:\n            return cache[idx]\n        except KeyError:\n            cols=['x', 'y', 'z', 'time', 'charge', 'auxiliary', 'qe']\n            #we need to add qe for the 7th column\n            eid = self.eids[idx]\n            start = time.time()\n            df = self.batch_df.loc[eid]\n            #print(f'{time.time()-start:8.4f} loc')\n            #Try random without sorting - we want to randomly reorder them even if we don't have more than mx_pulses,\n            #because later datasets will just take the first ones that they need and we need random sample of them\n            if len(df) > self.mx_pulses:\n                df=df.sample(n=self.mx_pulses)\n            #print(f'{time.time()-start:8.4f} sample')\n            #Could do faster if we check string_id==sensor_id//60, string_id<78 is REG sensor\n            #def mapper(x):\n            #    string_id, depth_id, typ = utils.sensor_type(x.sensor_id)\n            #    return string_id, typ\n            #rc = df.apply(mapper, axis=1, result_type='expand')\n            #print(rc)\n            df=df.copy()\n            #df['type']   = rc.iloc[:,1]\n            #df['qe'] = 1 + (df['type']>0)*.35\n            df['qe'] = 1 + ((df['sensor_id'].floordiv(60))>=78)*.35\n            df = df.merge(geometry, left_on=\"sensor_id\", right_index=True)\n            X = df[cols].to_numpy().astype(np.float32)    # 7 features\n            if MODE == 'train':\n                return X, eid, azimuth_true[eid], zenith_true[eid]\n            else:\n                return X, eid, None, None\n\nclass Dataset_data(torch.utils.data.Dataset):\n    def __init__(self, cache, pulses):\n        self.xs = cache\n        self.pulses = pulses\n        \n    def __len__(self):\n        return len(self.xs)\n    \n    def __getitem__(self, idx):\n        X, eid, az, zen = self.xs[idx]\n        num_pulses = len(X)\n        if num_pulses > self.pulses:\n            #ow_ids = random.sample(range(int(len(X))), self.pulses)\n            # = X[row_ids,:]\n            X = X[:self.pulses,:]\n        X = normalize_dyno(X)\n        #print(f'X shape {X.shape}')\n        if az is None:\n            az=0\n        if zen is None:\n            zen=0\n        data = geometric.data.Data(x=torch.tensor(X[:,:6],dtype=torch.float32),\n                                   eid=torch.tensor(eid,dtype=torch.int64),\n                                   az_zen=torch.tensor((az,zen), dtype=torch.float32).reshape(1,2),\n                                   n_pulses = torch.tensor(num_pulses, dtype=torch.int64))\n        return data\n    \n    \nclass Dataset_x(torch.utils.data.Dataset):\n    def __init__(self, cache, pulses):\n        self.xs = cache\n        self.pulses = pulses\n        \n    def __len__(self):\n        return len(self.xs)\n    \n    def __getitem__(self, idx):\n        X, eid, az, zen = self.xs[idx]\n        num_pulses = len(X)\n        #Need to pad/clip X\n        if num_pulses > self.pulses:\n            #need random sample\n            X=X[:self.pulses]\n            #ow_ids = random.sample(range(int(len(X))), self.pulses)\n            # = X[row_ids,:]\n        elif num_pulses < self.pulses:\n            X=np.pad(X, ((0, self.pulses - num_pulses), (0,0)), 'constant')\n            #A quirk causes use to require a QE of 1.0 for padded entries\n            mask = (X[:,6]==0)\n            X[mask,6]=1.0\n        X = normalize_dog(X)\n        #print(f'X shape {X.shape}')\n        if az is None:\n            az=0\n        if zen is None:\n            zen=0\n        return X, eid, np.array([az, zen]).reshape((2,))\n    \n#make a generator that uses Pools to get the answer more quickly\ndef fun_pool(n):\n    return dataset[n]\nstart = time.time()\ndef gen_pooled(batch_size, num_pulses, nb_inputs):\n    #cache data from dataset,\n    #but return batch_size x num_pulses x nb_inputs\n    print(f'gen_pooled {batch_size} x {num_pulses} x {nb_inputs}')\n    #if True:\n    with multiprocessing.Pool(2) as pool:\n        N=batch_size\n        batch_idx = 0\n        ret = np.zeros((N, num_pulses, nb_inputs))\n        ret[:,:,6]=1.0  # we set qe to 1 for padded entries\n        eid_ret = np.zeros((N,))\n        az_zen = np.zeros((N, 2))\n        for i, (X, eid, az, zen) in enumerate(pool.imap(fun_pool, range(len(dataset)), 10)):\n        #for i, (X, eid, az, zen) in enumerate(map(fun_pool, range(len(dataset)))):\n            \n            #print(f'{time.time()-start:8.4f} {X.shape}')\n            cache[i] = (X, eid, az, zen)\n            pulses = len(X)\n            if len(X) > num_pulses:\n                #ow_ids = random.sample(range(int(len(X))), num_pulses)\n                # = X[row_ids,:]\n                X=X[:num_pulses]   #dbg\n            ret[batch_idx,:len(X),:] = X[:,:nb_inputs]\n            eid_ret[batch_idx]=eid\n            az_zen[batch_idx]=(az, zen)\n            batch_idx += 1\n            #print(f'eid {eid} {len(X)} {X[:10]} {az} {zen}')\n            if batch_idx == N:\n                #print(az_zen)\n                ret=normalize_dog_batch(ret)\n                yield torch.tensor(ret,dtype=torch.float32), eid_ret, az_zen\n                ret = np.zeros((N, num_pulses, nb_inputs))\n                ret[:,:,6]=1.0  # we set qe to 1 for padded entries\n                eid_ret = np.zeros((N,))\n                az_zen = np.zeros((N, 2))\n                batch_idx = 0\n        #final one\n        if batch_idx > 0:\n            B=ret[:batch_idx,:,:]\n            B=normalize_dog_batch(B)\n            yield (torch.tensor(B,dtype=torch.float32),\n                   eid_ret[:batch_idx],\n                   az_zen[:batch_idx,:])\n        \n\ndef vector_wt(vecs, weights):\n    sum = np.zeros_like(vecs[0])\n    for vec, wt in zip(vecs, weights):\n        sum += vec * wt\n    azimuth_pred, zenith_pred = utils.cartesian_to_polar_ary(sum[:,0],sum[:,1],sum[:,2])\n    return azimuth_pred, zenith_pred\n\ndef score_it(vector, azimuth_true, zenith_true, s=''):\n    azimuth_pred, zenith_pred = utils.cartesian_to_polar_ary(vector[:,0],vector[:,1],vector[:,2])\n    score = utils.angular_dist_score(azimuth_true, zenith_true, azimuth_pred, zenith_pred)\n    print(f'{s:20} score {score:8.5f}')\n    return score\n","metadata":{"execution":{"iopub.status.busy":"2023-04-19T16:12:33.192522Z","iopub.execute_input":"2023-04-19T16:12:33.193521Z","iopub.status.idle":"2023-04-19T16:12:33.247258Z","shell.execute_reply.started":"2023-04-19T16:12:33.193468Z","shell.execute_reply":"2023-04-19T16:12:33.245941Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"meta_df = None\nif MODE == 'train':\n    #This takes a lot of memory, only do it in train mode\n    #Update: we will read the parsed out version instead with just 1 batch in it\n    meta_df = pd.read_parquet(\n        os.path.join(MODEL_DIR, f\"train_meta_601.parquet\")\n    )\n    #print(meta_df)\n    #just do batch 1, 2\n    #meta_df = meta_df[meta_df['batch_id'] <= 2]   #meta_1 does not have batch_id in it\n    meta_df.set_index('event_id', inplace=True)\n    true_vals = meta_df[['azimuth','zenith']].to_dict()\n    azimuth_true = true_vals['azimuth']\n    zenith_true = true_vals['zenith']\n    del meta_df\n    print(f'loaded true azimith and zenith values {len(azimuth_true)} {len(zenith_true)}')\n","metadata":{"execution":{"iopub.status.busy":"2023-04-19T16:12:33.251783Z","iopub.execute_input":"2023-04-19T16:12:33.253151Z","iopub.status.idle":"2023-04-19T16:12:33.800222Z","shell.execute_reply.started":"2023-04-19T16:12:33.253106Z","shell.execute_reply":"2023-04-19T16:12:33.799051Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if MODE == 'train':\n    ls = glob.glob(os.path.join(PATH_DATASET,MODE,'batch_601.parquet'))\nelse:\n    ls = glob.glob(os.path.join(PATH_DATASET,MODE,'*.parquet'))\nstart = time.time()\nprint(f'starting, {len(ls)} files')\nfor file in ls:\n    #set the dataset, clearing the cache\n    dataset = Dataset_idx(file, mx_pulses, mn_pulses, samples=SAMPLES)\n    vecs=[]\n    for i, (model_file,expl,pulses) in enumerate(use_model_names):\n        print(f'{i} trying to load {model_file} p{pulses}')\n        model_type = model_file.split('-')[0]\n        is_dyno = ( model_type == 'dyno' )\n        the_module = the_net_modules[model_type]\n        obj = torch.load(os.path.join(MODEL_DIR,model_file))\n        build_param=None\n        if 'osy' in model_file:\n            obj2 = torch.load(os.path.join(MODEL_DIR, 'dog-wnv-state-480-s1011-obj.pth'))\n            build_param = obj2['param']\n        if 'igl' in model_file:\n            #hack - use param we have set here\n            use_param = param\n        else:\n            assert 'model_state_dict' in obj, f'need state_dict and param in the model {model_name} {list(obj.keys())}'\n            new_param = obj['param']\n            use_param=copy.copy(param)\n            #update with module params\n            use_param.__dict__.update(the_module.param.__dict__)\n            #then update with ones from the file\n            use_param.__dict__.update(new_param.__dict__)\n            use_param.special_init = False\n            use_param.BATCH_SIZE = 128\n        #use_param.print()\n        model = the_module.Net(use_param if build_param is None else build_param)\n        if 'igl' in model_file:\n            model.load_state_dict(obj)\n        else:\n            model.load_state_dict(obj['model_state_dict'])\n        model = model.to(device)\n        model.eval()\n        if i==0:\n            dataloader = gen_pooled(BATCH_SIZE, pulses, use_param.nb_inputs)\n        else:\n            if is_dyno:\n                dataloader = geometric.loader.DataLoader(Dataset_data(cache, pulses),\n                                                         batch_size=BATCH_SIZE,\n                                                         shuffle=False,\n                                                         num_workers=0)\n            else:\n                dataloader = torch.utils.data.DataLoader(Dataset_x(cache,pulses),\n                                                         batch_size=BATCH_SIZE,\n                                                         shuffle=False,\n                                                         num_workers=0)\n\n        result = process_one(model,dataloader, is_dyno)\n        result = collapse(*result)\n        eid_ary, vec, az_zen_true = result\n        if MODE == 'train':\n            score_it(vec, az_zen_true[:,0], az_zen_true[:,1],s=model_file)\n        vecs.append(vec)\n        print(f'{time.time()-start:8.1f} results got {model_file} {len(vec)}')\n        #if MODE == 'train':\n        #    break\n    #average the results\n    print(f'finished {file}')\n    azimuths, zeniths = vector_wt(vecs, weights)\n    if MODE == 'train':\n        wtd = utils.angular_dist_score(az_zen_true[:,0], az_zen_true[:,1],azimuths,zeniths)\n        print(f'got weighted score {wtd:8.5f}')\n    if MODE != 'train':\n        for i, (eid, azimuth, zenith) in enumerate(zip(eid_ary, azimuths, zeniths)):\n            #this will increase size of ssub if eid does not exist!\n            #so it slows down as we add 200000 things \n            ssub.at[eid, \"azimuth\"] = azimuth\n            ssub.at[eid, \"zenith\"] = zenith\n","metadata":{"execution":{"iopub.status.busy":"2023-04-19T16:12:33.803344Z","iopub.execute_input":"2023-04-19T16:12:33.804367Z","iopub.status.idle":"2023-04-19T17:07:15.205410Z","shell.execute_reply.started":"2023-04-19T16:12:33.804311Z","shell.execute_reply":"2023-04-19T17:07:15.204024Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Write out submission file\nssub.fillna(0).to_csv('submission.csv', index=True)\nif len(ssub) < 5:\n    print(ssub)\n\n    \n","metadata":{"execution":{"iopub.status.busy":"2023-04-19T17:07:15.207478Z","iopub.execute_input":"2023-04-19T17:07:15.207943Z","iopub.status.idle":"2023-04-19T17:07:15.221240Z","shell.execute_reply.started":"2023-04-19T17:07:15.207901Z","shell.execute_reply":"2023-04-19T17:07:15.219458Z"},"trusted":true},"execution_count":null,"outputs":[]}]}