{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"},{"sourceId":4735408,"sourceType":"datasetVersion","datasetId":2740254},{"sourceId":8109933,"sourceType":"datasetVersion","datasetId":4771459}],"dockerImageVersionId":30683,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# 5th place solution of HMS competition - D.Imanishi's part\nThis notebook runs inference of the models trained in \"graph model training\" notebook.\n\nThe score is inferior because it is only D.Imanishi's part, not the whole team's.\n\n#### Other notebooks of my solution\n- [HMS waveform image generation](https://www.kaggle.com/code/dimanishi/hms-waveform-image-generation)\n- [HMS graph model training](https://www.kaggle.com/code/dimanishi/hms-graph-model-training)","metadata":{}},{"cell_type":"code","source":"!pip install -q --no-index --find-links=/kaggle/input/effnet-whl efficientnet","metadata":{"execution":{"iopub.status.busy":"2024-04-13T16:12:43.225480Z","iopub.execute_input":"2024-04-13T16:12:43.225843Z","iopub.status.idle":"2024-04-13T16:12:56.235148Z","shell.execute_reply.started":"2024-04-13T16:12:43.225812Z","shell.execute_reply":"2024-04-13T16:12:56.234105Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport glob\nimport os\nimport matplotlib.pyplot as plt\nimport cv2\nimport tensorflow as tf\nimport tensorflow.keras.layers as layers\nfrom tensorflow.keras.models import Model\nimport math\nfrom tqdm import tqdm\nfrom multiprocessing import Pool\nimport io\nimport librosa\nimport random\nfrom scipy.signal import butter, lfilter\nimport efficientnet.tfkeras as efn\n\nprint('TF version:', tf.__version__)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-04-13T16:12:56.237441Z","iopub.execute_input":"2024-04-13T16:12:56.238191Z","iopub.status.idle":"2024-04-13T16:13:09.681348Z","shell.execute_reply.started":"2024-04-13T16:12:56.238150Z","shell.execute_reply":"2024-04-13T16:13:09.680380Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DEBUG = False\n\nTARGET_COLS = ['seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']\n\nif DEBUG:\n    test_df = pd.read_csv('/kaggle/input/hms-harmful-brain-activity-classification/train.csv')\n    test_df = test_df[test_df['spectrogram_label_offset_seconds']==0]\n    test_df = test_df.sample(1000, random_state=0)[['spectrogram_id', 'eeg_id', 'patient_id'] + TARGET_COLS]\n    test_df = test_df.reset_index(drop=True)\n    \n    # prepare gts\n    sum_values = test_df[TARGET_COLS].sum(axis=1).values\n    for col in TARGET_COLS:\n        test_df[col] = test_df[col].astype(np.float32) / sum_values\n    gts = test_df[TARGET_COLS].values\n    test_df = test_df.drop(TARGET_COLS, axis=1)\n    \n    SPEC_DIR = '/kaggle/input/hms-harmful-brain-activity-classification/train_spectrograms'\n    EEG_DIR = '/kaggle/input/hms-harmful-brain-activity-classification/train_eegs'\nelse:\n    test_df = pd.read_csv('/kaggle/input/hms-harmful-brain-activity-classification/test.csv')\n    SPEC_DIR = '/kaggle/input/hms-harmful-brain-activity-classification/test_spectrograms'\n    EEG_DIR = '/kaggle/input/hms-harmful-brain-activity-classification/test_eegs'\n    \nsub_df = pd.read_csv('/kaggle/input/hms-harmful-brain-activity-classification/sample_submission.csv')\n\n\nGRAPH_SHAPE = (512, 1024, 3)\nSPEC_SHAPE = (100, 300, 4)\nBS = 4\nUSE_TTA = True\n\nGRAPH_SPEC_WEIGHTS = [\n    {'filepath': '/kaggle/input/hms-share/hms_graph_spec_model_effv2b0_fold0_score0.2278.h5',\n     'backbone': 'effv2b0'},\n    {'filepath': '/kaggle/input/hms-share/hms_graph_spec_model_effv2b0_fold1_score0.2229.h5',\n     'backbone': 'effv2b0'},\n    {'filepath': '/kaggle/input/hms-share/hms_graph_spec_model_effv2b0_fold2_score0.2171.h5',\n     'backbone': 'effv2b0'},\n    {'filepath': '/kaggle/input/hms-share/hms_graph_spec_model_effv2b0_fold3_score0.2325.h5',\n     'backbone': 'effv2b0'},\n    {'filepath': '/kaggle/input/hms-share/hms_graph_spec_model_effv2b0_fold4_score0.2365.h5',\n     'backbone': 'effv2b0'},\n]\n\nGRAPH_ONLY_WEIGHTS = [\n    {'filepath': '/kaggle/input/hms-share/hms_graph_only_model_effv2b0_fold0_score0.2441.h5',\n     'backbone': 'effv2b0'},\n    {'filepath': '/kaggle/input/hms-share/hms_graph_only_model_effv2b0_fold1_score0.2312.h5',\n     'backbone': 'effv2b0'},\n    {'filepath': '/kaggle/input/hms-share/hms_graph_only_model_effv2b0_fold2_score0.2321.h5',\n     'backbone': 'effv2b0'},\n    {'filepath': '/kaggle/input/hms-share/hms_graph_only_model_effv2b0_fold3_score0.2417.h5',\n     'backbone': 'effv2b0'},\n    {'filepath': '/kaggle/input/hms-share/hms_graph_only_model_effv2b0_fold4_score0.2479.h5',\n     'backbone': 'effv2b0'},\n]\n\nLL_PAIRS = [\n    ['Fp1', 'F7'],\n    ['F7', 'T3'],\n    ['T3', 'T5'],\n    ['T5', 'O1'],\n]\nRL_PAIRS = [\n    ['Fp2', 'F8'],\n    ['F8', 'T4'],\n    ['T4', 'T6'],\n    ['T6', 'O2'],\n]\nLP_PAIRS = [\n    ['Fp1', 'F3'],\n    ['F3', 'C3'],\n    ['C3', 'P3'],\n    ['P3', 'O1'],\n]\nRP_PAIRS = [\n    ['Fp2', 'F4'],\n    ['F4', 'C4'],\n    ['C4', 'P4'],\n    ['P4', 'O2'],\n]\nEEG_SAMPLE_SEC = 200\nEEG_TIMESPAN = 20.\nCUT_OFF = 20\nCLIP = 400\nV_OFFSET = 500\nLW = 0.5\nEDGE_CROP = 12\n\nGRAPH_MEAN = [0.5, 0.5, 0.5]\nGRAPH_STD = [0.5, 0.5, 0.5]\n\n\nTMP_DIR = '/kaggle_tmp'\n\ntest_df","metadata":{"execution":{"iopub.status.busy":"2024-04-13T16:13:09.682707Z","iopub.execute_input":"2024-04-13T16:13:09.683380Z","iopub.status.idle":"2024-04-13T16:13:09.731398Z","shell.execute_reply.started":"2024-04-13T16:13:09.683344Z","shell.execute_reply":"2024-04-13T16:13:09.730400Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Generate graph images","metadata":{}},{"cell_type":"code","source":"# prepare color of 19 signals\ncmap = plt.get_cmap('jet', 19)\ncmap = [cmap(i) for i in range(19)]\nrandom.seed(0)\nrandom.shuffle(cmap)","metadata":{"execution":{"iopub.status.busy":"2024-04-13T16:13:09.733567Z","iopub.execute_input":"2024-04-13T16:13:09.733931Z","iopub.status.idle":"2024-04-13T16:13:09.744437Z","shell.execute_reply.started":"2024-04-13T16:13:09.733905Z","shell.execute_reply":"2024-04-13T16:13:09.743521Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"SAVE_DIR = f'{TMP_DIR}/graphs'\nos.makedirs(f'{SAVE_DIR}/raw', exist_ok=True)\nos.makedirs(f'{SAVE_DIR}/lpf', exist_ok=True)\n\ndef butter_lowpass_filter(data, cutoff_freq=CUT_OFF, order=5):\n    nyquist = 0.5 * EEG_SAMPLE_SEC\n    normal_cutoff = cutoff_freq / nyquist\n    b, a = butter(order, normal_cutoff, btype='low', analog=False)\n    filtered_data = lfilter(b, a, data, axis=0)\n    return filtered_data\n\n\ndef clip_data(arr, is_negative=False, use_lpf=False):\n    if use_lpf:\n        if np.isnan(arr).sum() > 0:\n            arr = np.nan_to_num(arr, nan=np.nanmean(arr))\n        arr = butter_lowpass_filter(arr)\n    \n    arr = np.clip(arr, -CLIP, CLIP)\n    if is_negative:\n        arr = -arr\n    return arr\n\n\ndef draw_wave_graph(eeg_df_part, is_lrflip=False, is_negative=False, use_lpf=False):\n    if is_lrflip:\n        pairs = RL_PAIRS + LL_PAIRS + RP_PAIRS + LP_PAIRS\n    else:\n        pairs = LL_PAIRS + RL_PAIRS + LP_PAIRS + RP_PAIRS\n    \n    plt.figure(\n        figsize=((GRAPH_SHAPE[1]+EDGE_CROP*2)/100, (GRAPH_SHAPE[0]+EDGE_CROP*2)/100),\n        tight_layout=True)\n    \n    for i, pair in enumerate(pairs):\n        arr = clip_data(\n            eeg_df_part[pair[0]].values - eeg_df_part[pair[1]].values,\n            is_negative=is_negative, use_lpf=use_lpf)\n        plt.plot(arr - i*V_OFFSET, lw=LW, color=cmap[i])\n    \n    arr = clip_data(\n        eeg_df_part['Fz'].values - eeg_df_part['Cz'].values,\n        is_negative=is_negative, use_lpf=use_lpf)\n    plt.plot(arr - (i+1)*V_OFFSET, lw=LW, color=cmap[i+1])\n\n    arr = clip_data(\n        eeg_df_part['Cz'].values - eeg_df_part['Pz'].values,\n        is_negative=is_negative, use_lpf=use_lpf)\n    plt.plot(arr - (i+2)*V_OFFSET, lw=LW, color=cmap[i+2])\n\n    arr = clip_data(\n        eeg_df_part['EKG'].values,\n        is_negative=is_negative, use_lpf=use_lpf)\n    plt.plot(arr - (i+3)*V_OFFSET, lw=LW, color=cmap[i+3])\n\n    plt.xlim(0, EEG_SAMPLE_SEC*EEG_TIMESPAN)\n    plt.ylim(int(-19*V_OFFSET), int(0.5*V_OFFSET))\n    plt.axis('off')\n\n    buf = io.BytesIO()\n    plt.savefig(buf, format=\"png\")\n    buf.seek(0)\n    img_arr = np.frombuffer(buf.getvalue(), dtype=np.uint8)\n    buf.close()\n    img = cv2.imdecode(img_arr, 1)\n    plt.clf()\n    plt.close()\n    return img[EDGE_CROP:-EDGE_CROP, EDGE_CROP:-EDGE_CROP, :]\n\n\ndef generate_graph_files(eeg_id):\n    eeg_df = pd.read_parquet(f'{EEG_DIR}/{eeg_id}.parquet')\n    \n    start_idx = 0\n    timespan_start_idx = start_idx + int((25 - EEG_TIMESPAN/2)*EEG_SAMPLE_SEC)\n    timespan_end_idx = start_idx + int((25 + EEG_TIMESPAN/2)*EEG_SAMPLE_SEC) - 1\n    eeg_df_part = eeg_df.loc[timespan_start_idx:timespan_end_idx].copy()\n\n    lpf_ptns = [\n        dict(dir_path=f'{SAVE_DIR}/raw', use_lpf=False),\n        dict(dir_path=f'{SAVE_DIR}/lpf', use_lpf=True),\n    ]\n    for ptn in lpf_ptns:\n        use_lpf = ptn['use_lpf']\n        dir_path = ptn['dir_path']\n\n        img = draw_wave_graph(eeg_df_part, is_lrflip=False, is_negative=False, use_lpf=use_lpf)\n        cv2.imwrite(f'{dir_path}/{eeg_id}.png', img)\n\n        # for augmentation\n        img = draw_wave_graph(eeg_df_part, is_lrflip=True, is_negative=False, use_lpf=use_lpf)\n        cv2.imwrite(f'{dir_path}/{eeg_id}_lrflip.png', img)\n\n        img = draw_wave_graph(eeg_df_part, is_lrflip=False, is_negative=True, use_lpf=use_lpf)\n        cv2.imwrite(f'{dir_path}/{eeg_id}_neg.png', img)\n\n        img = draw_wave_graph(eeg_df_part, is_lrflip=True, is_negative=True, use_lpf=use_lpf)\n        cv2.imwrite(f'{dir_path}/{eeg_id}_lrflip_neg.png', img)\n    return None","metadata":{"execution":{"iopub.status.busy":"2024-04-13T16:13:09.745854Z","iopub.execute_input":"2024-04-13T16:13:09.746109Z","iopub.status.idle":"2024-04-13T16:13:09.768093Z","shell.execute_reply.started":"2024-04-13T16:13:09.746087Z","shell.execute_reply":"2024-04-13T16:13:09.767340Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for eeg_ids in np.array_split(test_df['eeg_id'].values, np.ceil(len(test_df)/1000)):\n    with Pool(4) as pool:\n        imap = pool.imap(generate_graph_files, eeg_ids)\n        res = list(imap)","metadata":{"execution":{"iopub.status.busy":"2024-04-13T16:13:09.769092Z","iopub.execute_input":"2024-04-13T16:13:09.769343Z","iopub.status.idle":"2024-04-13T16:13:12.195366Z","shell.execute_reply.started":"2024-04-13T16:13:09.769322Z","shell.execute_reply":"2024-04-13T16:13:12.194278Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Inference","metadata":{}},{"cell_type":"code","source":"AUTO = tf.data.experimental.AUTOTUNE\n\ndef build_test_dset(df, input_type='graph_spec', spec_time_flip=False, lr_flip=False, pn_flip=False):\n    def load_spec(idx):\n        spec_id = df.iloc[idx]['spectrogram_id']\n        spec_df = pd.read_parquet(f'{SPEC_DIR}/{spec_id}.parquet')\n        if DEBUG:\n            org_spec = spec_df.iloc[:SPEC_SHAPE[1], 1:].values\n        else:\n            org_spec = spec_df.iloc[:, 1:].values\n        org_spec = np.clip(org_spec, np.exp(-4), np.exp(8))\n        org_spec = org_spec.reshape([-1, 4, 100]) # [300, 4, 100]\n        org_spec = org_spec.transpose([2, 0, 1])  # [100, 300, 4]\n        org_spec = np.log(org_spec)\n        \n        mean = np.nanmean(org_spec)\n        std = np.nanstd(org_spec)\n        org_spec = (org_spec - mean) / (std + 1e-6)\n        org_spec = np.nan_to_num(org_spec, nan=0.)\n        return org_spec\n    \n    def load(idx, raw_path, lpf_path):\n        if input_type == 'graph_spec':\n            org_spec = tf.numpy_function(func=load_spec, inp=[idx], Tout=tf.float32)\n            org_spec = tf.ensure_shape(org_spec, SPEC_SHAPE)\n            # [LL, RL, LP, RP] => [LL, LP, RL, RP]\n            org_spec = tf.stack([org_spec[:, :, 0], org_spec[:, :, 2], org_spec[:, :, 1], org_spec[:, :, 3]], 2)\n            if spec_time_flip:\n                org_spec = org_spec[:, ::-1, :]\n            \n            if lr_flip:\n                org_spec = tf.stack([org_spec[:, :, 2], org_spec[:, :, 3], org_spec[:, :, 0], org_spec[:, :, 1]], 2)\n        elif input_type == 'graph_only':\n            org_spec = None\n            \n        img1 = tf.io.read_file(raw_path)\n        img1 = tf.image.decode_png(img1, channels=3)\n        img1 = tf.ensure_shape(img1, GRAPH_SHAPE)\n        \n        img2 = tf.io.read_file(lpf_path)\n        img2 = tf.image.decode_png(img2, channels=3)\n        img2 = tf.ensure_shape(img2, GRAPH_SHAPE)\n        \n        img = tf.concat([img1, img2], 0)\n        img = tf.cast(img, tf.float32) / 255.\n        if input_type == 'graph_spec':\n            img = (img - GRAPH_MEAN) / GRAPH_STD\n        return img, org_spec\n    \n    def concat_img(img, org_spec):\n        if input_type == 'graph_spec':\n            org_spec = tf.transpose(org_spec, [2, 0, 1]) # [4, 100, 300]\n            org_spec = tf.reshape(org_spec, [SPEC_SHAPE[0]*SPEC_SHAPE[2], SPEC_SHAPE[1], 1]) # [400, 300, 1]\n\n            pad_y_size = GRAPH_SHAPE[0]*2 - (SPEC_SHAPE[0]*SPEC_SHAPE[2])\n            org_spec = tf.pad(org_spec, [[0, pad_y_size], [0, 0], [0, 0]]) # [1024, 300, 1]\n            org_spec = tf.concat([org_spec]*3, -1) # [1024, 300, 3]\n\n            img = tf.concat([org_spec, img], 1)\n        return img\n    \n    raw_paths = df['eeg_id'].apply(lambda x: f'{SAVE_DIR}/raw/{int(x)}.png')\n    lpf_paths = df['eeg_id'].apply(lambda x: f'{SAVE_DIR}/lpf/{int(x)}.png')\n    if lr_flip:\n        raw_paths = [x.replace('.png', '_lrflip.png') for x in raw_paths]\n        lpf_paths = [x.replace('.png', '_lrflip.png') for x in lpf_paths]\n    if pn_flip:\n        raw_paths = [x.replace('.png', '_neg.png') for x in raw_paths]\n        lpf_paths = [x.replace('.png', '_neg.png') for x in lpf_paths]\n    \n    dset = tf.data.Dataset.from_tensor_slices((range(len(df)), raw_paths, lpf_paths))\n    dset = dset.map(load, num_parallel_calls=AUTO)\n    dset = dset.map(concat_img, num_parallel_calls=AUTO)\n    dset = dset.batch(BS, drop_remainder=False).prefetch(AUTO)\n    return dset","metadata":{"execution":{"iopub.status.busy":"2024-04-13T16:13:12.197158Z","iopub.execute_input":"2024-04-13T16:13:12.197462Z","iopub.status.idle":"2024-04-13T16:13:12.218962Z","shell.execute_reply.started":"2024-04-13T16:13:12.197435Z","shell.execute_reply":"2024-04-13T16:13:12.217977Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BACKBONES = {\n    'effv2b0': tf.keras.applications.EfficientNetV2B0,\n    'effv2b1': tf.keras.applications.EfficientNetV2B1,\n    'effv1b0': efn.EfficientNetB0,\n}\n\ndef build_model(backbone_name, input_type):\n    bb = BACKBONES[backbone_name]\n    if 'effv2' in backbone_name:\n        kwargs = {'include_preprocessing': False}\n    else:\n        kwargs = {}\n        \n    if input_type == 'graph_spec':\n        input_shape = (GRAPH_SHAPE[0]*2, SPEC_SHAPE[1] + GRAPH_SHAPE[1], 3)\n    elif input_type == 'graph_only':\n        input_shape = (GRAPH_SHAPE[0]*2, GRAPH_SHAPE[1], 3)\n        \n    base_model = bb(\n        include_top=False,\n        weights=None,\n        input_shape=input_shape,\n        pooling='avg',\n        **kwargs,\n    )\n\n    x = base_model.output\n    x = layers.Dense(len(TARGET_COLS), name='final_dense')(x)\n    # logits averaging\n    # x = layers.Activation('softmax', name='softmax', dtype=tf.float32)(x)\n    model = Model(inputs=base_model.input, outputs=x)\n    return model\n\n\nclass EnsembleModel(Model):\n    def __init__(self, models):\n        super().__init__()\n        self.models = models\n        \n    def call(self, inputs, training=None):\n        preds = []\n        for m in self.models:\n            preds.append(m(inputs))\n            \n        preds = tf.stack(preds, axis=0)\n        preds = tf.math.reduce_mean(preds, axis=0)\n        return preds","metadata":{"execution":{"iopub.status.busy":"2024-04-13T16:13:12.220150Z","iopub.execute_input":"2024-04-13T16:13:12.220438Z","iopub.status.idle":"2024-04-13T16:13:12.233977Z","shell.execute_reply.started":"2024-04-13T16:13:12.220414Z","shell.execute_reply":"2024-04-13T16:13:12.233121Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"input_type = 'graph_spec'\n\nif USE_TTA:\n    tta_patterns = [\n        dict(spec_time_flip=False, lr_flip=False, pn_flip=False),\n        dict(spec_time_flip=False, lr_flip=True, pn_flip=False),\n        dict(spec_time_flip=False, lr_flip=False, pn_flip=True),\n        dict(spec_time_flip=False, lr_flip=True, pn_flip=True),\n        dict(spec_time_flip=True, lr_flip=False, pn_flip=False),\n        dict(spec_time_flip=True, lr_flip=True, pn_flip=False),\n        dict(spec_time_flip=True, lr_flip=False, pn_flip=True),\n        dict(spec_time_flip=True, lr_flip=True, pn_flip=True),\n    ]\nelse:\n    tta_patterns = [\n        dict(spec_time_flip=False, lr_flip=False, pn_flip=False),\n    ]\n\nmodels = []\nfor w in GRAPH_SPEC_WEIGHTS:\n    model = build_model(w['backbone'], input_type)\n    model.load_weights(w['filepath'])\n    models.append(model)\nmodel = EnsembleModel(models)\n\ngraph_spec_preds = None\nfor ptn in tta_patterns:\n    dset = build_test_dset(test_df, input_type=input_type, **ptn)\n    preds = model.predict(dset)\n    if graph_spec_preds is None:\n        graph_spec_preds = preds\n    else:\n        graph_spec_preds += preds\ngraph_spec_preds /= len(tta_patterns)\n\n\nif DEBUG:\n    kl = tf.keras.losses.KLDivergence()\n    score = kl(gts, tf.nn.softmax(graph_spec_preds)).numpy()\n    print(f'Debug score: {score:.4f}')","metadata":{"execution":{"iopub.status.busy":"2024-04-13T16:13:12.235062Z","iopub.execute_input":"2024-04-13T16:13:12.235299Z","iopub.status.idle":"2024-04-13T16:13:57.212118Z","shell.execute_reply.started":"2024-04-13T16:13:12.235277Z","shell.execute_reply":"2024-04-13T16:13:57.211282Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"input_type = 'graph_only'\n\nif USE_TTA:\n    tta_patterns = [\n        dict(spec_time_flip=False, lr_flip=False),\n        dict(spec_time_flip=False, lr_flip=True),\n        dict(spec_time_flip=True, lr_flip=False),\n        dict(spec_time_flip=True, lr_flip=True),\n    ]\nelse:\n    tta_patterns = [\n        dict(spec_time_flip=False, lr_flip=False),\n    ]\n    \nmodels = []\nfor w in GRAPH_ONLY_WEIGHTS:\n    model = build_model(w['backbone'], input_type)\n    model.load_weights(w['filepath'])\n    models.append(model)\nmodel = EnsembleModel(models)\n\ngraph_only_preds = None\nfor ptn in tta_patterns:\n    dset = build_test_dset(test_df, input_type=input_type, **ptn)\n    preds = model.predict(dset)\n    if graph_only_preds is None:\n        graph_only_preds = preds\n    else:\n        graph_only_preds += preds\ngraph_only_preds /= len(tta_patterns)\n\n\nif DEBUG:\n    kl = tf.keras.losses.KLDivergence()\n    score = kl(gts, tf.nn.softmax(graph_only_preds)).numpy()\n    print(f'Debug score: {score:.4f}')","metadata":{"execution":{"iopub.status.busy":"2024-04-13T16:13:57.214487Z","iopub.execute_input":"2024-04-13T16:13:57.214795Z","iopub.status.idle":"2024-04-13T16:14:40.540892Z","shell.execute_reply.started":"2024-04-13T16:13:57.214770Z","shell.execute_reply":"2024-04-13T16:14:40.539968Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df[TARGET_COLS] = tf.nn.softmax((graph_spec_preds + graph_only_preds) / 2).numpy()\n\nsub_df = sub_df[['eeg_id']].merge(test_df[['eeg_id'] + TARGET_COLS])\nsub_df.to_csv('submission.csv', index=False)\nsub_df","metadata":{"execution":{"iopub.status.busy":"2024-04-13T16:14:40.542068Z","iopub.execute_input":"2024-04-13T16:14:40.542343Z","iopub.status.idle":"2024-04-13T16:14:40.587508Z","shell.execute_reply.started":"2024-04-13T16:14:40.542319Z","shell.execute_reply":"2024-04-13T16:14:40.586577Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n","metadata":{},"execution_count":null,"outputs":[]}]}