{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"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":"nvidiaTeslaT4","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"},{"sourceId":7392775,"sourceType":"datasetVersion","datasetId":4297782},{"sourceId":7447509,"sourceType":"datasetVersion","datasetId":4334995}],"dockerImageVersionId":30674,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Introduction\n\nThis notebook demonstrated the pre-training of a Vision Transformer model using spectrograms. Only for starter use purpose and not represent the best practice. \n\nWe used the ViTMAE model proposed in [Masked Autoencoders Are Scalable Vision Learners](http://https://arxiv.org/abs/2111.06377v2) by Kaiming He, Xinlei Chen, Saining Xie, Yanghao Li, Piotr Dollár, Ross Girshick. We will directly use the model from HuggingFace ([Link](http://https://huggingface.co/docs/transformers/en/model_doc/vit_mae)).\n\nThe ViTMAE model on HF was pretrained on ImageNet. The idea of this notebook is to have the ViTMAE pretrained on spectrograms to better learn the representation of spectrograms. Afterwards, the pretrained model can be used as the backbone to classify spectrograms. \n\nThe whole pipeline starts with defining the dataset and a collate function to reshape the input data. In our case, we reshape the input spectrograms into `3*224*224` to match with the default input size of ViTMAE. Then we can start training!\n\nWith only 5 epochs, we can see that the reconstruction of masked image was already imporved quite a lot. ","metadata":{}},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd \nfrom pathlib import Path\nimport matplotlib.pyplot as plt\nfrom typing import List, Dict\nfrom tqdm.notebook import tqdm\n\nimport torch\nfrom torch import nn\nfrom torch.nn import functional as F\nfrom transformers import ViTMAEForPreTraining\nfrom torch.utils.data import DataLoader, Dataset\nfrom torchvision.transforms import v2\nfrom torch.optim.lr_scheduler import OneCycleLR\nfrom time import time\n\n# Define Constants\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nprint(f\"Use Device: \", DEVICE)\n\nPRE_LOADED_SPECTOGRAMS = \"/kaggle/input/brain-spectrograms/specs.npy\"\nPRE_LOADED_EEGS = \"/kaggle/input/brain-eeg-spectrograms/eeg_specs.npy\"\nTRAIN_CSV = \"/kaggle/input/hms-harmful-brain-activity-classification/train.csv\"","metadata":{"execution":{"iopub.status.busy":"2024-04-08T14:43:14.895329Z","iopub.execute_input":"2024-04-08T14:43:14.895778Z","iopub.status.idle":"2024-04-08T14:43:21.958951Z","shell.execute_reply.started":"2024-04-08T14:43:14.895598Z","shell.execute_reply":"2024-04-08T14:43:21.957941Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load Data","metadata":{}},{"cell_type":"code","source":"def gen_non_overlap_samples(df_csv, targets):\n    # Reference Discussion:\n    # https://www.kaggle.com/competitions/hms-harmful-brain-activity-classification/discussion/467021\n\n    tgt_list = targets.tolist()\n    brain_activity = ['seizure', 'lpd', 'gpd', 'lrda', 'grda', 'other']\n\n    agg_dict = {\n        'spectrogram_id': 'first',\n        'spectrogram_label_offset_seconds': ['min', 'max'],\n        'patient_id': 'first',\n        'expert_consensus': 'first'\n    }\n\n    for tgt in tgt_list:\n        agg_dict[tgt] = 'sum'\n\n    groupby = df_csv.groupby(['eeg_id'])\n    train = groupby.agg(agg_dict)\n    train = train.reset_index()\n    train.columns = ['eeg_id', 'spectrogram_id', 'min', 'max', 'patient_id', 'target'] + tgt_list\n    train['total_votes'] = train[tgt_list].sum(axis=1)\n    train[tgt_list] = train[tgt_list].apply(lambda x: x / x.sum(), axis=1)\n    \n    return train","metadata":{"execution":{"iopub.status.busy":"2024-04-08T14:44:37.143334Z","iopub.execute_input":"2024-04-08T14:44:37.144341Z","iopub.status.idle":"2024-04-08T14:44:37.151998Z","shell.execute_reply.started":"2024-04-08T14:44:37.144302Z","shell.execute_reply":"2024-04-08T14:44:37.150957Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_csv = pd.read_csv(TRAIN_CSV)\ntarget_list = train_csv.columns[-6:]\nprint(target_list)\n\ntrain_df = gen_non_overlap_samples(train_csv, target_list)\n\ntrain_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-04-08T14:44:37.763231Z","iopub.execute_input":"2024-04-08T14:44:37.764046Z","iopub.status.idle":"2024-04-08T14:44:41.458394Z","shell.execute_reply.started":"2024-04-08T14:44:37.764011Z","shell.execute_reply":"2024-04-08T14:44:41.457468Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\nall_specs = np.load(PRE_LOADED_SPECTOGRAMS, allow_pickle=True).item()\nall_eegs = np.load(PRE_LOADED_EEGS, allow_pickle=True).item()","metadata":{"execution":{"iopub.status.busy":"2024-04-08T14:44:44.727242Z","iopub.execute_input":"2024-04-08T14:44:44.728057Z","iopub.status.idle":"2024-04-08T14:46:39.130528Z","shell.execute_reply.started":"2024-04-08T14:44:44.728020Z","shell.execute_reply":"2024-04-08T14:46:39.129494Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"def transform_spectrogram(spectrogram):\n\n    # Log transform spectogram\n    spectrogram = np.clip(spectrogram, np.exp(-4), np.exp(8))\n    spectrogram = np.log(spectrogram)\n\n    # Standarize per image\n    ep = 1e-6\n    mu = np.nanmean(spectrogram.flatten())\n    std = np.nanstd(spectrogram.flatten())\n    spectrogram = (spectrogram-mu) / (std+ep)\n    spectrogram = np.nan_to_num(spectrogram, nan=0.0)\n        \n    return spectrogram","metadata":{"execution":{"iopub.status.busy":"2024-04-08T14:47:41.231620Z","iopub.execute_input":"2024-04-08T14:47:41.232291Z","iopub.status.idle":"2024-04-08T14:47:41.238088Z","shell.execute_reply.started":"2024-04-08T14:47:41.232256Z","shell.execute_reply":"2024-04-08T14:47:41.237069Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class PreTrainDataset(Dataset):\n\n    def __init__(\n        self, \n        df: pd.DataFrame,\n        all_specs: Dict[str, np.ndarray],\n        all_eegs: Dict[str, np.ndarray],\n    ): \n        self.df = df\n        self.spectrograms = all_specs\n        self.eeg_spectrograms = all_eegs\n        \n    def __len__(self):\n        return len(self.df)\n        \n    def __getitem__(self, index):\n        X = self.__data_generation(index)\n        X = self.__transform(X)\n        return X\n    \n    def __data_generation(self, index): # --> [(C=8) x (H=128) x (W=256)]\n        \n        row = self.df.iloc[index]\n        r = int((row['min'] + row['max']) // 4)\n        \n        img_list = []\n        for region in range(4):\n            img = np.zeros((128, 256), dtype='float32')\n\n            spectrogram = self.spectrograms[row['spectrogram_id']][r:r+300, region*100:(region+1)*100].T\n            spectrogram = transform_spectrogram(spectrogram)\n            \n            img[14:-14, :] = spectrogram[:, 22:-22] / 2.0\n            img_list.append(img)\n\n        img = self.eeg_spectrograms[row['eeg_id']]\n        img_list += [img[:, :, i] for i in range(4)]\n      \n        X = np.array(img_list, dtype='float32')\n        X = torch.tensor(X, dtype=torch.float32)\n        \n        return X\n\n    def __transform(self, x):\n        # To be implemented...\n        return x \n\ndef reshape_input(x): #<- (N, C, H, W)\n    x = torch.stack(x, dim=0)\n    concat_p1 = torch.cat(torch.chunk(x[:, :4, :, :], 4, dim=1), dim=2)\n    concat_p2 = torch.cat(torch.chunk(x[:, 4:, :, :], 4, dim=1), dim=2)\n    x_concat = torch.cat((concat_p1, concat_p2), dim=3)\n   \n    resized = F.interpolate(x_concat, size=(224, 224), mode='bilinear', align_corners=False)\n    stacked = resized.repeat(1, 3, 1, 1)\n    \n    return stacked","metadata":{"execution":{"iopub.status.busy":"2024-04-08T14:46:52.593649Z","iopub.execute_input":"2024-04-08T14:46:52.594311Z","iopub.status.idle":"2024-04-08T14:46:52.608824Z","shell.execute_reply.started":"2024-04-08T14:46:52.594282Z","shell.execute_reply":"2024-04-08T14:46:52.607871Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Visualize Function","metadata":{}},{"cell_type":"code","source":"def visualize_model(x, model, fig_title=None, save_to=None):\n\n    if len(x.shape) == 3:\n        x = x.unsqueeze(0)\n\n    output = model(x)\n    loss = output.loss\n    y = output.logits\n    mask = output.mask\n\n    print(\"LOSS:\", loss.item())\n\n    y = model.unpatchify(y)\n    y = torch.einsum('nchw->nhwc', y).detach().cpu()\n\n    mask = mask.detach()\n    mask = mask.unsqueeze(-1).repeat(1, 1, model.config.patch_size**2 *3)  # (N, H*W, p*p*3)\n    mask = model.unpatchify(mask)  # 1 is removing, 0 is keeping\n    mask = torch.einsum('nchw->nhwc', mask).detach().cpu()\n\n    x = torch.einsum('nchw->nhwc', x)\n    x = x.detach().cpu()\n\n    # masked image\n    im_masked = x * (1 - mask)\n    # MAE reconstruction pasted with visible patches\n    im_paste = x * (1 - mask) + y * mask\n\n    fig, axes = plt.subplots(x.shape[0], 3, figsize=(15, 5*x.shape[0]))\n    for i in range(x.shape[0]):\n        if x.shape[0] == 1:\n            axes[0].imshow(x[i, :, :, 0].squeeze(), cmap='plasma', vmin=-1, vmax=2)\n            axes[1].imshow(im_paste[i, :, :, 0].squeeze(), cmap='plasma', vmin=-1, vmax=2)\n            axes[2].imshow(im_masked[i, :, :, 0].squeeze(), cmap='plasma', vmin=-1, vmax=2)\n        else:\n            axes[i, 0].imshow(x[i, :, :, 0].squeeze(), cmap='plasma', vmin=-1, vmax=2)\n            axes[i, 1].imshow(im_paste[i, :, :, 0].squeeze(), cmap='plasma', vmin=-1, vmax=2)\n            axes[i, 2].imshow(im_masked[i, :, :, 0].squeeze(), cmap='plasma', vmin=-1, vmax=2)\n\n    if x.shape[0] == 1:\n        axes[0].set_title('Original')\n        axes[1].set_title(f'{fig_title}\\nRecon.')\n        axes[2].set_title('Masked')\n    else:\n        axes[0, 0].set_title('Original')\n        axes[0, 1].set_title(f'{fig_title}\\nRecon.')\n        axes[0, 2].set_title('Masked')\n\n    for ax in axes.flatten():\n        ax.axis('off')\n    fig.tight_layout()\n\n    if save_to:\n        plt.savefig(save_to)\n\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-04-08T14:46:53.257108Z","iopub.execute_input":"2024-04-08T14:46:53.258032Z","iopub.status.idle":"2024-04-08T14:46:53.273230Z","shell.execute_reply.started":"2024-04-08T14:46:53.257998Z","shell.execute_reply":"2024-04-08T14:46:53.272131Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train Function","metadata":{}},{"cell_type":"code","source":"def train_mae(model, train_loader, config):\n    optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3, weight_decay=config['WEIGHT_DECAY'])\n    scheduler = OneCycleLR(\n        optimizer,\n        max_lr=1e-4,\n        epochs=config['EPOCHS'],\n        steps_per_epoch=len(train_loader),\n        pct_start=0.1,\n        anneal_strategy=\"cos\",\n        final_div_factor=100,\n    )\n\n    model.to(DEVICE)\n    best_loss = np.inf\n    early_stopping_counter = 0\n\n    loss_records = {}\n    for epoch in range(config['EPOCHS']):\n        start_time = time()\n        model.train()\n        losses = []\n        \n        with tqdm(train_loader, unit=\"batch\", desc='Train') as pbar:\n            for step, X in enumerate(pbar):\n                X = X.to(DEVICE)\n                optimizer.zero_grad()\n                output = model(X)\n                loss = output.loss\n                loss.backward()\n\n                grad_norm = nn.utils.clip_grad_norm_(model.parameters(), config['MAX_GRAD_NORM'])\n                optimizer.step()\n                scheduler.step()\n\n                losses.append(loss.item())\n\n                if step % config['PRINT_FREQ'] == 0 or step == len(train_loader) - 1:\n                    lr = scheduler.get_last_lr()[0]\n                    info = f\"Epoch: [{epoch + 1}][{step}/{len(train_loader)}]\"\n                    info += f\" | Loss: {loss:.4f} | Grad: {grad_norm:.4f} | LR: {lr:.4e}\"\n                    pbar.set_postfix_str(info)\n        \n        epoch_loss = np.mean(losses)\n        print(f\"{'-'*100}\\nEpoch {epoch + 1} - Avg Loss: {epoch_loss:.4f} - Time: {time() - start_time:.2f}s \\n{'-'*100}\")\n                \n        if epoch_loss < best_loss:\n            best_loss = epoch_loss\n            best_model_weights = model.state_dict()\n            torch.save(best_model_weights, f\"{config['OUTPUT_DIR']}/{config['MODEL_NAME']}_Best.pth\")\n            early_stopping_counter = 0\n        else:\n            early_stopping_counter += 1\n            if early_stopping_counter >= config['EARLY_STOPPING_PATIENCE']:\n                print(f\"Early stopping triggered. Stopping training after {epoch + 1} epochs.\")\n                break\n\n        loss_records[f\"Epoch_{epoch}\"] = losses\n                \n    return best_model_weights, loss_records\n","metadata":{"execution":{"iopub.status.busy":"2024-04-08T14:46:54.334547Z","iopub.execute_input":"2024-04-08T14:46:54.335237Z","iopub.status.idle":"2024-04-08T14:46:54.347126Z","shell.execute_reply.started":"2024-04-08T14:46:54.335205Z","shell.execute_reply":"2024-04-08T14:46:54.346235Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Pretrain ViTMAE using Spectrograms","metadata":{}},{"cell_type":"code","source":"pretrain_dataset = PreTrainDataset(train_df, all_specs, all_eegs)\npretrain_loader = DataLoader(pretrain_dataset, batch_size=16, shuffle=False, collate_fn=reshape_input)\n\n# Retrive config from HuggingFace\nmae_config = ViTMAEForPreTraining.config_class.from_pretrained('facebook/vit-mae-base')\nmae_config.attention_probs_dropout_prob = 0.00\nmae_config.hidden_dropout_prob = 0.00\n\n# Define model\nraw_mae_model = ViTMAEForPreTraining.from_pretrained('facebook/vit-mae-base', config=mae_config)","metadata":{"execution":{"iopub.status.busy":"2024-04-08T14:46:56.197820Z","iopub.execute_input":"2024-04-08T14:46:56.198490Z","iopub.status.idle":"2024-04-08T14:46:58.737463Z","shell.execute_reply.started":"2024-04-08T14:46:56.198456Z","shell.execute_reply":"2024-04-08T14:46:58.736530Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Visualize a few examples\nvisual_idx = [785, 4090, 1024, 2478, 13257]\nX = reshape_input([ pretrain_dataset[idx] for idx in visual_idx ])\nvisualize_model(X, raw_mae_model, fig_title='Non-trained model', save_to=\"/kaggle/working/raw_mae.png\")","metadata":{"execution":{"iopub.status.busy":"2024-04-08T14:48:50.099979Z","iopub.execute_input":"2024-04-08T14:48:50.100654Z","iopub.status.idle":"2024-04-08T14:48:54.999423Z","shell.execute_reply.started":"2024-04-08T14:48:50.100610Z","shell.execute_reply":"2024-04-08T14:48:54.998426Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Start Training Loop\n\nconfig = {\n    'OUTPUT_DIR': '/kaggle/working',\n    'MODEL_NAME': 'ViTMAE_PreTrained',\n    'WEIGHT_DECAY': 0.01,\n    'EPOCHS': 5,\n    'MAX_GRAD_NORM': 1e7,\n    'PRINT_FREQ': 50,\n    'EARLY_STOPPING_PATIENCE': 2\n}\n\nmodel_weights, loss_records = train_mae(raw_mae_model, pretrain_loader, config)","metadata":{"execution":{"iopub.status.busy":"2024-04-08T14:57:32.560181Z","iopub.execute_input":"2024-04-08T14:57:32.560871Z","iopub.status.idle":"2024-04-08T15:34:43.278293Z","shell.execute_reply.started":"2024-04-08T14:57:32.560837Z","shell.execute_reply":"2024-04-08T15:34:43.276726Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Visualize the result\ntrained_model = ViTMAEForPreTraining.from_pretrained(f\"{config['OUTPUT_DIR']}/{config['MODEL_NAME']}_Best.pth\", config=mae_config)\nX = reshape_input([ pretrain_dataset[idx] for idx in visual_idx ])\nvisualize_model(X, trained_model, fig_title='Trained model', save_to=\"/kaggle/working/trained_mae.png\")","metadata":{"execution":{"iopub.status.busy":"2024-04-08T15:40:40.053586Z","iopub.execute_input":"2024-04-08T15:40:40.053971Z","iopub.status.idle":"2024-04-08T15:40:45.377999Z","shell.execute_reply.started":"2024-04-08T15:40:40.053940Z","shell.execute_reply":"2024-04-08T15:40:45.377085Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Check the loss records\ndf_loss = pd.DataFrame(loss_records)\ndf_loss.to_csv(f\"{config['OUTPUT_DIR']}/pretrain_loss_records.csv\")\n\n# plot the loss with smoothing\nfig, ax = plt.subplots(figsize=(10, 5))\ndf_loss.rolling(window=50).mean().plot(\n    title='Training Loss', \n    xlabel='Steps', \n    ylabel='Loss', \n    ax=ax, \n    figsize=(10, 5)\n    )\n\nax.grid(True)\nplt.show()\n\nprint(\"Training Loss Records: \")\nprint(df_loss.mean())","metadata":{"execution":{"iopub.status.busy":"2024-04-08T15:41:04.114380Z","iopub.execute_input":"2024-04-08T15:41:04.114768Z","iopub.status.idle":"2024-04-08T15:41:04.445982Z","shell.execute_reply.started":"2024-04-08T15:41:04.114738Z","shell.execute_reply":"2024-04-08T15:41:04.445030Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}