{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"},{"sourceId":7392733,"sourceType":"datasetVersion","datasetId":4297749},{"sourceId":7392775,"sourceType":"datasetVersion","datasetId":4297782},{"sourceId":7447509,"sourceType":"datasetVersion","datasetId":4334995},{"sourceId":7585255,"sourceType":"datasetVersion","datasetId":4415285},{"sourceId":7637636,"sourceType":"datasetVersion","datasetId":4450985},{"sourceId":7930977,"sourceType":"datasetVersion","datasetId":4492125}],"dockerImageVersionId":30648,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Harmful Brain Activity Classification\n#### Student: Breno Freitas\n#### Class: Neuro 240","metadata":{}},{"cell_type":"markdown","source":"## Importing libraries","metadata":{}},{"cell_type":"code","source":"import os\nimport sys\nimport gc\nsys.path.append('/kaggle/input/kaggle-kl-div')\nfrom kaggle_kl_div import score\nimport math\n\nimport torch\nfrom torch import nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision.models import efficientnet_b0\nimport pytorch_lightning as pl\nimport pandas as pd, numpy as np\nimport matplotlib.pyplot as plt\nimport albumentations as albu\nfrom sklearn.model_selection import KFold, GroupKFold\nimport torch.nn as nn\nimport random\nimport time\n\nfrom glob import glob\nfrom tqdm import tqdm\nfrom typing import Dict, List","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-05-01T01:23:09.595549Z","iopub.execute_input":"2024-05-01T01:23:09.595915Z","iopub.status.idle":"2024-05-01T01:23:20.926968Z","shell.execute_reply.started":"2024-05-01T01:23:09.595888Z","shell.execute_reply":"2024-05-01T01:23:20.926194Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model 1 - EfficientNet","metadata":{}},{"cell_type":"markdown","source":"### Data loading","metadata":{}},{"cell_type":"code","source":"os.environ[\"CUDA_VISIBLE_DEVICES\"] = \"0,1\"\ndevice = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\nprint('Using', torch.cuda.device_count(), 'GPU(s)')","metadata":{"execution":{"iopub.status.busy":"2024-05-01T01:23:20.928433Z","iopub.execute_input":"2024-05-01T01:23:20.928861Z","iopub.status.idle":"2024-05-01T01:23:20.968511Z","shell.execute_reply.started":"2024-05-01T01:23:20.928834Z","shell.execute_reply":"2024-05-01T01:23:20.967644Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"VER = 5\n# IF THIS EQUALS NONE, THEN WE TRAIN NEW MODELS\n# IF THIS EQUALS DISK PATH, THEN WE LOAD PREVIOUSLY TRAINED MODELS\nLOAD_MODELS_FROM = '/kaggle/input/hms-efficientnetb0-pt-ckpts/'\n# LOAD_MODELS_FROM = None\n\nUSE_KAGGLE_SPECTROGRAMS = True\nUSE_EEG_SPECTROGRAMS = True","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-05-01T01:23:20.969631Z","iopub.execute_input":"2024-05-01T01:23:20.969914Z","iopub.status.idle":"2024-05-01T01:23:20.974606Z","shell.execute_reply.started":"2024-05-01T01:23:20.969890Z","shell.execute_reply":"2024-05-01T01:23:20.973676Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv('/kaggle/input/hms-harmful-brain-activity-classification/train.csv')\nTARGETS = df.columns[-6:]\nprint('Train shape:', df.shape )\nprint('Targets', list(TARGETS))\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2024-05-01T01:23:20.976665Z","iopub.execute_input":"2024-05-01T01:23:20.976948Z","iopub.status.idle":"2024-05-01T01:23:21.237243Z","shell.execute_reply.started":"2024-05-01T01:23:20.976917Z","shell.execute_reply":"2024-05-01T01:23:21.236365Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.tail()","metadata":{"execution":{"iopub.status.busy":"2024-05-01T01:23:21.238555Z","iopub.execute_input":"2024-05-01T01:23:21.238935Z","iopub.status.idle":"2024-05-01T01:23:21.255305Z","shell.execute_reply.started":"2024-05-01T01:23:21.238901Z","shell.execute_reply":"2024-05-01T01:23:21.254211Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# train_temp = df.groupby('eeg_id')[\n#     ['spectrogram_id', 'spectrogram_label_offset_seconds']\n# ].agg({'spectrogram_id': 'first', 'spectrogram_label_offset_seconds': 'min'})\n# train_temp.columns = ['spec_id', 'min']\n\n# train_temp.head()","metadata":{"execution":{"iopub.status.busy":"2024-05-01T01:23:21.256699Z","iopub.execute_input":"2024-05-01T01:23:21.256993Z","iopub.status.idle":"2024-05-01T01:23:21.262857Z","shell.execute_reply.started":"2024-05-01T01:23:21.256968Z","shell.execute_reply":"2024-05-01T01:23:21.261899Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = df.groupby('eeg_id')[\n    ['spectrogram_id', 'spectrogram_label_offset_seconds']\n].agg({'spectrogram_id': 'first', 'spectrogram_label_offset_seconds': 'min'})\ntrain.columns = ['spec_id', 'min']\n\ntmp = df.groupby('eeg_id')[\n    ['spectrogram_id','spectrogram_label_offset_seconds']\n].agg({'spectrogram_label_offset_seconds' :'max'})\ntrain['max'] = tmp\n\ntmp = df.groupby('eeg_id')[['patient_id']].agg('first')\ntrain['patient_id'] = tmp\n\ntmp = df.groupby('eeg_id')[TARGETS].agg('sum')\nfor t in TARGETS:\n    train[t] = tmp[t].values\n    \ny_data = train[TARGETS].values\ny_data = y_data / y_data.sum(axis=1, keepdims=True)\ntrain[TARGETS] = y_data\n\ntmp = df.groupby('eeg_id')[['expert_consensus']].agg('first')\ntrain['target'] = tmp\n\ntrain = train.reset_index()\nprint('Train non-overlapp eeg_id shape:', train.shape )\ntrain.head()","metadata":{"execution":{"iopub.status.busy":"2024-05-01T01:23:21.264036Z","iopub.execute_input":"2024-05-01T01:23:21.264414Z","iopub.status.idle":"2024-05-01T01:23:21.359342Z","shell.execute_reply.started":"2024-05-01T01:23:21.264377Z","shell.execute_reply":"2024-05-01T01:23:21.358411Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"READ_SPEC_FILES = False\n\n# READ ALL SPECTROGRAMS\nPATH = '/kaggle/input/hms-harmful-brain-activity-classification/train_spectrograms/'\nfiles = os.listdir(PATH)\nprint(f'There are {len(files)} spectrogram parquets')\n\nif READ_SPEC_FILES:    \n    spectrograms = {}\n    for i,f in enumerate(files):\n        if i % 100 == 0:\n            print(i, ', ', end='')\n        tmp = pd.read_parquet(f'{PATH}{f}')\n        name = int(f.split('.')[0])\n        spectrograms[name] = tmp.iloc[:,1:].values\nelse:\n    spectrograms = np.load('/kaggle/input/brain-spectrograms/specs.npy',allow_pickle=True).item()","metadata":{"execution":{"iopub.status.busy":"2024-05-01T01:23:21.360530Z","iopub.execute_input":"2024-05-01T01:23:21.360797Z","iopub.status.idle":"2024-05-01T01:24:18.502905Z","shell.execute_reply.started":"2024-05-01T01:23:21.360774Z","shell.execute_reply":"2024-05-01T01:24:18.502061Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"READ_EEG_SPEC_FILES = False\n\nif READ_EEG_SPEC_FILES:\n    all_eegs = {}\n    for i,e in enumerate(train.eeg_id.values):\n        if i % 100 == 0:\n            print(i, ', ', end='')\n        x = np.load(f'/kaggle/input/brain-eeg-spectrograms/EEG_Spectrograms/{e}.npy')\n        all_eegs[e] = x\nelse:\n    all_eegs = np.load('/kaggle/input/brain-eeg-spectrograms/eeg_specs.npy',allow_pickle=True).item()","metadata":{"execution":{"iopub.status.busy":"2024-05-01T01:24:18.504269Z","iopub.execute_input":"2024-05-01T01:24:18.504697Z","iopub.status.idle":"2024-05-01T01:25:28.744494Z","shell.execute_reply.started":"2024-05-01T01:24:18.504664Z","shell.execute_reply":"2024-05-01T01:25:28.743673Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TARS = {'Seizure':0, 'LPD':1, 'GPD':2, 'LRDA':3, 'GRDA':4, 'Other':5}\nTARS2 = {x: y for y, x in TARS.items()}\n\n\nclass EEGDataset(Dataset):\n    \n    def __init__(self, data, augment=False, mode='train', specs=spectrograms, eeg_specs=all_eegs): \n        self.data = data\n        self.augment = augment\n        self.mode = mode\n        self.specs = specs\n        self.eeg_specs = eeg_specs\n        \n    def __len__(self):\n        return len(self.data)\n    \n    def __getitem__(self, index):\n        return self.__getitems__([index])\n    \n    def __getitems__(self, indices):\n        X, y = self._generate_data(indices)\n        if self.augment:\n            X = self._augment(X) \n        if self.mode == 'train':\n            return list(zip(X, y))\n        else:\n            return X\n    \n    def _generate_data(self, indexes):\n        X = np.zeros((len(indexes), 128, 256, 8),dtype='float32')\n        y = np.zeros((len(indexes), 6),dtype='float32')\n        img = np.ones((128, 256),dtype='float32')\n        \n        for j, i in enumerate(indexes):\n            row = self.data.iloc[i]\n            if self.mode == 'test': \n                r = 0\n            else: \n                r = int((row['min'] + row['max'])//4)\n\n            for k in range(4):\n                # EXTRACT 300 ROWS OF SPECTROGRAM\n                img = self.specs[row.spec_id][r:r+300, k*100:(k+1)*100].T\n                \n                # LOG TRANSFORM SPECTROGRAM\n                img = np.clip(img, np.exp(-4), np.exp(8))\n                img = np.log(img)\n                \n                # STANDARDIZE PER IMAGE\n                ep = 1e-6\n                m = np.nanmean(img.flatten())\n                s = np.nanstd(img.flatten())\n                img = (img - m) / (s + ep)\n                img = np.nan_to_num(img, nan=0.0)\n                \n                # CROP TO 256 TIME STEPS\n                X[j, 14:-14, :, k] = img[:, 22:-22] / 2.0\n        \n            # EEG SPECTROGRAMS\n            img = self.eeg_specs[row.eeg_id]\n            X[j, :, :, 4:] = img\n                \n            if self.mode != 'test':\n                y[j,] = row[TARGETS]\n            \n        return X, y\n    \n    def _random_transform(self, img):\n        composition = albu.Compose([\n            albu.HorizontalFlip(p=0.5),\n            # albu.CoarseDropout(max_holes=8,max_height=32,max_width=32,fill_value=0,p=0.5),\n        ])\n        return composition(image=img)['image']\n            \n    def __augment(self, img_batch):\n        for i in range(img_batch.shape[0]):\n            img_batch[i,] = self._random_transform(img_batch[i,])\n        return img_batch","metadata":{"execution":{"iopub.status.busy":"2024-05-01T01:25:28.749112Z","iopub.execute_input":"2024-05-01T01:25:28.749414Z","iopub.status.idle":"2024-05-01T01:25:28.767210Z","shell.execute_reply.started":"2024-05-01T01:25:28.749391Z","shell.execute_reply":"2024-05-01T01:25:28.766197Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset = EEGDataset(train)\ndataloader = DataLoader(dataset, batch_size=32, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2024-05-01T01:25:28.768488Z","iopub.execute_input":"2024-05-01T01:25:28.768824Z","iopub.status.idle":"2024-05-01T01:25:28.783422Z","shell.execute_reply.started":"2024-05-01T01:25:28.768785Z","shell.execute_reply":"2024-05-01T01:25:28.782542Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ROWS = 2\nCOLS = 3\nBATCHES = 2\n\nfor i, (x, y) in enumerate(dataloader):\n    plt.figure(figsize=(20, 8))\n    for j in range(ROWS):\n        for k in range(COLS):\n            plt.subplot(ROWS, COLS, j*COLS + k + 1)\n            t = y[j*COLS + k]\n            img = torch.flip(x[j*COLS+k, :, :, 0], (0,))\n            mn = img.flatten().min()\n            mx = img.flatten().max()\n            img = (img-mn)/(mx-mn)\n            plt.imshow(img)\n            tars = f'[{t[0]:0.2f}]'\n            for s in t[1:]:\n                tars += f', {s:0.2f}'\n            eeg = train.eeg_id.values[i*32+j*COLS+k]\n            plt.title(f'EEG = {eeg}\\nTarget = {tars}',size=12)\n            plt.yticks([])\n            plt.ylabel('Frequencies (Hz)',size=14)\n            plt.xlabel('Time (sec)',size=16)\n    plt.show()\n    if i == BATCHES-1:\n        break","metadata":{"execution":{"iopub.status.busy":"2024-05-01T01:25:28.784778Z","iopub.execute_input":"2024-05-01T01:25:28.785566Z","iopub.status.idle":"2024-05-01T01:25:31.741075Z","shell.execute_reply.started":"2024-05-01T01:25:28.785505Z","shell.execute_reply":"2024-05-01T01:25:31.740125Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del dataset, dataloader\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2024-05-01T01:25:31.742337Z","iopub.execute_input":"2024-05-01T01:25:31.742668Z","iopub.status.idle":"2024-05-01T01:25:31.940089Z","shell.execute_reply.started":"2024-05-01T01:25:31.742639Z","shell.execute_reply":"2024-05-01T01:25:31.939060Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Training","metadata":{}},{"cell_type":"code","source":"# WEIGHTS_FILE = '/kaggle/input/efficientnet-pytorch/efficientnet_b4_rwightman-23ab8bcd.pth'\nWEIGHTS_FILE = '/kaggle/input/efficientnet-pytorch/efficientnet_b0_rwightman-7f5810bc.pth'\n\n\nclass EEGEffnetB0(pl.LightningModule):\n    \n    def __init__(self):\n        super().__init__()\n        self.base_model = efficientnet_b0()\n        self.base_model.load_state_dict(torch.load(WEIGHTS_FILE))\n        self.base_model.classifier[1] = nn.Linear(self.base_model.classifier[1].in_features, 6, dtype=torch.float32)\n        self.prob_out = nn.Softmax()\n        \n    def forward(self, x):\n        x1 = [x[:, :, :, i:i+1] for i in range(4)]\n        x1 = torch.concat(x1, dim=1)\n        x2 = [x[:, :, :, i+4:i+5] for i in range(4)]\n        x2 = torch.concat(x2, dim=1)\n        \n        if USE_KAGGLE_SPECTROGRAMS & USE_EEG_SPECTROGRAMS:\n            x = torch.concat([x1, x2], dim=2)\n        elif USE_EEG_SPECTROGRAMS:\n            x = x2\n        else:\n            x = x1\n        x = torch.concat([x, x, x], dim=3)\n        x = x.permute(0, 3, 1, 2)\n        \n        out = self.base_model(x)\n        return out\n    \n    def training_step(self, batch, batch_idx):\n        x, y = batch\n        out = self.forward(x)\n        out = F.log_softmax(out, dim=1)\n        kl_loss = nn.KLDivLoss(reduction='batchmean')\n        loss = kl_loss(out, y)\n        return loss\n    \n    def predict_step(self, batch, batch_idx, dataloader_idx=0):\n        return F.softmax(self(batch), dim=1)\n    \n    def configure_optimizers(self):\n        optimizer = torch.optim.Adam(self.parameters(), lr=1e-3)\n        return optimizer","metadata":{"execution":{"iopub.status.busy":"2024-05-01T01:25:31.941623Z","iopub.execute_input":"2024-05-01T01:25:31.942236Z","iopub.status.idle":"2024-05-01T01:25:31.955809Z","shell.execute_reply.started":"2024-05-01T01:25:31.942199Z","shell.execute_reply":"2024-05-01T01:25:31.954798Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import torch\n# from numba import cuda\n\n# def free_gpu_cache():\n#     torch.cuda.empty_cache()\n\n#     cuda.select_device(0)\n#     cuda.close()\n#     cuda.select_device(0)\n\n# free_gpu_cache()   ","metadata":{"execution":{"iopub.status.busy":"2024-05-01T01:25:31.957071Z","iopub.execute_input":"2024-05-01T01:25:31.957378Z","iopub.status.idle":"2024-05-01T01:25:31.967511Z","shell.execute_reply.started":"2024-05-01T01:25:31.957354Z","shell.execute_reply":"2024-05-01T01:25:31.966525Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"all_oof = []\nall_true = []\nvalid_loaders = []\n\ngkf = GroupKFold(n_splits=5)\nfor i, (train_index, valid_index) in enumerate(gkf.split(train, train.target, train.patient_id)):  \n    print('#'*25)\n    print(f'### Fold {i+1}')\n    \n    train_ds = EEGDataset(train.iloc[train_index])\n    train_loader = DataLoader(train_ds, shuffle=True, batch_size=32, num_workers=3)\n    valid_ds = EEGDataset(train.iloc[valid_index], mode='valid')\n    valid_loader = DataLoader(valid_ds, shuffle=False, batch_size=64, num_workers=3)\n    \n    print(f'### Train size: {len(train_index)}, Valid size: {len(valid_index)}')\n    print('#'*25)\n    \n    trainer = pl.Trainer(max_epochs=4)\n    model = EEGEffnetB0()\n    if LOAD_MODELS_FROM is None:\n        trainer.fit(model=model, train_dataloaders=train_loader)\n        trainer.save_checkpoint(f'EffNet_v{VER}_f{i}.ckpt')\n\n    valid_loaders.append(valid_loader)\n    all_true.append(train.iloc[valid_index][TARGETS].values)\n    del trainer, model\n    gc.collect()","metadata":{"execution":{"iopub.status.busy":"2024-05-01T01:25:31.968913Z","iopub.execute_input":"2024-05-01T01:25:31.969274Z","iopub.status.idle":"2024-05-01T01:25:34.930174Z","shell.execute_reply.started":"2024-05-01T01:25:31.969242Z","shell.execute_reply":"2024-05-01T01:25:34.929400Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{"execution":{"iopub.status.busy":"2024-05-01T01:25:34.931322Z","iopub.execute_input":"2024-05-01T01:25:34.931601Z","iopub.status.idle":"2024-05-01T01:25:34.935599Z","shell.execute_reply.started":"2024-05-01T01:25:34.931578Z","shell.execute_reply":"2024-05-01T01:25:34.934742Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for i in range(5):\n    print('#'*25)\n    print(f'### Validating Fold {i+1}')\n\n    ckpt_file = f'EffNet_v{VER}_f{i}.ckpt' if LOAD_MODELS_FROM is None else f'{LOAD_MODELS_FROM}/EffNet_v{VER}_f{i}.ckpt'\n    model = EEGEffnetB0.load_from_checkpoint(ckpt_file)\n    model.to(device).eval()\n    with torch.inference_mode():\n        for val_batch in valid_loaders[i]:\n            val_batch = val_batch.to(device)\n            oof = torch.softmax(model(val_batch), dim=1).cpu().numpy()\n            all_oof.append(oof)\n    del model\n    gc.collect()\n\nall_oof = np.concatenate(all_oof)\nall_true = np.concatenate(all_true)","metadata":{"execution":{"iopub.status.busy":"2024-05-01T01:25:34.937038Z","iopub.execute_input":"2024-05-01T01:25:34.937306Z","iopub.status.idle":"2024-05-01T01:27:30.507030Z","shell.execute_reply.started":"2024-05-01T01:25:34.937284Z","shell.execute_reply":"2024-05-01T01:27:30.506126Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"oof = pd.DataFrame(all_oof.copy())\noof['id'] = np.arange(len(oof))\n\ntrue = pd.DataFrame(all_true.copy())\ntrue['id'] = np.arange(len(true))\n\ncv = score(solution=true, submission=oof, row_id_column_name='id')\nprint('CV Score KL-Div for EfficientNetB0 =',cv)","metadata":{"execution":{"iopub.status.busy":"2024-05-01T01:27:30.509055Z","iopub.execute_input":"2024-05-01T01:27:30.509370Z","iopub.status.idle":"2024-05-01T01:27:30.571697Z","shell.execute_reply.started":"2024-05-01T01:27:30.509343Z","shell.execute_reply":"2024-05-01T01:27:30.570798Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Training analysis ","metadata":{}},{"cell_type":"code","source":"def calc_conditional_prob(\n    preds: np.ndarray,\n    gts: np.ndarray,\n    weights: np.ndarray | None = None,\n    normalize: bool = True,\n    norm_axis: int = 0,\n    eps=1e-4,\n):\n    \"\"\"\n    Parameters\n    ----------\n    preds: (N, C) array of predicted probabilities\n    gts: (N, C) array of ground truth probabilities\n    weights: (N, ) array of weights\n\n    Returns\n    -------\n    conditional_matrix: (C, C) array of conditional probabilities\n    \"\"\"\n    gts = gts[:, :, np.newaxis]\n    preds = preds[:, np.newaxis, :]\n    if weights is not None:\n        weights = weights[:, np.newaxis, np.newaxis]\n        mat = (preds * gts * weights).sum(axis=0) / weights.sum(axis=0)\n    else:\n        mat = (preds * gts).mean(axis=0)\n\n    if normalize:\n        mat = mat / (mat.sum(axis=norm_axis, keepdims=True) + eps)\n    return mat","metadata":{"execution":{"iopub.status.busy":"2024-05-01T01:27:30.572856Z","iopub.execute_input":"2024-05-01T01:27:30.573209Z","iopub.status.idle":"2024-05-01T01:27:30.581230Z","shell.execute_reply.started":"2024-05-01T01:27:30.573178Z","shell.execute_reply":"2024-05-01T01:27:30.580198Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"conditional_matrix = calc_conditional_prob(all_oof, all_true)\n\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\nclass_names = ['seizure', 'lpd', 'gpd', 'lrda', 'grda', 'other']\n\nplt.figure(figsize=(10, 8))\nsns.heatmap(conditional_matrix, annot=True, cmap='Blues', fmt=\".2f\", xticklabels=class_names, yticklabels=class_names, annot_kws={\"size\": 14})\nplt.title('EfficientNet: Conditional Probability Matrix', fontsize=16)\nplt.xlabel('Predicted Classes', fontsize=14)\nplt.ylabel('Actual Classes', fontsize=14)\n\n# Increase tick label size\nplt.xticks(fontsize=14)\nplt.yticks(fontsize=14)\n\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-05-01T01:27:30.582343Z","iopub.execute_input":"2024-05-01T01:27:30.582631Z","iopub.status.idle":"2024-05-01T01:27:31.042607Z","shell.execute_reply.started":"2024-05-01T01:27:30.582608Z","shell.execute_reply":"2024-05-01T01:27:31.041663Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"temp_true = pd.DataFrame(all_true)\ntemp_true.columns = class_names\nmean_probabilities = temp_true.mean()\n\n# Display the mean probabilities\nprint(mean_probabilities)\n\n# Plotting a stacked bar chart\nax = mean_probabilities.plot(kind='bar', color='skyblue')\nplt.title('Average Probability Distribution Across Classes', fontsize=16)\nplt.ylabel('Average Probability', fontsize=14)\nplt.xlabel('Classes', fontsize=14)\n\nax.set_ylim(0, max(mean_probabilities) + 0.05)  # Increase the limit to max value plus 0.5\n\nfor p in ax.patches:\n    ax.annotate(format(p.get_height(), '.2f'),  # Format the label\n                (p.get_x() + p.get_width() / 2., p.get_height()),  # Position\n                ha = 'center', va = 'center',  # Center alignment\n                xytext = (0, 9),  # 9 points vertical offset\n                textcoords = 'offset points',\n                fontsize=14)\n\n# Increase tick label size\nplt.xticks(fontsize=14)\nplt.yticks(fontsize=14)\n\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-05-01T01:30:37.726237Z","iopub.execute_input":"2024-05-01T01:30:37.727124Z","iopub.status.idle":"2024-05-01T01:30:37.957245Z","shell.execute_reply.started":"2024-05-01T01:30:37.727091Z","shell.execute_reply":"2024-05-01T01:30:37.956333Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Inference","metadata":{}},{"cell_type":"code","source":"del all_eegs, spectrograms\ngc.collect()\n\ntest = pd.read_csv('/kaggle/input/hms-harmful-brain-activity-classification/test.csv')\nprint('Test shape',test.shape)\ntest.head()","metadata":{"execution":{"iopub.status.busy":"2024-05-01T01:27:31.603216Z","iopub.status.idle":"2024-05-01T01:27:31.603569Z","shell.execute_reply.started":"2024-05-01T01:27:31.603404Z","shell.execute_reply":"2024-05-01T01:27:31.603418Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# READ ALL SPECTROGRAMS\nPATH2 = '/kaggle/input/hms-harmful-brain-activity-classification/test_spectrograms/'\nfiles2 = os.listdir(PATH2)\nprint(f'There are {len(files2)} test spectrogram parquets')\n    \nspectrograms2 = {}\nfor i, f in enumerate(files2):\n    if i % 100 == 0:\n        print(i, ', ',end='')\n    tmp = pd.read_parquet(f'{PATH2}{f}')\n    name = int(f.split('.')[0])\n    spectrograms2[name] = tmp.iloc[:, 1:].values\n    \n# RENAME FOR DATALOADER\ntest = test.rename({'spectrogram_id': 'spec_id'}, axis=1)","metadata":{"execution":{"iopub.status.busy":"2024-05-01T01:27:31.604655Z","iopub.status.idle":"2024-05-01T01:27:31.604976Z","shell.execute_reply.started":"2024-05-01T01:27:31.604822Z","shell.execute_reply":"2024-05-01T01:27:31.604834Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pywt, librosa\n\nUSE_WAVELET = None \n\nNAMES = ['LL','LP','RP','RR']\n\nFEATS = [['Fp1','F7','T3','T5','O1'],\n         ['Fp1','F3','C3','P3','O1'],\n         ['Fp2','F8','T4','T6','O2'],\n         ['Fp2','F4','C4','P4','O2']]\n\n\n# DENOISE FUNCTION\ndef maddest(d, axis=None):\n    return np.mean(np.absolute(d - np.mean(d, axis)), axis)\n\n\ndef denoise(x, wavelet='haar', level=1):    \n    coeff = pywt.wavedec(x, wavelet, mode=\"per\")\n    sigma = (1/0.6745) * maddest(coeff[-level])\n\n    uthresh = sigma * np.sqrt(2*np.log(len(x)))\n    coeff[1:] = (pywt.threshold(i, value=uthresh, mode='hard') for i in coeff[1:])\n\n    ret=pywt.waverec(coeff, wavelet, mode='per')\n    \n    return ret\n\n\ndef spectrogram_from_eeg(parquet_path, display=False):\n    \n    # LOAD MIDDLE 50 SECONDS OF EEG SERIES\n    eeg = pd.read_parquet(parquet_path)\n    middle = (len(eeg)-10_000)//2\n    eeg = eeg.iloc[middle:middle+10_000]\n    \n    # VARIABLE TO HOLD SPECTROGRAM\n    img = np.zeros((128,256,4),dtype='float32')\n    \n    if display: plt.figure(figsize=(10,7))\n    signals = []\n    for k in range(4):\n        COLS = FEATS[k]\n        \n        for kk in range(4):\n        \n            # COMPUTE PAIR DIFFERENCES\n            x = eeg[COLS[kk]].values - eeg[COLS[kk+1]].values\n\n            # FILL NANS\n            m = np.nanmean(x)\n            if np.isnan(x).mean()<1: x = np.nan_to_num(x,nan=m)\n            else: x[:] = 0\n\n            # DENOISE\n            if USE_WAVELET:\n                x = denoise(x, wavelet=USE_WAVELET)\n            signals.append(x)\n\n            # RAW SPECTROGRAM\n            mel_spec = librosa.feature.melspectrogram(y=x, sr=200, hop_length=len(x)//256, \n                  n_fft=1024, n_mels=128, fmin=0, fmax=20, win_length=128)\n\n            # LOG TRANSFORM\n            width = (mel_spec.shape[1]//32)*32\n            mel_spec_db = librosa.power_to_db(mel_spec, ref=np.max).astype(np.float32)[:,:width]\n\n            # STANDARDIZE TO -1 TO 1\n            mel_spec_db = (mel_spec_db+40)/40 \n            img[:,:,k] += mel_spec_db\n                \n        # AVERAGE THE 4 MONTAGE DIFFERENCES\n        img[:,:,k] /= 4.0\n        \n        if display:\n            plt.subplot(2,2,k+1)\n            plt.imshow(img[:,:,k],aspect='auto',origin='lower')\n            plt.title(f'EEG {eeg_id} - Spectrogram {NAMES[k]}')\n            \n    if display: \n        plt.show()\n        plt.figure(figsize=(10,5))\n        offset = 0\n        for k in range(4):\n            if k>0: offset -= signals[3-k].min()\n            plt.plot(range(10_000),signals[k]+offset,label=NAMES[3-k])\n            offset += signals[3-k].max()\n        plt.legend()\n        plt.title(f'EEG {eeg_id} Signals')\n        plt.show()\n        print(); print('#'*25); print()\n        \n    return img","metadata":{"execution":{"iopub.status.busy":"2024-05-01T01:27:31.606682Z","iopub.status.idle":"2024-05-01T01:27:31.607162Z","shell.execute_reply.started":"2024-05-01T01:27:31.606908Z","shell.execute_reply":"2024-05-01T01:27:31.606928Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# READ ALL EEG SPECTROGRAMS\nPATH2 = '/kaggle/input/hms-harmful-brain-activity-classification/test_eegs/'\nDISPLAY = 1\nEEG_IDS2 = test.eeg_id.unique()\nall_eegs2 = {}\n\nprint('Converting Test EEG to Spectrograms...'); print()\nfor i, eeg_id in enumerate(EEG_IDS2):\n        \n    # CREATE SPECTROGRAM FROM EEG PARQUET\n    img = spectrogram_from_eeg(f'{PATH2}{eeg_id}.parquet', i < DISPLAY)\n    all_eegs2[eeg_id] = img","metadata":{"execution":{"iopub.status.busy":"2024-05-01T01:27:31.609016Z","iopub.status.idle":"2024-05-01T01:27:31.609381Z","shell.execute_reply.started":"2024-05-01T01:27:31.609207Z","shell.execute_reply":"2024-05-01T01:27:31.609222Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# INFER EFFICIENTNET ON TEST\npreds = []\ntest_ds = EEGDataset(test, mode='test', specs=spectrograms2, eeg_specs=all_eegs2)\ntest_loader = DataLoader(test_ds, shuffle=False, batch_size=64, num_workers=3)\n\nfor i in range(5):\n    print('#'*25)\n    print(f'### Testing Fold {i+1}')\n\n    ckpt_file = f'EffNet_v{VER}_f{i}.ckpt' if LOAD_MODELS_FROM is None else f'{LOAD_MODELS_FROM}/EffNet_v{VER}_f{i}.ckpt'\n    model = EEGEffnetB0.load_from_checkpoint(ckpt_file)\n    model.to(device).eval()\n    fold_preds = []\n\n    with torch.inference_mode():\n        for test_batch in test_loader:\n            test_batch = test_batch.to(device)\n            pred = torch.softmax(model(test_batch), dim=1).cpu().numpy()\n            fold_preds.append(pred)\n        fold_preds = np.concatenate(fold_preds)\n\n    preds.append(fold_preds)\n\nefficientnet_predictions = np.mean(preds,axis=0)\nprint()\nprint('Test preds shape',efficientnet_predictions.shape)","metadata":{"execution":{"iopub.status.busy":"2024-05-01T01:27:31.611039Z","iopub.status.idle":"2024-05-01T01:27:31.611522Z","shell.execute_reply.started":"2024-05-01T01:27:31.611306Z","shell.execute_reply":"2024-05-01T01:27:31.611326Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model 2 - WaveNet","metadata":{}},{"cell_type":"markdown","source":"### Config and utils","metadata":{}},{"cell_type":"code","source":"class config:\n    BATCH_SIZE_TEST = 32\n    NUM_WORKERS = 0 # multiprocessing.cpu_count()\n    PRINT_FREQ = 20\n    SEED = 20\n    VISUALIZE = False\n    \n    \nclass paths:\n    OUTPUT_DIR = \"/kaggle/working/\"\n    TEST_CSV = \"/kaggle/input/hms-harmful-brain-activity-classification/test.csv\"\n    TEST_EEGS = \"/kaggle/input/hms-harmful-brain-activity-classification/test_eegs/\"\n    \n    \nmodel_weights = [x for x in glob(\"/kaggle/input/hms-wavenet/*.pth\")]\nmodel_weights","metadata":{"execution":{"iopub.status.busy":"2024-05-01T01:27:31.612812Z","iopub.status.idle":"2024-05-01T01:27:31.613134Z","shell.execute_reply.started":"2024-05-01T01:27:31.612976Z","shell.execute_reply":"2024-05-01T01:27:31.612990Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def eeg_from_parquet(parquet_path: str) -> np.ndarray:\n    \"\"\"\n    This function reads a parquet file and extracts the middle 50 seconds of readings. Then it fills NaN values\n    with the mean value (ignoring NaNs).\n    :param parquet_path: path to parquet file.\n    :param display: whether to display EEG plots or not.\n    :return data: np.array of shape  (time_steps, eeg_features) -> (10_000, 8)\n    \"\"\"\n    # === Extract middle 50 seconds ===\n    eeg = pd.read_parquet(parquet_path, columns=eeg_features)\n    rows = len(eeg)\n    offset = (rows - 10_000) // 2 # 50 * 200 = 10_000\n    eeg = eeg.iloc[offset:offset+10_000] # middle 50 seconds, has the same amount of readings to left and right\n    # === Convert to numpy ===\n    data = np.zeros((10_000, len(eeg_features))) # create placeholder of same shape with zeros\n    for index, feature in enumerate(eeg_features):\n        x = eeg[feature].values.astype('float32') # convert to float32\n        mean = np.nanmean(x) # arithmetic mean along the specified axis, ignoring NaNs\n        nan_percentage = np.isnan(x).mean() # percentage of NaN values in feature\n        # === Fill nan values ===\n        if nan_percentage < 1: # if some values are nan, but not all\n            x = np.nan_to_num(x, nan=mean)\n        else: # if all values are nan\n            x[:] = 0\n        data[:, index] = x\n   \n    return data\n\n\ndef seed_everything(seed: int):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed) \n    \n    \ndef sep():\n    print(\"-\"*100)\n\n    \ntarget_preds = [x + \"_pred\" for x in ['seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']]\nlabel_to_num = {'Seizure': 0, 'LPD': 1, 'GPD': 2, 'LRDA': 3, 'GRDA': 4, 'Other':5}\nnum_to_label = {v: k for k, v in label_to_num.items()}\nseed_everything(config.SEED)","metadata":{"execution":{"iopub.status.busy":"2024-05-01T01:27:31.614026Z","iopub.status.idle":"2024-05-01T01:27:31.614359Z","shell.execute_reply.started":"2024-05-01T01:27:31.614198Z","shell.execute_reply":"2024-05-01T01:27:31.614212Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Data loading","metadata":{}},{"cell_type":"code","source":"# Prepare raw EEG\n\ntest_df = pd.read_csv(paths.TEST_CSV)\neeg_parquet_paths = glob(paths.TEST_EEGS + \"*.parquet\")\neeg_df = pd.read_parquet(eeg_parquet_paths[0])\neeg_features = eeg_df.columns\nprint(f'There are {len(eeg_features)} raw eeg features')\nprint(list(eeg_features))\neeg_features = ['Fp1','T3','C3','O1','Fp2','C4','T4','O2']\nfeature_to_index = {x:y for x,y in zip(eeg_features, range(len(eeg_features)))}\n\nCREATE_EEGS = False\nall_eegs = {}\nvisualize = 1\neeg_paths = glob(paths.TEST_EEGS + \"*.parquet\")\neeg_ids = test_df.eeg_id.unique()\n\nfor i, eeg_id in tqdm(enumerate(eeg_ids)):  \n    # Save EEG to Python dictionary of numpy arrays\n    eeg_path = paths.TEST_EEGS + str(eeg_id) + \".parquet\"\n    data = eeg_from_parquet(eeg_path)              \n    all_eegs[eeg_id] = data","metadata":{"execution":{"iopub.status.busy":"2024-05-01T01:27:31.616548Z","iopub.status.idle":"2024-05-01T01:27:31.616996Z","shell.execute_reply.started":"2024-05-01T01:27:31.616755Z","shell.execute_reply":"2024-05-01T01:27:31.616781Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Butter low-pass filter","metadata":{}},{"cell_type":"code","source":"from scipy.signal import butter, lfilter\n\ndef butter_lowpass_filter(data, cutoff_freq: int = 20, sampling_rate: int = 200, order: int = 4):\n    nyquist = 0.5 * sampling_rate\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","metadata":{"execution":{"iopub.status.busy":"2024-05-01T01:27:31.618317Z","iopub.status.idle":"2024-05-01T01:27:31.618752Z","shell.execute_reply.started":"2024-05-01T01:27:31.618523Z","shell.execute_reply":"2024-05-01T01:27:31.618542Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Dataset","metadata":{}},{"cell_type":"code","source":"class CustomDataset(Dataset):\n    def __init__(\n        self, df: pd.DataFrame, config,\n        eegs: Dict[int, np.ndarray] = all_eegs, downsample: int = 5\n    ): \n        self.df = df\n        self.config = config\n        self.batch_size = self.config.BATCH_SIZE_TEST\n        self.eegs = eegs\n        self.downsample = downsample\n        \n    def __len__(self):\n        \"\"\"\n        Length of dataset.\n        \"\"\"\n        return len(self.df)\n        \n    def __getitem__(self, index):\n        \"\"\"\n        Get one item.\n        \"\"\"\n        X = self.__data_generation(index)\n        X = X[::self.downsample, :]\n        output = {\n            \"X\": torch.tensor(X, dtype=torch.float32)\n        }\n        return output\n                        \n    def __data_generation(self, index):\n        row = self.df.iloc[index]\n        X = np.zeros((10_000, 8), dtype='float32')\n        data = self.eegs[row.eeg_id]\n\n        # === Feature engineering ===\n        X[:,0] = data[:,feature_to_index['Fp1']] - data[:,feature_to_index['T3']]\n        X[:,1] = data[:,feature_to_index['T3']] - data[:,feature_to_index['O1']]\n\n        X[:,2] = data[:,feature_to_index['Fp1']] - data[:,feature_to_index['C3']]\n        X[:,3] = data[:,feature_to_index['C3']] - data[:,feature_to_index['O1']]\n\n        X[:,4] = data[:,feature_to_index['Fp2']] - data[:,feature_to_index['C4']]\n        X[:,5] = data[:,feature_to_index['C4']] - data[:,feature_to_index['O2']]\n\n        X[:,6] = data[:,feature_to_index['Fp2']] - data[:,feature_to_index['T4']]\n        X[:,7] = data[:,feature_to_index['T4']] - data[:,feature_to_index['O2']]\n\n        # === Standarize ===\n        X = np.clip(X,-1024, 1024)\n        X = np.nan_to_num(X, nan=0) / 32.0\n\n        # === Butter Low-pass Filter ===\n        X = butter_lowpass_filter(X)\n            \n        return X","metadata":{"execution":{"iopub.status.busy":"2024-05-01T01:27:31.620380Z","iopub.status.idle":"2024-05-01T01:27:31.620824Z","shell.execute_reply.started":"2024-05-01T01:27:31.620586Z","shell.execute_reply":"2024-05-01T01:27:31.620605Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### DataLoader","metadata":{}},{"cell_type":"code","source":"test_dataset = CustomDataset(test_df, config)\ntest_loader = DataLoader(\n    test_dataset,\n    batch_size=config.BATCH_SIZE_TEST,\n    shuffle=False,\n    num_workers=config.NUM_WORKERS,\n    pin_memory=True,\n    drop_last=False\n)\noutput = test_dataset[0]\nX = output[\"X\"]\nprint(f\"X shape: {X.shape}\")","metadata":{"execution":{"iopub.status.busy":"2024-05-01T01:27:31.622301Z","iopub.status.idle":"2024-05-01T01:27:31.622745Z","shell.execute_reply.started":"2024-05-01T01:27:31.622508Z","shell.execute_reply":"2024-05-01T01:27:31.622526Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Model","metadata":{}},{"cell_type":"code","source":"class Wave_Block(nn.Module):\n    def __init__(self, in_channels: int, out_channels: int, dilation_rates: int, kernel_size: int = 3):\n        \"\"\"\n        WaveNet building block.\n        :param in_channels: number of input channels.\n        :param out_channels: number of output channels.\n        :param dilation_rates: how many levels of dilations are used.\n        :param kernel_size: size of the convolving kernel.\n        \"\"\"\n        super(Wave_Block, self).__init__()\n        self.num_rates = dilation_rates\n        self.convs = nn.ModuleList()\n        self.filter_convs = nn.ModuleList()\n        self.gate_convs = nn.ModuleList()\n        self.convs.append(nn.Conv1d(in_channels, out_channels, kernel_size=1, bias=True))\n        \n        dilation_rates = [2 ** i for i in range(dilation_rates)]\n        for dilation_rate in dilation_rates:\n            self.filter_convs.append(\n                nn.Conv1d(out_channels, out_channels, kernel_size=kernel_size,\n                          padding=int((dilation_rate*(kernel_size-1))/2), dilation=dilation_rate))\n            self.gate_convs.append(\n                nn.Conv1d(out_channels, out_channels, kernel_size=kernel_size,\n                          padding=int((dilation_rate*(kernel_size-1))/2), dilation=dilation_rate))\n            self.convs.append(nn.Conv1d(out_channels, out_channels, kernel_size=1, bias=True))\n        \n        for i in range(len(self.convs)):\n            nn.init.xavier_uniform_(self.convs[i].weight, gain=nn.init.calculate_gain('relu'))\n            nn.init.zeros_(self.convs[i].bias)\n\n        for i in range(len(self.filter_convs)):\n            nn.init.xavier_uniform_(self.filter_convs[i].weight, gain=nn.init.calculate_gain('relu'))\n            nn.init.zeros_(self.filter_convs[i].bias)\n\n        for i in range(len(self.gate_convs)):\n            nn.init.xavier_uniform_(self.gate_convs[i].weight, gain=nn.init.calculate_gain('relu'))\n            nn.init.zeros_(self.gate_convs[i].bias)\n\n    def forward(self, x):\n        x = self.convs[0](x)\n        res = x\n        for i in range(self.num_rates):\n            tanh_out = torch.tanh(self.filter_convs[i](x))\n            sigmoid_out = torch.sigmoid(self.gate_convs[i](x))\n            x = tanh_out * sigmoid_out\n            x = self.convs[i + 1](x) \n            res = res + x\n        return res\n    \nclass WaveNet(nn.Module):\n    def __init__(self, input_channels: int = 1, kernel_size: int = 3):\n        super(WaveNet, self).__init__()\n        self.model = nn.Sequential(\n                Wave_Block(input_channels, 8, 12, kernel_size),\n                Wave_Block(8, 16, 8, kernel_size),\n                Wave_Block(16, 32, 4, kernel_size),\n                Wave_Block(32, 64, 1, kernel_size) \n        )\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        x = x.permute(0, 2, 1) \n        output = self.model(x)\n        return output\n\n\nclass CustomModel(nn.Module):\n    def __init__(self):\n        super(CustomModel, self).__init__()\n        self.model = WaveNet()\n        self.global_avg_pooling = nn.AdaptiveAvgPool1d(1)\n        self.dropout = 0.0\n        self.head = nn.Sequential(\n            nn.Linear(256, 64),\n            nn.BatchNorm1d(64),\n            nn.ReLU(),\n            nn.Dropout(self.dropout),\n            nn.Linear(64, 6)\n        )\n        \n    def forward(self, x: torch.Tensor):\n        \"\"\"\n        Forwward pass.\n        \"\"\"\n        x1 = self.model(x[:, :, 0:1])\n        x1 = self.global_avg_pooling(x1)\n        x1 = x1.squeeze(dim=2)\n        x2 = self.model(x[:, :, 1:2])\n        x2 = self.global_avg_pooling(x2)\n        x2 = x2.squeeze(dim=2)\n        z1 = torch.mean(torch.stack([x1, x2]), dim=0)\n\n        x1 = self.model(x[:, :, 2:3])\n        x1 = self.global_avg_pooling(x1)\n        x1 = x1.squeeze(dim=2)\n        x2 = self.model(x[:, :, 3:4])\n        x2 = self.global_avg_pooling(x2)\n        x2 = x2.squeeze(dim=2)\n        z2 = torch.mean(torch.stack([x1, x2]), dim=0)\n        \n        x1 = self.model(x[:, :, 4:5])\n        x1 = self.global_avg_pooling(x1)\n        x1 = x1.squeeze(dim=2)\n        x2 = self.model(x[:, :, 5:6])\n        x2 = self.global_avg_pooling(x2)\n        x2 = x2.squeeze(dim=2)\n        z3 = torch.mean(torch.stack([x1, x2]), dim=0)\n        \n        x1 = self.model(x[:, :, 6:7])\n        x1 = self.global_avg_pooling(x1)\n        x1 = x1.squeeze(dim=2)\n        x2 = self.model(x[:, :, 7:8])\n        x2 = self.global_avg_pooling(x2)\n        x2 = x2.squeeze(dim=2)\n        z4 = torch.mean(torch.stack([x1, x2]), dim=0)\n        \n        y = torch.cat([z1, z2, z3, z4], dim=1)\n        y = self.head(y)\n        \n        return y\n\nmodel = CustomModel()\ntotal_params = sum(p.numel() for p in model.parameters())\nprint(f\"Total number of parameters: {total_params}\")","metadata":{"execution":{"iopub.status.busy":"2024-05-01T01:27:31.624509Z","iopub.status.idle":"2024-05-01T01:27:31.624954Z","shell.execute_reply.started":"2024-05-01T01:27:31.624718Z","shell.execute_reply":"2024-05-01T01:27:31.624735Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Inference","metadata":{}},{"cell_type":"code","source":"def inference_function(test_loader, model, device):\n    model.eval() # set model in evaluation mode\n    softmax = nn.Softmax(dim=1)\n    prediction_dict = {}\n    preds = []\n    with tqdm(test_loader, unit=\"test_batch\", desc='Inference') as tqdm_test_loader:\n        for step, batch in enumerate(tqdm_test_loader):\n            X = batch.pop(\"X\").to(device) # send inputs to `device`\n            batch_size = X.size(0)\n            with torch.no_grad():\n                y_preds = model(X) # forward propagation pass\n            y_preds = softmax(y_preds)\n            preds.append(y_preds.to('cpu').numpy()) # save predictions\n                \n    prediction_dict[\"predictions\"] = np.concatenate(preds) # np.array() of shape (fold_size, target_cols)\n    return prediction_dict","metadata":{"execution":{"iopub.status.busy":"2024-05-01T01:27:31.626567Z","iopub.status.idle":"2024-05-01T01:27:31.627011Z","shell.execute_reply.started":"2024-05-01T01:27:31.626784Z","shell.execute_reply":"2024-05-01T01:27:31.626803Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predictions = []\n\nfor model_weight in model_weights:\n    test_dataset = CustomDataset(test_df, config)\n    train_loader = DataLoader(\n        test_dataset,\n        batch_size=config.BATCH_SIZE_TEST,\n        shuffle=False,\n        num_workers=config.NUM_WORKERS,\n        pin_memory=True,\n        drop_last=False\n    )\n    model = CustomModel()\n    checkpoint = torch.load(model_weight)\n    model.load_state_dict(checkpoint[\"model\"])\n    model.to(device)\n    prediction_dict = inference_function(test_loader, model, device)\n    predictions.append(prediction_dict[\"predictions\"])\n    torch.cuda.empty_cache()\n    gc.collect()\n    \nwavenet_predictions = np.array(predictions)\nwavenet_predictions = np.mean(wavenet_predictions, axis=0)","metadata":{"execution":{"iopub.status.busy":"2024-05-01T01:27:31.629161Z","iopub.status.idle":"2024-05-01T01:27:31.629507Z","shell.execute_reply.started":"2024-05-01T01:27:31.629346Z","shell.execute_reply":"2024-05-01T01:27:31.629360Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Ensemble - Submission","metadata":{}},{"cell_type":"code","source":"# Calculate weighted average of predictions and submit\nsub = pd.DataFrame({'eeg_id': test.eeg_id.values})\nsub[TARGETS] = 0.75*efficientnet_predictions + 0.25*wavenet_predictions\nsub.to_csv('submission.csv',index=False)\nprint('Submissionn shape',sub.shape)\nsub.head()","metadata":{"execution":{"iopub.status.busy":"2024-05-01T01:27:31.631364Z","iopub.status.idle":"2024-05-01T01:27:31.631668Z","shell.execute_reply.started":"2024-05-01T01:27:31.631518Z","shell.execute_reply":"2024-05-01T01:27:31.631530Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# SANITY CHECK TO CONFIRM PREDICTIONS SUM TO ONE\nsub.iloc[:,-6:].sum(axis=1)","metadata":{"execution":{"iopub.status.busy":"2024-05-01T01:27:31.632690Z","iopub.status.idle":"2024-05-01T01:27:31.632996Z","shell.execute_reply.started":"2024-05-01T01:27:31.632845Z","shell.execute_reply":"2024-05-01T01:27:31.632858Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}