{"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-06-08T23:17:43.227015Z","iopub.execute_input":"2023-06-08T23:17:43.228036Z","iopub.status.idle":"2023-06-08T23:17:43.232578Z","shell.execute_reply.started":"2023-06-08T23:17:43.227992Z","shell.execute_reply":"2023-06-08T23:17:43.230810Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# FAST = False\nSAMPLE = 4","metadata":{"execution":{"iopub.status.busy":"2023-06-08T23:17:43.234650Z","iopub.execute_input":"2023-06-08T23:17:43.235038Z","iopub.status.idle":"2023-06-08T23:17:43.243778Z","shell.execute_reply.started":"2023-06-08T23:17:43.234995Z","shell.execute_reply":"2023-06-08T23:17:43.242788Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# DATA_PATH = '/data/'\nDATASET = 'walkdata4'\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-06-08T23:17:43.245492Z","iopub.execute_input":"2023-06-08T23:17:43.245888Z","iopub.status.idle":"2023-06-08T23:17:43.253089Z","shell.execute_reply.started":"2023-06-08T23:17:43.245857Z","shell.execute_reply":"2023-06-08T23:17:43.252151Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls $CODE_PATH","metadata":{"execution":{"iopub.status.busy":"2023-06-08T23:17:43.255483Z","iopub.execute_input":"2023-06-08T23:17:43.256234Z","iopub.status.idle":"2023-06-08T23:17:44.280444Z","shell.execute_reply.started":"2023-06-08T23:17:43.256201Z","shell.execute_reply":"2023-06-08T23:17:44.279284Z"},"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-06-08T23:17:44.282273Z","iopub.execute_input":"2023-06-08T23:17:44.282895Z","iopub.status.idle":"2023-06-08T23:18:15.375501Z","shell.execute_reply.started":"2023-06-08T23:17:44.282857Z","shell.execute_reply":"2023-06-08T23:18:15.374083Z"},"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-06-08T23:18:15.379456Z","iopub.execute_input":"2023-06-08T23:18:15.379815Z","iopub.status.idle":"2023-06-08T23:18:15.386923Z","shell.execute_reply.started":"2023-06-08T23:18:15.379782Z","shell.execute_reply":"2023-06-08T23:18:15.385810Z"},"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-06-08T23:18:15.388569Z","iopub.execute_input":"2023-06-08T23:18:15.389007Z","iopub.status.idle":"2023-06-08T23:18:17.426866Z","shell.execute_reply.started":"2023-06-08T23:18:15.388975Z","shell.execute_reply":"2023-06-08T23:18:17.425907Z"},"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-06-08T23:18:17.428169Z","iopub.execute_input":"2023-06-08T23:18:17.428962Z","iopub.status.idle":"2023-06-08T23:18:17.447655Z","shell.execute_reply.started":"2023-06-08T23:18:17.428929Z","shell.execute_reply":"2023-06-08T23:18:17.446780Z"},"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-06-08T23:18:17.448943Z","iopub.execute_input":"2023-06-08T23:18:17.449944Z","iopub.status.idle":"2023-06-08T23:18:17.456618Z","shell.execute_reply.started":"2023-06-08T23:18:17.449910Z","shell.execute_reply":"2023-06-08T23:18:17.455676Z"},"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-06-08T23:18:17.458004Z","iopub.execute_input":"2023-06-08T23:18:17.458525Z","iopub.status.idle":"2023-06-08T23:18:17.470043Z","shell.execute_reply.started":"2023-06-08T23:18:17.458495Z","shell.execute_reply":"2023-06-08T23:18:17.469122Z"},"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-06-08T23:18:17.474466Z","iopub.execute_input":"2023-06-08T23:18:17.474763Z","iopub.status.idle":"2023-06-08T23:18:17.726659Z","shell.execute_reply.started":"2023-06-08T23:18:17.474721Z","shell.execute_reply":"2023-06-08T23:18:17.725733Z"},"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-06-08T23:18:17.728167Z","iopub.execute_input":"2023-06-08T23:18:17.728502Z","iopub.status.idle":"2023-06-08T23:18:17.749151Z","shell.execute_reply.started":"2023-06-08T23:18:17.728470Z","shell.execute_reply":"2023-06-08T23:18:17.748107Z"},"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-06-08T23:18:17.750788Z","iopub.execute_input":"2023-06-08T23:18:17.751141Z","iopub.status.idle":"2023-06-08T23:18:17.761400Z","shell.execute_reply.started":"2023-06-08T23:18:17.751110Z","shell.execute_reply":"2023-06-08T23:18:17.760402Z"},"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-06-08T23:18:17.762679Z","iopub.execute_input":"2023-06-08T23:18:17.763709Z","iopub.status.idle":"2023-06-08T23:18:17.775000Z","shell.execute_reply.started":"2023-06-08T23:18:17.763650Z","shell.execute_reply":"2023-06-08T23:18:17.774132Z"},"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-06-08T23:18:17.776234Z","iopub.execute_input":"2023-06-08T23:18:17.776672Z","iopub.status.idle":"2023-06-08T23:18:17.785926Z","shell.execute_reply.started":"2023-06-08T23:18:17.776640Z","shell.execute_reply":"2023-06-08T23:18:17.785003Z"},"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-06-08T23:18:17.787556Z","iopub.execute_input":"2023-06-08T23:18:17.787906Z","iopub.status.idle":"2023-06-08T23:18:18.872923Z","shell.execute_reply.started":"2023-06-08T23:18:17.787875Z","shell.execute_reply":"2023-06-08T23:18:18.871830Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!du -sh {DATA_PATH}cache/test*","metadata":{"execution":{"iopub.status.busy":"2023-06-08T23:18:18.877382Z","iopub.execute_input":"2023-06-08T23:18:18.880353Z","iopub.status.idle":"2023-06-08T23:18:19.890926Z","shell.execute_reply.started":"2023-06-08T23:18:18.880313Z","shell.execute_reply":"2023-06-08T23:18:19.889686Z"},"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-06-08T23:18:19.893190Z","iopub.execute_input":"2023-06-08T23:18:19.893584Z","iopub.status.idle":"2023-06-08T23:18:19.913498Z","shell.execute_reply.started":"2023-06-08T23:18:19.893548Z","shell.execute_reply":"2023-06-08T23:18:19.912659Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%run -i {CODE_PATH}/model.py","metadata":{"execution":{"iopub.status.busy":"2023-06-08T23:18:19.914742Z","iopub.execute_input":"2023-06-08T23:18:19.915070Z","iopub.status.idle":"2023-06-08T23:18:19.935238Z","shell.execute_reply.started":"2023-06-08T23:18:19.915047Z","shell.execute_reply":"2023-06-08T23:18:19.934373Z"},"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-06-08T23:18:19.936691Z","iopub.execute_input":"2023-06-08T23:18:19.937093Z","iopub.status.idle":"2023-06-08T23:18:19.943727Z","shell.execute_reply.started":"2023-06-08T23:18:19.937060Z","shell.execute_reply":"2023-06-08T23:18:19.942677Z"},"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-06-08T23:18:19.945390Z","iopub.execute_input":"2023-06-08T23:18:19.946091Z","iopub.status.idle":"2023-06-08T23:18:19.953060Z","shell.execute_reply.started":"2023-06-08T23:18:19.946059Z","shell.execute_reply":"2023-06-08T23:18:19.952124Z"},"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-06-08T23:18:19.954740Z","iopub.execute_input":"2023-06-08T23:18:19.955643Z","iopub.status.idle":"2023-06-08T23:18:19.970809Z","shell.execute_reply.started":"2023-06-08T23:18:19.955558Z","shell.execute_reply":"2023-06-08T23:18:19.969898Z"},"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-06-08T23:18:19.973687Z","iopub.execute_input":"2023-06-08T23:18:19.974098Z","iopub.status.idle":"2023-06-08T23:18:19.992253Z","shell.execute_reply.started":"2023-06-08T23:18:19.974074Z","shell.execute_reply":"2023-06-08T23:18:19.991345Z"},"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-06-08T23:18:19.993739Z","iopub.execute_input":"2023-06-08T23:18:19.994939Z","iopub.status.idle":"2023-06-08T23:18:20.001631Z","shell.execute_reply.started":"2023-06-08T23:18:19.994905Z","shell.execute_reply":"2023-06-08T23:18:20.000687Z"},"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-06-08T23:18:20.003020Z","iopub.execute_input":"2023-06-08T23:18:20.003489Z","iopub.status.idle":"2023-06-08T23:18:20.244467Z","shell.execute_reply.started":"2023-06-08T23:18:20.003458Z","shell.execute_reply":"2023-06-08T23:18:20.243434Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"select_ms = '0c1908, 38f97e, 6724d4, b904fe, 4e7844, b792a2, d64c60, d85bf1, 798bc8, c5be66, cd590d, d7b7b6, 2233bd, 2c1d4a, b42aae, ef4805, f41ae6, 18da64, 37898f, 9805ce, f65f09, 53b767, 7033ef, 90be5b, fdbb0a, 36b30a, 94cae9, cbf5ad, fdd535, 25f680, b8cf10, c8d569, faed9c, 91cee9, a33a0c, a4745d, a64923, c89baf, 1ac968, 22693b, e1d1a2, fc8f37'\nselect_ms = select_ms.split(', ')","metadata":{"execution":{"iopub.status.busy":"2023-06-08T23:18:20.249512Z","iopub.execute_input":"2023-06-08T23:18:20.249826Z","iopub.status.idle":"2023-06-08T23:18:20.254272Z","shell.execute_reply.started":"2023-06-08T23:18:20.249800Z","shell.execute_reply":"2023-06-08T23:18:20.253247Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from collections import defaultdict\nmgroups = defaultdict(list)\nmpoints = defaultdict(int)\nfor m, p in ms:\n    if m not in select_ms: continue\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-06-08T23:18:34.893889Z","iopub.execute_input":"2023-06-08T23:18:34.894250Z","iopub.status.idle":"2023-06-08T23:18:34.907042Z","shell.execute_reply.started":"2023-06-08T23:18:34.894221Z","shell.execute_reply":"2023-06-08T23:18:34.905917Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.hist(mpoints.values(), bins =50);","metadata":{"execution":{"iopub.status.busy":"2023-06-08T23:22:02.265125Z","iopub.execute_input":"2023-06-08T23:22:02.265486Z","iopub.status.idle":"2023-06-08T23:22:02.600468Z","shell.execute_reply.started":"2023-06-08T23:22:02.265459Z","shell.execute_reply":"2023-06-08T23:22:02.599577Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sum(mpoints.values())","metadata":{"execution":{"iopub.status.busy":"2023-06-08T23:20:41.326038Z","iopub.execute_input":"2023-06-08T23:20:41.326412Z","iopub.status.idle":"2023-06-08T23:20:41.333688Z","shell.execute_reply.started":"2023-06-08T23:20:41.326382Z","shell.execute_reply":"2023-06-08T23:20:41.332798Z"},"trusted":true},"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-06-08T23:19:10.054672Z","iopub.execute_input":"2023-06-08T23:19:10.055582Z","iopub.status.idle":"2023-06-08T23:19:10.062429Z","shell.execute_reply.started":"2023-06-08T23:19:10.055536Z","shell.execute_reply":"2023-06-08T23:19:10.061464Z"},"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-06-08T23:19:10.844558Z","iopub.execute_input":"2023-06-08T23:19:10.844925Z","iopub.status.idle":"2023-06-08T23:19:10.852315Z","shell.execute_reply.started":"2023-06-08T23:19:10.844896Z","shell.execute_reply":"2023-06-08T23:19:10.851407Z"},"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-06-08T23:19:12.247175Z","iopub.execute_input":"2023-06-08T23:19:12.247535Z","iopub.status.idle":"2023-06-08T23:19:12.252526Z","shell.execute_reply.started":"2023-06-08T23:19:12.247500Z","shell.execute_reply":"2023-06-08T23:19:12.251520Z"},"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[_] ** 0.8)\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-06-08T23:19:43.365186Z","iopub.execute_input":"2023-06-08T23:19:43.365551Z","iopub.status.idle":"2023-06-08T23:20:13.322164Z","shell.execute_reply.started":"2023-06-08T23:19:43.365519Z","shell.execute_reply":"2023-06-08T23:20:13.320839Z"},"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-06-08T23:20:13.324513Z","iopub.execute_input":"2023-06-08T23:20:13.324839Z","iopub.status.idle":"2023-06-08T23:20:13.676802Z","shell.execute_reply.started":"2023-06-08T23:20:13.324809Z","shell.execute_reply":"2023-06-08T23:20:13.675694Z"},"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-06-08T23:20:13.678202Z","iopub.execute_input":"2023-06-08T23:20:13.678699Z","iopub.status.idle":"2023-06-08T23:20:13.688858Z","shell.execute_reply.started":"2023-06-08T23:20:13.678662Z","shell.execute_reply":"2023-06-08T23:20:13.687923Z"},"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-06-08T23:20:13.692372Z","iopub.execute_input":"2023-06-08T23:20:13.692687Z","iopub.status.idle":"2023-06-08T23:20:14.306070Z","shell.execute_reply.started":"2023-06-08T23:20:13.692662Z","shell.execute_reply":"2023-06-08T23:20:14.305180Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preds.tail(5)","metadata":{"execution":{"iopub.status.busy":"2023-06-08T23:20:14.307363Z","iopub.execute_input":"2023-06-08T23:20:14.307818Z","iopub.status.idle":"2023-06-08T23:20:14.318899Z","shell.execute_reply.started":"2023-06-08T23:20:14.307782Z","shell.execute_reply":"2023-06-08T23:20:14.317983Z"},"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-06-08T23:20:14.320468Z","iopub.execute_input":"2023-06-08T23:20:14.321183Z","iopub.status.idle":"2023-06-08T23:20:14.372117Z","shell.execute_reply.started":"2023-06-08T23:20:14.321149Z","shell.execute_reply":"2023-06-08T23:20:14.370886Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preds.to_csv('submission.csv')#, index = False)","metadata":{"execution":{"iopub.status.busy":"2023-06-08T23:20:14.373884Z","iopub.execute_input":"2023-06-08T23:20:14.374306Z","iopub.status.idle":"2023-06-08T23:20:15.498491Z","shell.execute_reply.started":"2023-06-08T23:20:14.374271Z","shell.execute_reply":"2023-06-08T23:20:15.497314Z"},"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-06-08T23:20:15.500792Z","iopub.execute_input":"2023-06-08T23:20:15.501499Z","iopub.status.idle":"2023-06-08T23:20:15.926984Z","shell.execute_reply.started":"2023-06-08T23:20:15.501464Z","shell.execute_reply":"2023-06-08T23:20:15.925948Z"},"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-06-08T23:20:15.928335Z","iopub.execute_input":"2023-06-08T23:20:15.928923Z","iopub.status.idle":"2023-06-08T23:20:16.543686Z","shell.execute_reply.started":"2023-06-08T23:20:15.928889Z","shell.execute_reply":"2023-06-08T23:20:16.542685Z"},"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-06-08T23:02:19.028107Z","iopub.execute_input":"2023-06-08T23:02:19.028454Z","iopub.status.idle":"2023-06-08T23:02:19.038362Z","shell.execute_reply.started":"2023-06-08T23:02:19.028420Z","shell.execute_reply":"2023-06-08T23:02:19.037063Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}