{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# packages = ('pytorch_lightning timm==0.6.12 ipywidgets==7.7.1 opencv-python zstandard awscli ' \n#                     +  'transformers librosa torchlibrosa torchaudio torchvision ' \n#                     + 'lion-pytorch segmentation-models-pytorch==0.3.2 ' \n#                     'fcwt') # fsspec[s3] albumentations==1.3.0 \n# !pip install -q -U $packages","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:08:03.399264Z","iopub.execute_input":"2023-07-20T01:08:03.399618Z","iopub.status.idle":"2023-07-20T01:08:03.404663Z","shell.execute_reply.started":"2023-07-20T01:08:03.399590Z","shell.execute_reply":"2023-07-20T01:08:03.403676Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# FAST = False\nSAMPLE = 4","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:08:03.410964Z","iopub.execute_input":"2023-07-20T01:08:03.412261Z","iopub.status.idle":"2023-07-20T01:08:03.416683Z","shell.execute_reply.started":"2023-07-20T01:08:03.412222Z","shell.execute_reply":"2023-07-20T01:08:03.415629Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# DATA_PATH = '/data/'\nDATASET = 'walkdata5'\nTAG = 'tlvmc-parkinsons-freezing-gait-prediction'\nINPUT_PATH = '/kaggle/input/' + TAG\nDATA_PATH = '/tmp/'\nCODE_PATH = '/kaggle/input/' + DATASET\nPREFIX = 'walk/'\nOFFLINE = True","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:08:03.418644Z","iopub.execute_input":"2023-07-20T01:08:03.419705Z","iopub.status.idle":"2023-07-20T01:08:03.427146Z","shell.execute_reply.started":"2023-07-20T01:08:03.419670Z","shell.execute_reply":"2023-07-20T01:08:03.426170Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls $CODE_PATH","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:08:03.429003Z","iopub.execute_input":"2023-07-20T01:08:03.429762Z","iopub.status.idle":"2023-07-20T01:08:04.395815Z","shell.execute_reply.started":"2023-07-20T01:08:03.429711Z","shell.execute_reply":"2023-07-20T01:08:04.394608Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"WHEEL_PATH = CODE_PATH + '/whl'\nimport os\nfor p in os.listdir(WHEEL_PATH):\n    if 'segment' in p: continue; \n    !pip install {WHEEL_PATH}/{p}","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:08:04.398277Z","iopub.execute_input":"2023-07-20T01:08:04.399001Z","iopub.status.idle":"2023-07-20T01:08:38.207494Z","shell.execute_reply.started":"2023-07-20T01:08:04.398958Z","shell.execute_reply":"2023-07-20T01:08:38.206300Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nsys.path.append(CODE_PATH + '/smp')\n# import segmentation_models_pytorch as smp\n# smp.Unet(encoder_weights = None);","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:08:38.210086Z","iopub.execute_input":"2023-07-20T01:08:38.210881Z","iopub.status.idle":"2023-07-20T01:08:38.216476Z","shell.execute_reply.started":"2023-07-20T01:08:38.210816Z","shell.execute_reply":"2023-07-20T01:08:38.215505Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import zipfile\nimport boto3\ns3 = boto3.client('s3')\n\nimport os\nimport io\nfrom joblib import Parallel, delayed\n\nimport json\nimport pickle\nimport zstandard as zstd\nzd = zstd.ZstdDecompressor()\nzc = zstd.ZstdCompressor()\n\nimport math\nimport random\nimport datetime","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:08:38.220106Z","iopub.execute_input":"2023-07-20T01:08:38.221013Z","iopub.status.idle":"2023-07-20T01:08:40.666908Z","shell.execute_reply.started":"2023-07-20T01:08:38.220983Z","shell.execute_reply":"2023-07-20T01:08:40.665981Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# walk and list all files\nobjs = []\nfor root, path, files in os.walk(INPUT_PATH):\n    if 'train' in root: continue\n    objs.extend([os.path.join(root, f).replace(INPUT_PATH + '/', '') for f in files])\n    print(root, len(files))\nlen(objs)","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:08:40.668202Z","iopub.execute_input":"2023-07-20T01:08:40.668561Z","iopub.status.idle":"2023-07-20T01:08:40.917790Z","shell.execute_reply.started":"2023-07-20T01:08:40.668530Z","shell.execute_reply":"2023-07-20T01:08:40.916847Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def getFile(file, save = True):\n#     s3 = boto3.client('s3')\n#     zd = zstd.ZstdDecompressor()\n#     file_path = DATA_PATH + 'data/' + file + '.zstd'\n#     if not os.path.exists(file_path) or os.path.getsize(file_path) == 0:\n#         if save: os.makedirs(os.path.dirname(file_path), exist_ok=True)\n#         fin = s3.get_object(Bucket = DATA_BUCKET, Key = PREFIX + 'data/' + file + '.zstd')['Body']#.read()        \n#         if save:\n#             with open(file_path, 'wb') as f: f.write(fin.read())\n    file_path = os.path.join(INPUT_PATH, file)# + '.zstd'\n    out = io.BytesIO()\n    if os.path.exists(file_path) and os.path.getsize(file_path) > 0:\n        with open(file_path, 'rb') as fin:\n#             zd.copy_stream(fin, out)\n#     else:        \n            out.write(fin.read())\n    out.seek(0)\n    # print(out.getbuffer().nbytes)\n    return out","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:08:40.919174Z","iopub.execute_input":"2023-07-20T01:08:40.919546Z","iopub.status.idle":"2023-07-20T01:08:40.928824Z","shell.execute_reply.started":"2023-07-20T01:08:40.919511Z","shell.execute_reply":"2023-07-20T01:08:40.927554Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nplt.rcParams['figure.figsize'] = (7, 4)\n\nfrom IPython.display import display\nnp.random.seed(datetime.datetime.now().microsecond)","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:08:40.930562Z","iopub.execute_input":"2023-07-20T01:08:40.931598Z","iopub.status.idle":"2023-07-20T01:08:40.938503Z","shell.execute_reply.started":"2023-07-20T01:08:40.931563Z","shell.execute_reply":"2023-07-20T01:08:40.937376Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"files = sorted(objs)#o)[0] for o in objs])\ndisplay(files[:10])\n\nevents = pd.read_csv(getFile('events.csv'))\nsubjects = pd.read_csv(getFile('subjects.csv'))\ntasks = pd.read_csv(getFile('tasks.csv'))\ndaily_metadata = pd.read_csv(getFile('daily_metadata.csv'))\ndefog_metadata = pd.read_csv(getFile('defog_metadata.csv'))\ntdcsfog_metadata = pd.read_csv(getFile('tdcsfog_metadata.csv'))\nsample = pd.read_csv(getFile('sample_submission.csv'), index_col = 0, \n                    ).astype(np.float32)","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:08:40.940122Z","iopub.execute_input":"2023-07-20T01:08:40.940632Z","iopub.status.idle":"2023-07-20T01:08:41.340967Z","shell.execute_reply.started":"2023-07-20T01:08:40.940598Z","shell.execute_reply":"2023-07-20T01:08:41.339985Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def prep_metadata(defog_metadata, tdcsfog_metadata, daily_metadata, subjects, full = True,\n                    expanded = True):\n    m1 = defog_metadata.copy()\n    m1.insert(3, 'Test', 0)\n    m1 = m1.merge(subjects, on = ['Subject', 'Visit', ], how = 'inner')\n    assert len(m1) == len(defog_metadata)\n\n    m2 = tdcsfog_metadata.copy()\n    m2 = m2.merge(subjects.drop(columns = 'Visit'), on = ['Subject', ], how = 'inner')\n    assert len(m2) == len(tdcsfog_metadata)\n\n    m3 = daily_metadata.copy()\n    m3 = m3.merge(subjects, on = ['Subject', 'Visit' ], how = 'inner')\n    m3.insert(3, 'Test', 0)\n    m3.insert(4, 'Medication', (defog_metadata.Medication == 'on').mean())\n    m3.drop(columns = [c for c in m3.columns if 'recording' in c], inplace = True)\n    assert len(m3) == len(daily_metadata)\n\n    metadata = pd.concat([m1, m2, #m3\n                            ], axis = 0)\n    metadata.Medication = 1 * (metadata.Medication == 'on')\n    metadata.Sex = 1 * (metadata.Sex == 'M')\n    \n    if expanded:\n        metadata['num_tests'] = metadata.groupby('Subject').transform(lambda x: x.nunique()).Test\n        metadata['max_visit'] = metadata.groupby('Subject').transform(lambda x: x.max()).Visit\n        metadata['visit_medications'] = metadata.groupby(['Subject', 'Visit']).transform('nunique').Medication \n        metadata['UPDRS_On_vs_Off'] = metadata.UPDRSIII_On - metadata.UPDRSIII_Off\n        # add 4 columns to dmetadata\n    \n\n    if full:\n        # null fix\n        metadata['Uon_null'] = 1 * (metadata.UPDRSIII_On.isnull())\n        metadata['Uoff_null'] = 1 * (metadata.UPDRSIII_Off.isnull())\n        metadata.UPDRSIII_On = metadata.UPDRSIII_On.fillna(metadata.UPDRSIII_On.mean())\n        metadata.UPDRSIII_Off = metadata.UPDRSIII_Off.fillna(metadata.UPDRSIII_Off.mean())\n        if expanded:\n            metadata.UPDRS_On_vs_Off = metadata.UPDRS_On_vs_Off.fillna(metadata.UPDRS_On_vs_Off.mean())   \n\n    metadata.set_index('Id', inplace = True)\n    \n    if full:\n        metadata['Test_Nonzero'] = 1. * (metadata.Test > 0)\n        # for i in range(1):\n        #     metadata['Test{}'.format(i)] = metadata.Test == i\n        metadata.iloc[:, 1:] = metadata.iloc[:, 1:].astype(np.float32)\n        metadata.iloc[:, 1:] = (metadata.iloc[:, 1:] - metadata.iloc[:, 1:].mean(0)) / metadata.iloc[:, 1:].std(0)  \n        metadata.iloc[:, 1:] = metadata.iloc[:, 1:].clip(-3, 3)\n\n    msubject = metadata.Subject\n    metadata.drop(columns = 'Subject', inplace = True)\n    if full:\n        metadata = metadata.astype(np.float32)\n    \n    m3 = m3.set_index('Id')\n    assert m3.shape[1] <= metadata.shape[1]; i = 0\n    while m3.shape[1] < metadata.shape[1]:\n        m3.insert(m3.shape[1], 'dummy_{}'.format(i), 0); i += 1\n    m3.iloc[:, -1] = metadata.iloc[:, -1].min() # yes, hack, for default_metadata in dataset.py \n\n    return metadata, msubject, m3","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:08:41.342230Z","iopub.execute_input":"2023-07-20T01:08:41.342569Z","iopub.status.idle":"2023-07-20T01:08:41.363786Z","shell.execute_reply.started":"2023-07-20T01:08:41.342538Z","shell.execute_reply":"2023-07-20T01:08:41.362583Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fog_files = sorted([f for f in files if 'test/' in f and 'fog' in f])\n# if FAST and len(fog_files) == 2:\n#     fog_files = [f for f in fog_files if 'tdcs' in f]\nprint(len(fog_files))\n\nunlabeled_files = [f for f in files if 'unlabeled/' in f]# and not any([z in f for z in common_files])]\nprint(len(unlabeled_files))","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:08:41.369581Z","iopub.execute_input":"2023-07-20T01:08:41.370166Z","iopub.status.idle":"2023-07-20T01:08:41.378172Z","shell.execute_reply.started":"2023-07-20T01:08:41.370134Z","shell.execute_reply":"2023-07-20T01:08:41.377048Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# load and process data\ndef load(f):\n    return np.load(DATA_PATH + 'cache/' + f + '.npy')\n\ndef process(f):\n    # if exists, return stats\n    cache_file = DATA_PATH + 'cache/' + f + '.npy'\n    if os.path.exists(cache_file) and os.path.getsize(cache_file) > 0: \n        return load(f).shape, os.path.getsize(cache_file)\n    \n    # if not, load array,\n    df = pd.read_csv(getFile(f, ), )\n#     print(df)\n    assert (df.index == df.Time).all()\n    # verify\n    assert (df.Time == df.index).all()\n    assert all(df.columns == ['Time',\n                           'AccV', 'AccML', 'AccAP',\n                           'StartHesitation', 'Turn', 'Walking',\n                            'Valid', 'Task'][:len(df.columns)])\n    if 'train/' not in f: assert len(df.columns) in [4]\n    \n    v = np.zeros((len(df), 12), dtype = np.float32)\n    v[:, 6:8] = 1\n    v[:, :df.shape[1] - 1] = df.iloc[:, 1:]\n\n#     fid = f.split('/')[-1].split('.')[0]\n#     mult = 100 if 'tdcs' not in f else 128\n#     for e in events[events.Id == fid].itertuples():\n#         v[int(round(e.Init * mult)): int(round(e.Completion * mult)), 8] = 1\n#         v[int(round(e.Init * mult)): int(round(e.Completion * mult)), 9] = e.Kinetic\n#         v[int(round(e.Init * mult)): int(round(e.Completion * mult)), 10] = 1 - e.Kinetic\n#     for e in tasks[tasks.Id == fid].itertuples():\n#         v[int(round(e.Begin * mult)): int(round(e.End * mult)), 11] = task_dict[e.Task]\n\n    # store as compresssed;\n    assert v.dtype == np.float32\n    # zc = zstd.ZstdCompressor()\n    # compr = zc.compress(pickle.dumps(v))\n    os.makedirs(os.path.dirname(cache_file), exist_ok = True)\n    np.save(cache_file, v)\n    # with open(cache_file, 'wb') as f:\n    #     f.write(compr)\n\n    return v.shape, os.path.getsize(cache_file)#len(compr)\n","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:08:41.379537Z","iopub.execute_input":"2023-07-20T01:08:41.380396Z","iopub.status.idle":"2023-07-20T01:08:41.392246Z","shell.execute_reply.started":"2023-07-20T01:08:41.380362Z","shell.execute_reply":"2023-07-20T01:08:41.391284Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# fog_files = sorted(fog_files)[-1:]\n# fog_files # ***","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:08:41.393714Z","iopub.execute_input":"2023-07-20T01:08:41.394202Z","iopub.status.idle":"2023-07-20T01:08:41.405875Z","shell.execute_reply.started":"2023-07-20T01:08:41.394168Z","shell.execute_reply":"2023-07-20T01:08:41.404811Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n# process all files\nr = Parallel(os.cpu_count())(delayed(process)(f) for f in fog_files[:])\n\n# display counts;\nfogcount = dict(zip(fog_files, [e[0][0] for e in r]))\n[sum([v for k, v in fogcount.items() if s in k]) / 1e6\n        for s in ['/defog/', '/tdcsfog/', ]]","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:08:41.407511Z","iopub.execute_input":"2023-07-20T01:08:41.407913Z","iopub.status.idle":"2023-07-20T01:08:43.297041Z","shell.execute_reply.started":"2023-07-20T01:08:41.407872Z","shell.execute_reply":"2023-07-20T01:08:43.295933Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!du -sh {DATA_PATH}cache/test*","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:08:43.301972Z","iopub.execute_input":"2023-07-20T01:08:43.304686Z","iopub.status.idle":"2023-07-20T01:08:44.253274Z","shell.execute_reply.started":"2023-07-20T01:08:43.304643Z","shell.execute_reply":"2023-07-20T01:08:44.252052Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%run -i {CODE_PATH}/dataset.py","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:08:44.256098Z","iopub.execute_input":"2023-07-20T01:08:44.257187Z","iopub.status.idle":"2023-07-20T01:09:00.323823Z","shell.execute_reply.started":"2023-07-20T01:08:44.257142Z","shell.execute_reply":"2023-07-20T01:09:00.322885Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%run -i {CODE_PATH}/model.py","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:09:00.328789Z","iopub.execute_input":"2023-07-20T01:09:00.331663Z","iopub.status.idle":"2023-07-20T01:09:10.685640Z","shell.execute_reply.started":"2023-07-20T01:09:00.331620Z","shell.execute_reply":"2023-07-20T01:09:10.684671Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def getData(params, m, lcount, metadata):\n    test_data = WalkDataset(\n        {k: v for k, v in lcount.items()},\n                      metadata, load, None, -1,\n                      test = True,\n                      **getParams(WalkDataset, params))\n    test_loader = DataLoader(test_data, batch_size = 32,\n                              num_workers = os.cpu_count())\n    print(len(test_data), len(test_loader))\n    return test_data, test_loader","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:09:10.688913Z","iopub.execute_input":"2023-07-20T01:09:10.689721Z","iopub.status.idle":"2023-07-20T01:09:10.695656Z","shell.execute_reply.started":"2023-07-20T01:09:10.689692Z","shell.execute_reply":"2023-07-20T01:09:10.694738Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def getModel(params, m):\n    model = WalkNetwork(params, **getParams(WalkNetwork, params)).to(device)\n    model.load_state_dict(\n        pickle.load(open(CODE_PATH + '/models/' + m + '.pt', 'rb')), \n    )\n    model.to(device);\n    model.eval();\n    return model","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:09:10.697397Z","iopub.execute_input":"2023-07-20T01:09:10.698028Z","iopub.status.idle":"2023-07-20T01:09:10.709517Z","shell.execute_reply.started":"2023-07-20T01:09:10.697994Z","shell.execute_reply":"2023-07-20T01:09:10.708339Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def infer(model, test_loader):\n    yps, ss, fs, idxs = [], [], [], []\n    for batch in test_loader:\n        x, y, s, frac, m_, f, i, flen = batch\n        with torch.no_grad():\n            yp = model(*[e.to(device)\n                                     for e in [x, m_, frac, flen]], \n                                 )[0]\n        yps.append(yp[..., :3].cpu())\n        ss.append(s.cpu())\n        fs.append(f)\n        idxs.append(i.cpu())\n#         print(yp.shape)\n    \n    # process\n#     n_models = yps[0].shape[-2]\n#     print(n_models)\n    \n    fs = np.concatenate(fs)\n    idxs = torch.cat(idxs)\n    slen = yps[0].shape[1]\n    s = torch.cat(ss, 0).reshape(-1, slen, 1).cpu().numpy()\n    yps = torch.cat(yps, 0)\n    yps = yps.reshape(-1, slen, 3).cpu().numpy()        \n    print(yps.shape, s.shape, len(fs), len(idxs))\n    assert all([len(e) == len(yps) for e in [fs, idxs, s, yps]])\n\n    pred_dict, ct_dict, target_dict = {}, {}, {}\n    for f in list(set(fs)):\n        l = idxs[fs == f].max() + slen // SAMPLE\n        pred_dict[f] = np.zeros((l, 3), dtype = np.float32)\n        ct_dict[f] = np.zeros((l, 3), dtype = np.float32)\n\n    for i in range(len(yps)):\n        pred_dict[fs[i]][idxs[i]:idxs[i] + slen] += yps[i] * s[i]\n        ct_dict[fs[i]][idxs[i]:idxs[i] + slen] += s[i]\n\n    # compile and nrm\n    spred_dict = {}; final_yps = []\n    for k in ct_dict:\n        print(ct_dict[k].shape)\n        f = ct_dict[k].sum(1) > 0\n        spred_dict[k] = pred_dict[k][f] / (ct_dict[k][f] + 1e-5)\n        final_yps.append(spred_dict[k])\n\n    final_yps = np.concatenate(final_yps)\n    avg = final_yps.mean(0)\n    print(avg)\n    \n    spred_dict = {k: v/avg for k, v in spred_dict.items()}\n    return spred_dict\n    ","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:09:10.712904Z","iopub.execute_input":"2023-07-20T01:09:10.713202Z","iopub.status.idle":"2023-07-20T01:09:10.729725Z","shell.execute_reply.started":"2023-07-20T01:09:10.713175Z","shell.execute_reply":"2023-07-20T01:09:10.728732Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def inferParallel(model, test_loader):\n    yps, ss, fs, idxs = [], [], [], []\n    for batch in test_loader:\n        x, y, s, frac, m_, f, i, flen = batch\n        with torch.no_grad():\n            yp = model(*[e.to(device)\n                                     for e in [x, m_, frac, flen]], \n                                 )#[0]\n#         SAMPLE = 10\n        yps.append(yp[:, ::SAMPLE, :, :3].cpu())\n        ss.append(s[:, ::SAMPLE].cpu())\n        fs.append(f)\n        idxs.append((i / SAMPLE).long().cpu())\n#         print(yp.shape)\n    \n    # process\n    n_models = yps[0].shape[-2]\n    print(n_models)\n    \n    fs = np.concatenate(fs)\n    idxs = torch.cat(idxs)\n    slen = yps[0].shape[1]\n    s = torch.cat(ss, 0).reshape(-1, slen, 1, 1).cpu().numpy()\n    yp = torch.cat(yps, 0); del yps\n    yps = yp.reshape(-1, slen, n_models, 3).cpu().numpy(); del yp        \n    print(yps.shape, s.shape, len(fs), len(idxs))\n    assert all([len(e) == len(yps) for e in [fs, idxs, s, yps]])\n\n    pred_dict, ct_dict, target_dict = {}, {}, {}\n    for f in list(set(fs)):\n        l = idxs[fs == f].max() + slen\n        pred_dict[f] = np.zeros((l, n_models, 3), dtype = np.float32)\n        ct_dict[f] = np.zeros((l, n_models, 3), dtype = np.float32)\n\n    for i in range(len(yps)):\n        pred_dict[fs[i]][idxs[i]:idxs[i] + slen] += yps[i] * s[i]\n        ct_dict[fs[i]][idxs[i]:idxs[i] + slen] += s[i]\n\n    # compile and nrm\n    spred_dict = {}; final_yps = []\n    for k in ct_dict:\n        print(ct_dict[k].shape)\n        f = ct_dict[k].sum(2).sum(1) > 0\n        spred_dict[k] = pred_dict[k][f] / (ct_dict[k][f] + 1e-5)\n        final_yps.append(spred_dict[k])\n\n    final_yps = np.concatenate(final_yps) # B, M, C\n    avg = final_yps.mean(0)  # M, C\n    print(avg)\n#     avg = avg / avg.mean(0, keepdims = True) # each class should sum to 1;\n#     print(avg)\n#     print((k/avg).sum(0))\n    \n    spred_dict = {k: (v/avg).mean(-2) for k, v in spred_dict.items()}\n    return spred_dict\n    ","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:09:10.733486Z","iopub.execute_input":"2023-07-20T01:09:10.734188Z","iopub.status.idle":"2023-07-20T01:09:10.751140Z","shell.execute_reply.started":"2023-07-20T01:09:10.734151Z","shell.execute_reply":"2023-07-20T01:09:10.750146Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ms = [f.split('.')[0]#, json.load(open(CODE_PATH + '/params/' + f)))\n#               for f in os.listdir(CODE_PATH + '/params')]\n# print(len(ms))","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:09:10.754153Z","iopub.execute_input":"2023-07-20T01:09:10.754485Z","iopub.status.idle":"2023-07-20T01:09:10.765217Z","shell.execute_reply.started":"2023-07-20T01:09:10.754459Z","shell.execute_reply":"2023-07-20T01:09:10.764078Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ms = [(f.split('.')[0], json.load(open(CODE_PATH + '/params/' + f)))\n              for f in os.listdir(CODE_PATH + '/params')]\nprint(len(ms))","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:09:10.767063Z","iopub.execute_input":"2023-07-20T01:09:10.767466Z","iopub.status.idle":"2023-07-20T01:09:10.828846Z","shell.execute_reply.started":"2023-07-20T01:09:10.767432Z","shell.execute_reply":"2023-07-20T01:09:10.827544Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# select_ms = '0c1908, 38f97e, 6724d4, b904fe, 4e7844, b792a2, d64c60, d85bf1, 2233bd, 2c1d4a, b42aae, ef4805, f41ae6, 18da64, 37898f, 9805ce, f65f09, 53b767, 7033ef, 90be5b, fdbb0a'\n# select_ms = select_ms.split(', ')","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:09:10.830343Z","iopub.execute_input":"2023-07-20T01:09:10.830784Z","iopub.status.idle":"2023-07-20T01:09:10.836177Z","shell.execute_reply.started":"2023-07-20T01:09:10.830748Z","shell.execute_reply":"2023-07-20T01:09:10.835088Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from collections import defaultdict\nmgroups = defaultdict(list)\nmpoints = defaultdict(int)\nall_params = []\nfor m, p in ms:\n#     if m not in select_ms: continue\n    all_params.append(p)\n    # if p.get('xformer_init_1') == 0.1: continue;\n    # if p.get('xformer_layers', 2) > 3: continue;\n    # if p['frac_pwr_mult'] != f: continue\n    # if 'seg' in p: continue# or 'melspec' in p: continue\n    # if p['batch_size'] >= 20: continue;\n    r_ = {k: v for k, v in p.items() \n              if k not in ['fold', 'seed' ]#'seed', 'n_folds' ]\n         }\n    mgroups[json.dumps(r_)].append(m)\n    mpoints[json.dumps(r_)] += 1 #- select_ms.index(m)/len(select_ms)\nlen(mgroups)","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:09:38.816086Z","iopub.execute_input":"2023-07-20T01:09:38.816472Z","iopub.status.idle":"2023-07-20T01:09:38.828100Z","shell.execute_reply.started":"2023-07-20T01:09:38.816443Z","shell.execute_reply":"2023-07-20T01:09:38.827188Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.hist(mpoints.values(), bins =250);","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:09:45.996517Z","iopub.execute_input":"2023-07-20T01:09:45.996918Z","iopub.status.idle":"2023-07-20T01:09:46.647706Z","shell.execute_reply.started":"2023-07-20T01:09:45.996881Z","shell.execute_reply":"2023-07-20T01:09:46.646779Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ParallelModel(nn.Module):\n    def __init__(self, models):\n        super().__init__()\n        self.models = nn.ModuleList(models)\n    \n    def forward(self, x, m, frac, flen):\n        return torch.stack([model(x, m, frac, flen)[0]\n                            for model in self.models], -2)\n","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:09:53.884644Z","iopub.execute_input":"2023-07-20T01:09:53.885199Z","iopub.status.idle":"2023-07-20T01:09:53.893932Z","shell.execute_reply.started":"2023-07-20T01:09:53.885163Z","shell.execute_reply":"2023-07-20T01:09:53.892485Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# numpy add two arrays, expanding the smaller one on axis 0\ndef eadd(a, b):\n    if isinstance(a, int): return b.copy();\n    if a.shape[0] < b.shape[0]:\n        a = np.pad(a, ((0, b.shape[0] - a.shape[0]), (0, 0)), mode = 'constant')\n    elif a.shape[0] > b.shape[0]:\n        b = np.pad(b, ((0, a.shape[0] - b.shape[0]), (0, 0)), mode = 'constant')\n    return a + b","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:09:54.615843Z","iopub.execute_input":"2023-07-20T01:09:54.616520Z","iopub.status.idle":"2023-07-20T01:09:54.623449Z","shell.execute_reply.started":"2023-07-20T01:09:54.616487Z","shell.execute_reply":"2023-07-20T01:09:54.622448Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import time","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:09:55.611915Z","iopub.execute_input":"2023-07-20T01:09:55.612675Z","iopub.status.idle":"2023-07-20T01:09:55.619946Z","shell.execute_reply.started":"2023-07-20T01:09:55.612642Z","shell.execute_reply":"2023-07-20T01:09:55.618928Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"all_pred_dict = {}; ct = 0\n# for m in ms:#_, \nfor _, ms in mgroups.items():\n    models = []; m = ms[0];# print(json.loads(_))\n    params = json.load(open(CODE_PATH + '/params/' + m + '.json', 'r'))\n#     if params['frac_pwr_mult'] not in [2.36, 2.04, 1.81, 1.91, 1.88]: continue\n    for m in ms:\n        models.append(getModel(params, m ))    \n    model = ParallelModel(models).to(device).eval()\n    metadata, msubject, md = prep_metadata(defog_metadata, tdcsfog_metadata, daily_metadata, subjects,\n                                      expanded = params['expanded'])\n#     model = getModel(params, m)\n    test_data, test_loader = getData(params, m, fogcount, metadata)\n    start = time.time()\n    spred_dict = inferParallel(model, test_loader)\n#     spred_dict = infer(model, test_loader)\n    print(time.time() - start)\n    \n    for k, v in spred_dict.items():\n        all_pred_dict[k] = eadd(all_pred_dict.get(k, 0), \n                                v * mpoints[_])\n    del test_data, test_loader, spred_dict\n        \n    ct += 1;\n    print()\n    del model#, models\n","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:09:56.068033Z","iopub.execute_input":"2023-07-20T01:09:56.068405Z","iopub.status.idle":"2023-07-20T01:10:33.068885Z","shell.execute_reply.started":"2023-07-20T01:09:56.068375Z","shell.execute_reply":"2023-07-20T01:10:33.067600Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preds = []\nfor k, v in all_pred_dict.items():\n    true_len = len(load(k))\n    print(k, len(v), true_len)\n    v = cv2.resize(v, None, fx = 1, fy = SAMPLE)\n    v = v[:true_len]\n    v = np.pad(v, ((0, true_len - len(v) ), (0, 0)), mode = 'edge')\n    \n    df = pd.DataFrame(v / ct, columns = sample.columns, \n                index = ['{}_{}'.format(k.split('/')[-1].split('.')[0], i)\n                             for i in range(len(v))],\n                     dtype = np.float32)\n    df.index.name = 'Id'\n    preds.append(df)\npreds = pd.concat(preds)#.reset_index()\nprint(len(preds))\npreds.tail()","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:10:33.071368Z","iopub.execute_input":"2023-07-20T01:10:33.072004Z","iopub.status.idle":"2023-07-20T01:10:33.455441Z","shell.execute_reply.started":"2023-07-20T01:10:33.071964Z","shell.execute_reply":"2023-07-20T01:10:33.454452Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if len(fog_files) == 2:\n    preds = preds[:-1]#.head(200)","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:10:33.457206Z","iopub.execute_input":"2023-07-20T01:10:33.457977Z","iopub.status.idle":"2023-07-20T01:10:33.467169Z","shell.execute_reply.started":"2023-07-20T01:10:33.457938Z","shell.execute_reply":"2023-07-20T01:10:33.465924Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if len(preds) != len(sample): \n    preds = pd.concat((preds.reindex(\n        sorted(list(set(sample.index) & set(preds.index))) ), \n                    sample.reindex(\n        sorted(list(set(sample.index) - set(preds.index))) ) ))\nlen(preds)","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:10:33.470362Z","iopub.execute_input":"2023-07-20T01:10:33.470772Z","iopub.status.idle":"2023-07-20T01:10:34.217584Z","shell.execute_reply.started":"2023-07-20T01:10:33.470737Z","shell.execute_reply":"2023-07-20T01:10:34.216632Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preds.tail(5)","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:10:34.218999Z","iopub.execute_input":"2023-07-20T01:10:34.219451Z","iopub.status.idle":"2023-07-20T01:10:34.231114Z","shell.execute_reply.started":"2023-07-20T01:10:34.219416Z","shell.execute_reply":"2023-07-20T01:10:34.229917Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"assert len(preds) == len(sample) and all(preds.columns == sample.columns)\nassert sorted(preds.index.tolist()) == sorted(sample.index.tolist())","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:10:34.233082Z","iopub.execute_input":"2023-07-20T01:10:34.233899Z","iopub.status.idle":"2023-07-20T01:10:34.289929Z","shell.execute_reply.started":"2023-07-20T01:10:34.233847Z","shell.execute_reply":"2023-07-20T01:10:34.288904Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preds.to_csv('submission.csv')#, index = False)","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:10:34.291544Z","iopub.execute_input":"2023-07-20T01:10:34.292016Z","iopub.status.idle":"2023-07-20T01:10:36.071874Z","shell.execute_reply.started":"2023-07-20T01:10:34.291982Z","shell.execute_reply":"2023-07-20T01:10:36.070868Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preds2 = pd.read_csv('submission.csv')\nsample2 = pd.read_csv(getFile('sample_submission.csv'))\nassert len(preds2) == len(sample2) and all(preds2.columns == sample2.columns)\nassert sorted(preds2.index.tolist()) == sorted(sample2.index.tolist())","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:10:36.073421Z","iopub.execute_input":"2023-07-20T01:10:36.073787Z","iopub.status.idle":"2023-07-20T01:10:36.556158Z","shell.execute_reply.started":"2023-07-20T01:10:36.073752Z","shell.execute_reply":"2023-07-20T01:10:36.555166Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for k in fog_files[:2]:\n    plt.figure()\n    plt.plot(all_pred_dict[k] / ct)","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:10:36.557607Z","iopub.execute_input":"2023-07-20T01:10:36.557994Z","iopub.status.idle":"2023-07-20T01:10:37.191698Z","shell.execute_reply.started":"2023-07-20T01:10:36.557958Z","shell.execute_reply":"2023-07-20T01:10:37.190581Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"np.concatenate(list(all_pred_dict.values())).mean(0)","metadata":{"execution":{"iopub.status.busy":"2023-07-20T01:10:37.194942Z","iopub.execute_input":"2023-07-20T01:10:37.195337Z","iopub.status.idle":"2023-07-20T01:10:37.206369Z","shell.execute_reply.started":"2023-07-20T01:10:37.195301Z","shell.execute_reply":"2023-07-20T01:10:37.204672Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}