{"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":"tpu1vmV38","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"},{"sourceId":8104095,"sourceType":"datasetVersion","datasetId":4786224},{"sourceId":8109933,"sourceType":"datasetVersion","datasetId":4771459}],"dockerImageVersionId":30581,"isInternetEnabled":true,"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\nAll my models can be trained in this notebook. Graph images need to be generated in advance using the \"waveform image generation\" notebook.\n#### Usage\n1. train both \"graph_spec\" model and \"graph_only\" model on \"effv2b0\" backbone without pseudo label (`USE_PSEUDO_LABEL=False`) and get pseudo label csv file.\n2. train both \"graph_spec\" model and \"graph_only\" model on each backbone with the pseudo label generated in step1 (`USE_PSEUDO_LABEL=True`)\n\nThis notebook is intended to run on the TPU kernel.\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 inference](https://www.kaggle.com/code/dimanishi/hms-graph-model-inference)","metadata":{}},{"cell_type":"code","source":"!pip -q install transformers pyarrow efficientnet","metadata":{"execution":{"iopub.status.busy":"2024-04-14T01:53:43.857806Z","iopub.execute_input":"2024-04-14T01:53:43.858070Z","iopub.status.idle":"2024-04-14T01:53:54.048803Z","shell.execute_reply.started":"2024-04-14T01:53:43.858022Z","shell.execute_reply":"2024-04-14T01:53:54.047869Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nos.environ['TF_CPP_MIN_LOG_LEVEL'] = '3'\n\nimport numpy as np\nimport pandas as pd\nimport glob\nimport matplotlib.pyplot as plt\nimport cv2\nimport tensorflow as tf\nimport tensorflow.keras.layers as layers\nfrom tensorflow.keras.models import Model\nimport tensorflow.keras.backend as K\nimport math\nfrom transformers import AdamWeightDecay\nfrom tqdm import tqdm\nfrom IPython.display import FileLink\nimport efficientnet.tfkeras as efn\n\nprint('TF version:', tf.__version__)\n","metadata":{"execution":{"iopub.status.busy":"2024-04-14T01:53:54.050326Z","iopub.execute_input":"2024-04-14T01:53:54.050573Z","iopub.status.idle":"2024-04-14T01:54:33.668792Z","shell.execute_reply.started":"2024-04-14T01:53:54.050545Z","shell.execute_reply":"2024-04-14T01:54:33.667840Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#################### Change here as necessary ####################\nBACKBONE_NAME = 'effv2b0'   # choose from {'effv2b0', 'effv2b1', 'effv1b0'}\nINPUT_TYPE = 'graph_spec'   # choose from {'graph_spec', 'graph_only'}\nUSE_PSEUDO_LABEL = True    # 1st training => False, 2nd training => True\n##################################################################\n\n\nGRAPH_RAW_DIR = '/kaggle/input/hms-graph-images/graphs/raw'\nGRAPH_LPF_DIR = '/kaggle/input/hms-graph-images/graphs/lpf'\nSPEC_DIR = '/kaggle/input/hms-harmful-brain-activity-classification/train_spectrograms'\n\nGRAPH_SHAPE = (512, 1024, 3)\nSPEC_SHAPE = (100, 300, 4)\n\nGRAPH_MEAN = [0.5, 0.5, 0.5]\nGRAPH_STD = [0.5, 0.5, 0.5]\n\nTARGET_COLS = ['seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']\nGRAPH_TYPES = ['raw', 'lpf']\nGRAPH_AUGS = ['neg', 'lrflip', 'lrflip_neg']\n\nN_VOTES_TH = 10\nPSEUDO_WEIGHT = 5.\nPSEUDO_CSVS = [\n    '/kaggle/input/hms-share/pseudo_graph_only_model_effv2b0.csv',\n    '/kaggle/input/hms-share/pseudo_graph_spec_model_effv2b0.csv',\n]\n\nFP16 = True","metadata":{"execution":{"iopub.status.busy":"2024-04-14T01:54:33.669899Z","iopub.execute_input":"2024-04-14T01:54:33.670480Z","iopub.status.idle":"2024-04-14T01:54:33.677199Z","shell.execute_reply.started":"2024-04-14T01:54:33.670445Z","shell.execute_reply":"2024-04-14T01:54:33.676516Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"try:\n    resolver = tf.distribute.cluster_resolver.TPUClusterResolver(tpu='local')\n    tf.config.experimental_connect_to_cluster(resolver)\n    # This is the TPU initialization code that has to be at the beginning.\n    tf.tpu.experimental.initialize_tpu_system(resolver)\n    print(\"All devices: \", tf.config.list_logical_devices('TPU'))\n    strategy = tf.distribute.TPUStrategy(resolver)\n    TPU = True\nexcept:\n    if len(tf.config.list_physical_devices('GPU')) > 1:\n        strategy = tf.distribute.MirroredStrategy()\n    else:\n        strategy = tf.distribute.get_strategy()\n    TPU = False\nprint(f\"Running on {strategy.num_replicas_in_sync} replicas\")\n\nif TPU:\n    MBS = 16\n    BS = strategy.num_replicas_in_sync * MBS\n    if FP16:\n        tf.keras.mixed_precision.set_global_policy('mixed_bfloat16')\nelse:\n    BS = 32\n    if FP16:\n        tf.keras.mixed_precision.set_global_policy('mixed_float16')\nprint(f'batch size: {BS}')","metadata":{"execution":{"iopub.status.busy":"2024-04-14T01:54:33.678812Z","iopub.execute_input":"2024-04-14T01:54:33.679067Z","iopub.status.idle":"2024-04-14T01:54:43.214351Z","shell.execute_reply.started":"2024-04-14T01:54:33.679025Z","shell.execute_reply":"2024-04-14T01:54:43.213354Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_df = pd.read_csv('/kaggle/input/hms-harmful-brain-activity-classification/train.csv')\n\n# Pre-created group kfold (5 folds)\nfold_df = pd.read_csv('/kaggle/input/hms-share/fold.csv')\ndata_df = data_df.merge(fold_df[['eeg_id', 'fold']])\n\n# data sampling\ndata_df['n_votes'] = data_df[TARGET_COLS].sum(axis=1)\nrows = []\nfor _, df in tqdm(data_df.groupby('eeg_id')):\n    df_mean_target = df[TARGET_COLS].values.sum(axis=0) / df[TARGET_COLS].values.sum()\n    if len(df['expert_consensus'].unique()) > 1:\n        for _, sdf in df.groupby('expert_consensus'):\n            df2 = sdf[sdf['n_votes']==sdf['n_votes'].max()]\n            this_target = df2.iloc[len(df2)//2][TARGET_COLS].values\n            this_target /= this_target.sum()\n            row = df2.iloc[len(df2) // 2].copy()\n            row[TARGET_COLS] = (df_mean_target + this_target) / 2\n            rows.append(row)\n    else:\n        df2 = df[df['n_votes']==df['n_votes'].max()]\n        this_target = df2.iloc[len(df2)//2][TARGET_COLS].values\n        this_target /= this_target.sum()\n        row = df2.iloc[len(df2) // 2].copy()\n        row[TARGET_COLS] = (df_mean_target + this_target) / 2\n        rows.append(row)\ndata_df = pd.DataFrame(rows).sort_index()\n\ndata_df['img_idx'] = 0\nfor _, df in tqdm(data_df.groupby('eeg_id')):\n    if len(df) > 1:\n        data_df.loc[df.index, 'img_idx'] = range(len(df))\n        \ndata_df[['n_votes'] + TARGET_COLS] = data_df[['n_votes'] + TARGET_COLS].astype(np.float32)\n\ndata_df['img_path_raw'] = \\\n    f'{GRAPH_RAW_DIR}/' + data_df['eeg_id'].astype(str) + '_' + data_df['img_idx'].astype(str) + '.png'\ndata_df['img_path_lpf'] = \\\n    f'{GRAPH_LPF_DIR}/' + data_df['eeg_id'].astype(str) + '_' + data_df['img_idx'].astype(str) + '.png'\nfor typ in GRAPH_TYPES:\n    for aug in GRAPH_AUGS:\n        data_df[f'img_path_{typ}_{aug}'] = data_df[f'img_path_{typ}'].apply(\n                                                lambda x: x.replace('.png', f'_{aug}.png'))\n\nif USE_PSEUDO_LABEL:\n    targets = None\n    for c in PSEUDO_CSVS:\n        pseudo_df = pd.read_csv(c)\n        if targets is None:\n            targets = pseudo_df[TARGET_COLS].values\n        else:\n            targets += pseudo_df[TARGET_COLS].values\n\n    pseudo_df[TARGET_COLS] = targets / len(PSEUDO_CSVS)\n    pseudo_df[TARGET_COLS] = pseudo_df[TARGET_COLS].astype(np.float32)\n    pseudo_df = pseudo_df.rename(columns={x: f'p_{x}' for x in TARGET_COLS})\n    data_df = data_df.merge(pseudo_df)","metadata":{"execution":{"iopub.status.busy":"2024-04-14T01:54:43.215605Z","iopub.execute_input":"2024-04-14T01:54:43.215879Z","iopub.status.idle":"2024-04-14T01:55:21.700019Z","shell.execute_reply.started":"2024-04-14T01:54:43.215850Z","shell.execute_reply":"2024-04-14T01:55:21.699025Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data loader","metadata":{}},{"cell_type":"code","source":"AUTO = tf.data.experimental.AUTOTUNE\n\ndef build_dset(df, is_train=True, with_label=True):\n    def load_spec(idx):\n        spec_id = df.iloc[idx]['spectrogram_id']\n        spec_ofs_sec = df.iloc[idx]['spectrogram_label_offset_seconds']\n        \n        spec_df = pd.read_parquet(f'{SPEC_DIR}/{spec_id}.parquet')\n        # 10 min.\n        spec_df = spec_df[(spec_ofs_sec <= spec_df['time'])  & (spec_df['time'] < spec_ofs_sec + 600)]\n        \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_paths, lpf_paths, label):\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        else:\n            org_spec = None\n        return raw_paths, lpf_paths, org_spec, label\n    \n    def decode_train(raw_paths, lpf_paths, org_spec, label):\n        img_idx = tf.random.uniform([], 0, 4, tf.int64)\n        \n        img1 = tf.io.read_file(raw_paths[img_idx])\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_paths[img_idx])\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, label\n    \n    def decode_val(raw_paths, lpf_paths, org_spec, label):\n        img_idx = 0\n        \n        img1 = tf.io.read_file(raw_paths[img_idx])\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_paths[img_idx])\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, label\n    \n    def augment(img, org_spec, label):\n        # time inverting (image lr-flip)\n        p = tf.random.uniform([], 0., 1., tf.float32)\n        if p < 0.5:\n            org_spec = org_spec[:, ::-1, :]\n\n        # lr swap\n        p = tf.random.uniform([], 0., 1., tf.float32)\n        if p < 0.5:\n            org_spec = tf.stack([org_spec[:, :, 2], org_spec[:, :, 3], org_spec[:, :, 0], org_spec[:, :, 1]], 2)\n        return img, org_spec, label\n    \n    def concat_img(img, org_spec, label):\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        \n        if with_label:\n            return img, label\n        else:\n            return img\n    \n    def mixup(img, label):\n        p = tf.random.uniform([], 0., 1., tf.float32)\n        if p < 0.5:\n            n = tf.shape(img)[0]\n            lam = tf.random.uniform([n], 0., 1., tf.float32)\n            \n            img_slide = tf.concat([img[1:, :, :, :], img[:1, :, :, :]], 0)\n            lam_img = lam[:, None, None, None]\n            img = img * lam_img + img_slide * (1 - lam_img)\n            \n            label_slide = tf.concat([label[1:, :], label[:1, :]], 0)\n            lam_label = lam[:, None]\n            label = label * lam_label + label_slide * (1 - lam_label)\n        return img, label\n    \n        \n    raw_paths = df[['img_path_raw'] + [f'img_path_raw_{x}' for x in GRAPH_AUGS]]\n    lpf_paths = df[['img_path_lpf'] + [f'img_path_lpf_{x}' for x in GRAPH_AUGS]]\n        \n    dset = tf.data.Dataset.from_tensor_slices(\n        (range(len(df)), raw_paths, lpf_paths, df[TARGET_COLS]))\n    dset = dset.shuffle(buffer_size=BS*8) if is_train else dset\n    dset = dset.map(load, num_parallel_calls=AUTO)\n    dset = dset.cache() if TPU else dset\n    dset = dset.repeat() if is_train else dset\n    if is_train:\n        dset = dset.map(decode_train, num_parallel_calls=AUTO)\n    else:\n        dset = dset.map(decode_val, num_parallel_calls=AUTO)\n    if is_train and (INPUT_TYPE == 'graph_spec'):\n        dset = dset.map(augment, num_parallel_calls=AUTO)\n    dset = dset.map(concat_img, num_parallel_calls=AUTO)\n    dset = dset.batch(BS, drop_remainder=is_train).prefetch(AUTO)\n    dset = dset.map(mixup, num_parallel_calls=AUTO) if is_train else dset\n    return dset\n    ","metadata":{"execution":{"iopub.status.busy":"2024-04-14T01:55:21.701245Z","iopub.execute_input":"2024-04-14T01:55:21.701556Z","iopub.status.idle":"2024-04-14T01:55:21.723811Z","shell.execute_reply.started":"2024-04-14T01:55:21.701513Z","shell.execute_reply":"2024-04-14T01:55:21.723172Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Check data loading pipeline\nfold = 0\ntrain_df = data_df[data_df['fold']!=fold].sample(frac=1, random_state=0)\nval_df = data_df[data_df['fold']==fold]\n\ntrain_dset = build_dset(train_df, is_train=True, with_label=True)\nval_dset = build_dset(val_df, is_train=False, with_label=True)\n\nfor imgs, labels in train_dset:\n#for imgs, labels in val_dset:\n    break\n\nfig, axes = plt.subplots(3, 3, figsize=(8, 8), tight_layout=True)\naxes = axes.flatten()\n\nfor i, ax in enumerate(axes):\n    if i >= imgs.shape[0]:\n        break\n    \n    ax.imshow(imgs[i, :, :, :])\n    label_text = labels[i].numpy().round(2).astype(str).tolist()\n    label_text = '|'.join(label_text)\n    ax.set_title(label_text, fontsize=10)\n    ax.axis('off')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-04-14T01:55:21.724642Z","iopub.execute_input":"2024-04-14T01:55:21.724920Z","iopub.status.idle":"2024-04-14T01:55:36.879235Z","shell.execute_reply.started":"2024-04-14T01:55:21.724889Z","shell.execute_reply":"2024-04-14T01:55:36.878300Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model","metadata":{}},{"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():\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    else:\n        input_shape = (GRAPH_SHAPE[0]*2, GRAPH_SHAPE[1], 3)\n        \n    base_model = bb(\n        include_top=False,\n        weights='imagenet',\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    x = layers.Activation('softmax', name='softmax', dtype=tf.float32)(x)\n    model = Model(inputs=base_model.input, outputs=x)\n    return model","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training","metadata":{}},{"cell_type":"code","source":"# LR schedule\nclass CosineDecayWarmUp(tf.keras.optimizers.schedules.LearningRateSchedule):\n    def __init__(self, init_lr, total_steps, alpha, warmup_steps=0, warmup_lr=1e-6):\n        self.init_lr = init_lr\n        self.total_steps = total_steps\n        self.alpha = alpha\n        self.warmup_steps = warmup_steps\n        self.warmup_lr = warmup_lr\n        self.decay_steps = total_steps - warmup_steps\n\n    def __call__(self, step):\n        def cosine_decay():\n            completed_fraction = (step - self.warmup_steps) / self.decay_steps\n            cosine_decayed = 0.5 * (1.0 + tf.cos(tf.constant(math.pi, dtype=tf.float32) * completed_fraction))\n            decayed = (1 - self.alpha) * cosine_decayed + self.alpha\n            return tf.multiply(self.init_lr, decayed)\n\n        def warmup():\n            return (self.init_lr - self.warmup_lr) * step / self.warmup_steps + self.warmup_lr\n        \n        step = tf.cast(step, tf.float32)\n        step = tf.minimum(step, self.total_steps)\n        lr = tf.cond(tf.less(step, self.warmup_steps), warmup, cosine_decay)\n        return lr","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# one-fold training function\ndef run_train(fold):\n    print(f'=======================================')\n    print(f'           Fold {fold} start           ')\n    print(f'=======================================')\n    train_df = data_df[data_df['fold']!=fold].sample(frac=1, random_state=0).copy()\n    train_df_2nd = train_df[train_df['n_votes']>=N_VOTES_TH]\n    val_df = data_df[data_df['fold']==fold].copy()\n    val_df = val_df[val_df['n_votes']>=N_VOTES_TH]\n    \n    if USE_PSEUDO_LABEL:\n        # label mixing\n        tmp_df = train_df.loc[train_df['n_votes']<N_VOTES_TH]\n        org_labels = tmp_df[TARGET_COLS].values\n        pseudo_labels = tmp_df[[f'p_{x}' for x in TARGET_COLS]].values\n        weights = tmp_df['n_votes'].values[:, None]\n        train_df.loc[train_df['n_votes']<N_VOTES_TH, TARGET_COLS] = \\\n            (org_labels*weights + pseudo_labels*PSEUDO_WEIGHT) / (weights + PSEUDO_WEIGHT)\n\n    train_dset_1st = build_dset(train_df, is_train=True, with_label=True)\n    train_dset_2nd = build_dset(train_df_2nd, is_train=True, with_label=True)\n    val_dset = build_dset(val_df, is_train=False, with_label=True)\n    \n    steps_per_epoch_1st = len(train_df) // BS\n    steps_per_epoch_2nd = len(train_df_2nd) // BS\n    total_steps = EPOCHS_1ST * steps_per_epoch_1st + EPOCHS_2ND * steps_per_epoch_2nd\n\n    lr_decayed_fn = CosineDecayWarmUp(\n        INIT_LR, total_steps, ALPHA, warmup_steps=int(steps_per_epoch_1st/2), warmup_lr=1e-6)\n    \n    with strategy.scope():\n        model = build_model()\n        model.compile(\n            loss=tf.keras.losses.KLDivergence(),\n            optimizer=AdamWeightDecay(learning_rate=lr_decayed_fn, weight_decay_rate=0.01),\n            metrics=[],\n        )\n        \n    save_path = f'./hms_{INPUT_TYPE}_model_{BACKBONE_NAME}_fold{fold}_score.h5'\n    cp_cb = tf.keras.callbacks.ModelCheckpoint(\n        filepath = save_path,\n        verbose=1,\n        monitor='val_loss',\n        mode='min',\n        save_best_only=True,\n        save_weights_only=True,\n        save_freq='epoch',\n    )\n    \n    print(f'============== 1st stage ==============')\n    history = model.fit(\n        train_dset_1st,\n        epochs=EPOCHS_1ST,\n        validation_data=val_dset,\n        callbacks=[cp_cb],\n        steps_per_epoch=steps_per_epoch_1st,\n    )\n    \n    if EPOCHS_2ND != 0:\n        print(f'============== 2nd stage ==============')\n        history = model.fit(\n            train_dset_2nd,\n            epochs=EPOCHS_2ND,\n            validation_data=val_dset,\n            callbacks=[cp_cb],\n            steps_per_epoch=steps_per_epoch_2nd,\n        )\n    \n    best_loss =  min(history.history['val_loss'])\n    new_path = save_path.replace('.h5', f'{best_loss:.4f}.h5')\n    os.rename(save_path, new_path)\n    \n    # prediction of OOF (create pseudo label)\n    if not USE_PSEUDO_LABEL:\n        oof_df = data_df[data_df['fold']==fold]\n        oof_dset = build_dset(oof_df, is_train=False, with_label=False)\n        model.load_weights(new_path)\n        oof_preds = model.predict(oof_dset)\n        oof_df = oof_df[['eeg_id', 'eeg_sub_id']].copy()\n        oof_df[TARGET_COLS] = oof_preds\n    else:\n        oof_df = None\n    return best_loss, new_path, oof_df\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Run training\n\n# training parameters\nEPOCHS_1ST = 3\nEPOCHS_2ND = 5\nINIT_LR = 0.002\nALPHA = 0.01\n\ncv_scores = []\noof_dfs = []\nfor fold in range(5):\n    best_loss, new_path, oof_df = run_train(fold)\n    display(FileLink(new_path))\n    cv_scores.append(best_loss)\n    oof_dfs.append(oof_df)\n    \ntotal_cv_score = np.mean(cv_scores)\nprint(f'Total CV score: {total_cv_score:.4f}')\n\n\nif not USE_PSEUDO_LABEL:\n    oof_dfs = pd.concat(oof_dfs)\n    pseudo_csv_path = f'./pseudo_{INPUT_TYPE}_model_{BACKBONE_NAME}.csv'\n    oof_dfs.to_csv(pseudo_csv_path, index=False)\n    display(FileLink(pseudo_csv_path))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}