{"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":4533,"sourceType":"modelInstanceVersion","modelInstanceId":3325}],"dockerImageVersionId":30648,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"This is based on the [EfficientNetB0 Starter](https://www.kaggle.com/code/cdeotte/efficientnetb0-starter-lb-0-43) notebook, modified for Pytorch. The original notebook is implemented in Tensorflow.\n\n* Change to pytorch's Dataset and Dataloader\n* Use efficientnet_b0 from torchvision\n* Use pytorch lightning for building the model and training\n* Inference using Trainer on multiple GPUs (DDP strategy) requires adding predictions gathering code, otherwise it will hang waiting for other nodes. So the raw pytorch's inference loop is used in the CV part. This needs to be done after training all the folds since the manual torch's GPU device initialization can't be mixed with lightning's DDP strategy.","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\n\nimport torch\nfrom torch import nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport pytorch_lightning as pl\nimport pandas as pd, numpy as np\nimport matplotlib.pyplot as plt\nfrom torchvision.models import inception_v3\nimport albumentations as albu\nfrom sklearn.model_selection import KFold, GroupKFold","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-03-03T06:43:48.008847Z","iopub.execute_input":"2024-03-03T06:43:48.009405Z","iopub.status.idle":"2024-03-03T06:43:48.016656Z","shell.execute_reply.started":"2024-03-03T06:43:48.009373Z","shell.execute_reply":"2024-03-03T06:43:48.015531Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Loading","metadata":{}},{"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\n#LOAD_MODELS_FROM = '/kaggle/input/hms-efficientnetb0-pt-ckpts/'\n\nUSE_KAGGLE_SPECTROGRAMS = True\nUSE_EEG_SPECTROGRAMS = True","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-03-03T06:43:48.018402Z","iopub.execute_input":"2024-03-03T06:43:48.018675Z","iopub.status.idle":"2024-03-03T06:43:48.026668Z","shell.execute_reply.started":"2024-03-03T06:43:48.018651Z","shell.execute_reply":"2024-03-03T06:43:48.025710Z"},"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-03-03T06:43:48.027844Z","iopub.execute_input":"2024-03-03T06:43:48.028211Z","iopub.status.idle":"2024-03-03T06:43:48.302588Z","shell.execute_reply.started":"2024-03-03T06:43:48.028176Z","shell.execute_reply":"2024-03-03T06:43:48.301603Z"},"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-03-03T06:43:48.305016Z","iopub.execute_input":"2024-03-03T06:43:48.305338Z","iopub.status.idle":"2024-03-03T06:43:48.397253Z","shell.execute_reply.started":"2024-03-03T06:43:48.305313Z","shell.execute_reply":"2024-03-03T06:43:48.396334Z"},"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-03-03T06:43:48.398338Z","iopub.execute_input":"2024-03-03T06:43:48.398600Z","iopub.status.idle":"2024-03-03T06:44:42.742275Z","shell.execute_reply.started":"2024-03-03T06:43:48.398578Z","shell.execute_reply":"2024-03-03T06:44:42.741338Z"},"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-03-03T06:44:42.743840Z","iopub.execute_input":"2024-03-03T06:44:42.744234Z","iopub.status.idle":"2024-03-03T06:45:49.392151Z","shell.execute_reply.started":"2024-03-03T06:44:42.744199Z","shell.execute_reply":"2024-03-03T06:45:49.391317Z"},"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\n","metadata":{"execution":{"iopub.status.busy":"2024-03-03T06:45:49.393451Z","iopub.execute_input":"2024-03-03T06:45:49.393745Z","iopub.status.idle":"2024-03-03T06:45:49.411233Z","shell.execute_reply.started":"2024-03-03T06:45:49.393720Z","shell.execute_reply":"2024-03-03T06:45:49.410298Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset = EEGDataset(train)\ndataloader = DataLoader(dataset, batch_size=64, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2024-03-03T06:45:49.412449Z","iopub.execute_input":"2024-03-03T06:45:49.412782Z","iopub.status.idle":"2024-03-03T06:45:49.423123Z","shell.execute_reply.started":"2024-03-03T06:45:49.412751Z","shell.execute_reply":"2024-03-03T06:45:49.422372Z"},"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\n","metadata":{"execution":{"iopub.status.busy":"2024-03-03T06:45:49.427358Z","iopub.execute_input":"2024-03-03T06:45:49.428067Z","iopub.status.idle":"2024-03-03T06:45:52.769107Z","shell.execute_reply.started":"2024-03-03T06:45:49.428020Z","shell.execute_reply":"2024-03-03T06:45:52.768194Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del dataset, dataloader\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2024-03-03T06:45:52.770394Z","iopub.execute_input":"2024-03-03T06:45:52.770776Z","iopub.status.idle":"2024-03-03T06:45:52.980350Z","shell.execute_reply.started":"2024-03-03T06:45:52.770748Z","shell.execute_reply":"2024-03-03T06:45:52.979401Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import collections\nclass Dinov2PatchEmbeddings(nn.Module):\n    def __init__(self):\n        super().__init__()\n        image_size, patch_size = 224, 14\n        image_size = image_size if isinstance(image_size, collections.abc.Iterable) else (image_size, image_size)\n        patch_size = patch_size if isinstance(patch_size, collections.abc.Iterable) else (patch_size, patch_size)\n        num_patches = (image_size[1] // patch_size[1]) * (image_size[0] // patch_size[0])\n        self.image_size = image_size\n        self.patch_size = patch_size\n        self.num_channels = 4\n        self.num_patches = num_patches\n\n        self.projection = nn.Conv2d(self.num_channels, 384, kernel_size=patch_size, stride=patch_size)\n\n    def forward(self, pixel_values: torch.Tensor) -> torch.Tensor:\n        num_channels = pixel_values.shape[1]\n        if num_channels != self.num_channels:\n            raise ValueError(\n                \"Make sure that the channel dimension of the pixel values match with the one set in the configuration.\"\n                f\" Expected {self.num_channels} but got {num_channels}.\"\n            )\n        embeddings = self.projection(pixel_values).flatten(2).transpose(1, 2)\n        return embeddings","metadata":{"execution":{"iopub.status.busy":"2024-03-03T06:45:52.981771Z","iopub.execute_input":"2024-03-03T06:45:52.982178Z","iopub.status.idle":"2024-03-03T06:45:52.992225Z","shell.execute_reply.started":"2024-03-03T06:45:52.982144Z","shell.execute_reply":"2024-03-03T06:45:52.991402Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import AutoImageProcessor, AutoModel\nimport torch.nn as nn\n\nlearnable_modules = [\n                    'encoder.layer.9',\n                    'encoder.layer.10',\n                    'encoder.layer.11']\nprocessor = AutoImageProcessor.from_pretrained('/kaggle/input/dinov2/pytorch/small/1')\ndinov2_vits14 = AutoModel.from_pretrained('/kaggle/input/dinov2/pytorch/small/1')\ndinov2_vits14.embeddings.patch_embeddings = Dinov2PatchEmbeddings()\n#for c, m in dinov2_vits14.named_modules():\n#    print(c)\ndinov2_vits14.requires_grad_(False)\nmodules = dict(dinov2_vits14.named_modules())\nfor m in learnable_modules:\n    modules[m].requires_grad_(True)\n#dinov2_vits14 = torch.hub.load('facebookresearch/dinov2', 'dinov2_vits14')","metadata":{"execution":{"iopub.status.busy":"2024-03-03T06:45:52.993479Z","iopub.execute_input":"2024-03-03T06:45:52.993994Z","iopub.status.idle":"2024-03-03T06:46:05.787873Z","shell.execute_reply.started":"2024-03-03T06:45:52.993956Z","shell.execute_reply":"2024-03-03T06:46:05.786840Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DinoV2(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.encoder = dinov2_vits14\n        self.classifier = nn.Sequential(nn.LazyLinear(6),\n                                        nn.Dropout(0.1))\n    def forward(self, x):\n        hook = self.encoder.encoder.register_forward_hook(forward_hook)\n        meta_logits = self.encoder(x)\n        hook.remove()\n        cls_logits = self.classifier(linear_input)\n        return cls_logits","metadata":{"execution":{"iopub.status.busy":"2024-03-03T06:46:05.789207Z","iopub.execute_input":"2024-03-03T06:46:05.789907Z","iopub.status.idle":"2024-03-03T06:46:05.796987Z","shell.execute_reply.started":"2024-03-03T06:46:05.789863Z","shell.execute_reply":"2024-03-03T06:46:05.795915Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def forward_hook(module, input, output):\n    sequence_output = output[0]\n    cls_token = sequence_output[:, 0]\n    patch_tokens = sequence_output[:, 1:]\n    global linear_input\n    linear_input = cls_token\n    return output\n#h = dinov2_vits14.encoder.register_forward_hook(forward_hook)","metadata":{"execution":{"iopub.status.busy":"2024-03-03T06:46:05.798250Z","iopub.execute_input":"2024-03-03T06:46:05.798520Z","iopub.status.idle":"2024-03-03T06:46:05.954659Z","shell.execute_reply.started":"2024-03-03T06:46:05.798498Z","shell.execute_reply":"2024-03-03T06:46:05.953675Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class KLDivLossWithLogits(nn.KLDivLoss):\n\n    def __init__(self):\n        super().__init__(reduction=\"batchmean\")\n\n    def forward(self, y, t):\n        y = nn.functional.softmax(y,  dim=1)\n        loss = super().forward(y, t)\n\n        return loss","metadata":{"execution":{"iopub.status.busy":"2024-03-03T06:46:05.956237Z","iopub.execute_input":"2024-03-03T06:46:05.956573Z","iopub.status.idle":"2024-03-03T06:46:05.964683Z","shell.execute_reply.started":"2024-03-03T06:46:05.956546Z","shell.execute_reply":"2024-03-03T06:46:05.963717Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tqdm.notebook import tqdm\nimport time\nimport torchvision.transforms as transforms\n\ndef train_model(model, train_loader, num_epochs, criterion, optimizer, path = None, scheduler = None):\n    print('beginning to train model')\n    if path and not os.path.exists(path):\n        os.makedirs(path)\n    model.to(device)\n    for epoch in tqdm(range(1, num_epochs + 1)):\n        model.train()\n        total_loss = 0\n        start_time = time.perf_counter()\n        for inputs, labels in tqdm(train_loader, leave=False):\n            B, S, _, __ = inputs.shape\n            inputs = inputs.reshape(B, 4, 256, 256)\n            inputs = inputs[:, :, :224, :224]\n            #inp = inp[:, :, 0:224, 0:224]\n            inputs, labels = inputs.to(device), labels.to(device)\n            optimizer.zero_grad()\n            #hook = dinov2_vits14.encoder.register_forward_hook(forward_hook)\n            output = model(inputs)\n            #print(linear_input.shape)\n            loss = criterion(output, labels)\n            loss.backward()\n            optimizer.step()\n            total_loss += loss\n        if path:\n            torch.save(model.state_dict(), f'{path}/model_ep_{epoch:2d}.pth')\n        end_time = time.perf_counter()\n        duration = end_time - start_time\n                        \n        #train_acc = compute_accuracy(model, train_loader)\n\n\n\n        current_lr = optimizer.param_groups[0]['lr']\n\n        if scheduler and current_lr > 5e-5:\n            scheduler.step()\n\n        print(f'epoch {epoch:2}',\n              f'loss: {total_loss:.3f}',\n              f'time: {duration:.3f}',\n              f'val acc: {0:.4f}',\n              sep='\\t')","metadata":{"execution":{"iopub.status.busy":"2024-03-03T06:54:40.551800Z","iopub.execute_input":"2024-03-03T06:54:40.552198Z","iopub.status.idle":"2024-03-03T06:54:40.564077Z","shell.execute_reply.started":"2024-03-03T06:54:40.552160Z","shell.execute_reply":"2024-03-03T06:54:40.563088Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"gkf = GroupKFold(n_splits=5)\ndevice = \"cuda:0\"\ndinov2 = DinoV2().cpu()\nimg = torch.randn(32, 4, 224, 224).cpu()\nprint(dinov2(img).shape)\nmodel = (dinov2)\n\noptimizer = torch.optim.Adam(model.parameters(), lr=1e-6)\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=10, eta_min=0)\ncriterion = KLDivLossWithLogits()\n\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=32, num_workers=3)\n    train_model(model, train_loader, 5, criterion, optimizer, scheduler=scheduler)\n    \n    gc.collect()","metadata":{"execution":{"iopub.status.busy":"2024-03-03T06:55:42.007509Z","iopub.execute_input":"2024-03-03T06:55:42.008367Z","iopub.status.idle":"2024-03-03T06:56:29.311802Z","shell.execute_reply.started":"2024-03-03T06:55:42.008317Z","shell.execute_reply":"2024-03-03T06:56:29.310600Z"},"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-03-03T06:46:11.851365Z","iopub.status.idle":"2024-03-03T06:46:11.851695Z","shell.execute_reply.started":"2024-03-03T06:46:11.851535Z","shell.execute_reply":"2024-03-03T06:46:11.851551Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference","metadata":{}},{"cell_type":"code","source":"\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-03-03T06:46:11.852799Z","iopub.status.idle":"2024-03-03T06:46:11.853209Z","shell.execute_reply.started":"2024-03-03T06:46:11.853018Z","shell.execute_reply":"2024-03-03T06:46:11.853050Z"},"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-03-03T06:46:11.854622Z","iopub.status.idle":"2024-03-03T06:46:11.855102Z","shell.execute_reply.started":"2024-03-03T06:46:11.854836Z","shell.execute_reply":"2024-03-03T06:46:11.854863Z"},"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-03-03T06:46:11.856658Z","iopub.status.idle":"2024-03-03T06:46:11.857139Z","shell.execute_reply.started":"2024-03-03T06:46:11.856887Z","shell.execute_reply":"2024-03-03T06:46:11.856906Z"},"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-03-03T06:46:11.858870Z","iopub.status.idle":"2024-03-03T06:46:11.859319Z","shell.execute_reply.started":"2024-03-03T06:46:11.859089Z","shell.execute_reply":"2024-03-03T06:46:11.859108Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preds = []\ntest_ds = EEGDataset(test, mode='test', specs=spectrograms2, eeg_specs=all_eegs2)\ntest_loader = DataLoader(test_ds, shuffle=False, batch_size=1, num_workers=3)\n\nfor i in range(5):\n    print('#'*25)\n    print(f'### Testing Fold {i+1}')\n\n    model.to(device).eval()\n    fold_preds = []\n\n    with torch.no_grad():\n        for test_batch in test_loader:\n            B, S, _, __ = test_batch.shape\n            test_batch = test_batch.reshape(B, 4, 256, 256)\n            test_batch = test_batch[:, :, :224, :224]\n            test_batch = test_batch.to(device)\n            pred = model(test_batch).detach().cpu()\n            pred = F.softmax(pred)\n            print(pred)\n            fold_preds.append(pred)\n        fold_preds = np.concatenate(fold_preds)\n\n    preds.append(fold_preds)\n\npred = np.mean(preds,axis=0)\nprint(\"pred\", pred)\nprint()\nprint('Test preds shape',pred.shape)","metadata":{"execution":{"iopub.status.busy":"2024-03-03T06:46:11.860545Z","iopub.status.idle":"2024-03-03T06:46:11.861028Z","shell.execute_reply.started":"2024-03-03T06:46:11.860762Z","shell.execute_reply":"2024-03-03T06:46:11.860782Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub = pd.DataFrame({'eeg_id': test.eeg_id.values})\nsub[TARGETS] = pred\nsub.to_csv('submission.csv',index=False)\nprint('Submissionn shape',sub.shape)\nsub.head()","metadata":{"execution":{"iopub.status.busy":"2024-03-03T06:46:11.863716Z","iopub.status.idle":"2024-03-03T06:46:11.864087Z","shell.execute_reply.started":"2024-03-03T06:46:11.863896Z","shell.execute_reply":"2024-03-03T06:46:11.863914Z"},"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-03-03T06:46:11.865823Z","iopub.status.idle":"2024-03-03T06:46:11.866214Z","shell.execute_reply.started":"2024-03-03T06:46:11.866005Z","shell.execute_reply":"2024-03-03T06:46:11.866049Z"},"trusted":true},"execution_count":null,"outputs":[]}]}