{"metadata":{"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"},{"sourceId":7392775,"sourceType":"datasetVersion","datasetId":4297782},{"sourceId":7465251,"sourceType":"datasetVersion","datasetId":4317718},{"sourceId":7570342,"sourceType":"datasetVersion","datasetId":4407194},{"sourceId":7752462,"sourceType":"datasetVersion","datasetId":4382744},{"sourceId":7818976,"sourceType":"datasetVersion","datasetId":4417235}],"dockerImageVersionId":30635,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true},"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.10.12"},"papermill":{"default_parameters":{},"duration":98.666323,"end_time":"2024-02-28T19:49:51.36372","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2024-02-28T19:48:12.697397","version":"2.4.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Features+Head Ensemble Starter [LB 0.34] for HMS Brain Comp\nThis is Features+Head is a combination and ensemble Starter notebook for Kaggle's HMS brain comp. We can train 4 different models using:\n- Kaggle's spectrograms (CV 0.6123 – LB 0.41)\n- Chris's EEG spectrograms(modified version) (CV 0.6288 – LB 0.39)\n- Both Kaggle and EEG spectrograms (CV 0.5768 – LB 0.37)\n- Chris's [WaveNet][4] (CV 0.6992 - LB 0.41)\n\n**The Ensemble achieves LB 0.34** \n\nGreat discussion [here][5] by @KOLOO that led to the latest score!\n\nFeatures+Head Starter uses Chris Deotte's Kaggle dataset [here][1]. Also Uses Chris's EEG spectrograms [here][3] (modified version) \n\n### Train and Infer Tips\n\nThis notebook can be used both to train and submit (infer) to Kaggle LB. When training, you can set variable `submission = False` , you can also set `TEST_MODE = TRUE` to upload 500 samples queckly instead of the whole dataset for testing. \n\nTo train a specific model type, you should set `DATA_TYPE = 'both|eeg|kaggle|raw'`, `kaggle` to train on Kaggle's spectrograms, `eeg` to train on EEG's spectrograms, `both` to train on Kaggle's and EEG's spectrograms, `raw` to train on EEG's signal with WaveNet,\n\nFor submission after training models, you should save them in the LOAD_MODELS_FROM dataset, then run this notebook with `submission = True`.\n\nOnce we have all the models saved to LOAD_MODELS_FROM and ready ensemble, we should set `submission = True` and `ENSEMBLE = True` and set the models versions that we prior specified, as well as their `LBs` for weighted ensemble.\n\nThis notebook is made as generic as possible to expand and try different experiments.\n\nWhat you could do:\n- Change EfficientNetB(0-7) with `LOAD_BACKBONE_FROM`\n- Data augmentation by setting DataGenerator's parameter to `augment = True`\n- Different image configurations as input.\n- WaveNet model tuning.\n\n\nThis notebook is a direct descendent of Chris's notebook [here][2]\n\n[1]: https://www.kaggle.com/datasets/cdeotte/brain-spectrograms\n[2]: https://www.kaggle.com/code/cdeotte/efficientnetb2-starter-lb-0-57\n[3]: https://www.kaggle.com/datasets/nartaa/eeg-spectrograms\n[4]: https://www.kaggle.com/competitions/hms-harmful-brain-activity-classification/discussion/468684\n[5]: https://www.kaggle.com/competitions/hms-harmful-brain-activity-classification/discussion/477461","metadata":{"papermill":{"duration":0.00905,"end_time":"2024-02-28T19:48:16.229318","exception":false,"start_time":"2024-02-28T19:48:16.220268","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import os, random\nimport tensorflow as tf\nimport tensorflow\nimport tensorflow.keras.backend as K\nimport pandas as pd, numpy as np\nimport matplotlib.pyplot as plt\nfrom tensorflow.keras.models import load_model\n\n\nVER = 45\nDATA_TYPE = 'both' # both|eeg|kaggle|raw\nTEST_MODE = False\nsubmission = False\n\n\nnp.random.seed(21)\nrandom.seed(21)\ntf.random.set_seed(21)\n\n# USE SINGLE GPU, MULTIPLE GPUS \ngpus = tf.config.list_physical_devices('GPU')\n# WE USE MIXED PRECISION\ntf.config.optimizer.set_experimental_options({\"auto_mixed_precision\": True})\nif len(gpus)>1:\n    strategy = tf.distribute.MirroredStrategy()\n    print(f'Using {len(gpus)} GPUs')\nelse:\n    strategy = tf.distribute.OneDeviceStrategy(device=\"/gpu:0\")\n    print(f'Using {len(gpus)} GPU')","metadata":{"papermill":{"duration":13.504974,"end_time":"2024-02-28T19:48:29.742388","exception":false,"start_time":"2024-02-28T19:48:16.237414","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-13T18:40:57.821740Z","iopub.execute_input":"2024-03-13T18:40:57.822080Z","iopub.status.idle":"2024-03-13T18:41:18.184612Z","shell.execute_reply.started":"2024-03-13T18:40:57.822051Z","shell.execute_reply":"2024-03-13T18:41:18.183592Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load and create Non-Overlapping Eeg Id Train Data\nThe competition data description says that test data does not have multiple crops from the same `eeg_id`. Therefore we will train and validate using only 1 crop per `eeg_id`. There is a discussion about this [here][1].\n[1]: https://www.kaggle.com/competitions/hms-harmful-brain-activity-classification/discussion/467021","metadata":{"papermill":{"duration":0.011883,"end_time":"2024-02-28T19:48:29.765736","exception":false,"start_time":"2024-02-28T19:48:29.753853","status":"completed"},"tags":[]}},{"cell_type":"code","source":"TARGETS = ['seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']\nFEATS2 = ['Fp1','T3','C3','O1','Fp2','C4','T4','O2']\nFEAT2IDX = {x:y for x,y in zip(FEATS2,range(len(FEATS2)))}\n\ndef eeg_from_parquet(parquet_path):\n\n    eeg = pd.read_parquet(parquet_path, columns=FEATS2)\n    rows = len(eeg)\n    offset = (rows-10_000)//2\n    eeg = eeg.iloc[offset:offset+10_000]\n    data = np.zeros((10_000,len(FEATS2)))\n    for j,col in enumerate(FEATS2):\n        \n        # FILL NAN\n        x = eeg[col].values.astype('float32')\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        data[:,j] = x\n\n    return data\n\ndef add_kl(data):\n    import torch\n    labels = data[TARGETS].values + 1e-5\n\n    # compute kl-loss with uniform distribution by pytorch\n    data['kl'] = torch.nn.functional.kl_div(\n        torch.log(torch.tensor(labels)),\n        torch.tensor([1 / 6] * 6),\n        reduction='none'\n    ).sum(dim=1).numpy()\n    return data\n\ndef reset_seed(seed):\n    np.random.seed(seed)\n    random.seed(seed)\n    tf.random.set_seed(seed)\n    \nif not submission:\n    train = pd.read_csv('/kaggle/input/hms-harmful-brain-activity-classification/train.csv')\n    TARGETS = ['seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']\n    META = ['spectrogram_id','spectrogram_label_offset_seconds','patient_id','expert_consensus']\n    train = train.groupby('eeg_id')[META+TARGETS\n                           ].agg({**{m:'first' for m in META},**{t:'sum' for t in TARGETS}}).reset_index() \n    train[TARGETS] = train[TARGETS]/train[TARGETS].values.sum(axis=1,keepdims=True)\n    train.columns = ['eeg_id','spec_id','offset','patient_id','target'] + TARGETS\n    train = add_kl(train)\n    print(train.head(1).to_string())","metadata":{"papermill":{"duration":0.025421,"end_time":"2024-02-28T19:48:29.800372","exception":false,"start_time":"2024-02-28T19:48:29.774951","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-13T18:41:18.186506Z","iopub.execute_input":"2024-03-13T18:41:18.187055Z","iopub.status.idle":"2024-03-13T18:41:23.422558Z","shell.execute_reply.started":"2024-03-13T18:41:18.187026Z","shell.execute_reply":"2024-03-13T18:41:23.421343Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Read Train Spectrograms and EEGs\n\nWe can read 3 file from Chris's [Kaggle dataset here][1] which contains all the 11k spectrograms. From Chris's modified EEG spectrogram [here][2]. From Chris's EEG signals [here][3]\n\n[1]: https://www.kaggle.com/datasets/cdeotte/brain-spectrograms\n[2]: https://www.kaggle.com/datasets/nartaa/eeg-spectrograms\n[3]: https://www.kaggle.com/datasets/cdeotte/brain-eegs","metadata":{"papermill":{"duration":0.007453,"end_time":"2024-02-28T19:48:29.815771","exception":false,"start_time":"2024-02-28T19:48:29.808318","status":"completed"},"tags":[]}},{"cell_type":"code","source":"%%time\nif not submission:\n    # FOR TESTING SET TEST_MODE TO TRUE\n    if TEST_MODE:\n        train = train.sample(500,random_state=42).reset_index(drop=True)\n        spectrograms = {}\n        for i,e in enumerate(train.spec_id.values):\n            if i%100==0: print(i,', ',end='')\n            x = pd.read_parquet(f'/kaggle/input/hms-harmful-brain-activity-classification/train_spectrograms/{e}.parquet')\n            spectrograms[e] = x.values\n        all_eegs = {}\n        for i,e in enumerate(train.eeg_id.values):\n            if i%100==0: print(i,', ',end='')\n            x = np.load(f'/kaggle/input/eeg-spectrograms/EEG_Spectrograms/{e}.npy')\n            all_eegs[e] = x\n        all_raw_eegs = {}\n        for i,e in enumerate(train.eeg_id.values):\n            if i%100==0: print(i,', ',end='')\n            x = eeg_from_parquet(f'/kaggle/input/hms-harmful-brain-activity-classification/train_eegs/{e}.parquet')              \n            all_raw_eegs[e] = x\n    else:\n        spectrograms = None\n        all_eegs = None\n        all_raw_eegs = None\n        if DATA_TYPE=='both' or DATA_TYPE=='kaggle':\n            spectrograms = np.load('/kaggle/input/brain-spectrograms/specs.npy',allow_pickle=True).item()\n        if DATA_TYPE=='both' or DATA_TYPE=='eeg':\n            all_eegs = np.load('/kaggle/input/eeg-spectrograms/eeg_specs.npy',allow_pickle=True).item()\n        if DATA_TYPE=='raw':\n            all_raw_eegs = np.load('/kaggle/input/brain-eegs/eegs.npy',allow_pickle=True).item()","metadata":{"papermill":{"duration":0.02106,"end_time":"2024-02-28T19:48:29.844607","exception":false,"start_time":"2024-02-28T19:48:29.823547","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-13T18:41:23.423772Z","iopub.execute_input":"2024-03-13T18:41:23.424381Z","iopub.status.idle":"2024-03-13T18:44:06.078451Z","shell.execute_reply.started":"2024-03-13T18:41:23.424350Z","shell.execute_reply":"2024-03-13T18:44:06.077550Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# DATA GENERATOR\nThis data generator outputs 512x512x3, the spectrogram and eeg images are concatenated all togother in a single image. For using data augmention you can set `augment = True` when creating the train data generator.","metadata":{"papermill":{"duration":0.007648,"end_time":"2024-02-28T19:48:29.860169","exception":false,"start_time":"2024-02-28T19:48:29.852521","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import albumentations as albu\nfrom scipy.signal import butter, lfilter\n\nclass DataGenerator():\n    'Generates data for Keras'\n    def __init__(self, data, specs=None, eeg_specs=None, raw_eegs=None, augment=False, mode='train', data_type=DATA_TYPE): \n        self.data = data\n        self.augment = augment\n        self.mode = mode\n        self.data_type = data_type\n        self.specs = specs\n        self.eeg_specs = eeg_specs\n        self.raw_eegs = raw_eegs\n        self.on_epoch_end()\n        \n    def __len__(self):\n        return self.data.shape[0]\n\n    def __getitem__(self, index):\n        X, y = self.data_generation(index)\n        if self.augment: X = self.augmentation(X)\n        return X, y\n    \n    def __call__(self):\n        for i in range(self.__len__()):\n            yield self.__getitem__(i)\n            \n            if i == self.__len__()-1:\n                self.on_epoch_end()\n                \n    def on_epoch_end(self):\n        if self.mode=='train': \n            self.data = self.data.sample(frac=1).reset_index(drop=True)\n    \n    def data_generation(self, index):\n        if self.data_type == 'both':\n            X,y = self.generate_all_specs(index)\n        elif self.data_type == 'eeg' or self.data_type == 'kaggle':\n            X,y = self.generate_specs(index)\n        elif self.data_type == 'raw':\n            X,y = self.generate_raw(index)\n\n        return X,y\n    \n    def generate_all_specs(self, index):\n        X = np.zeros((512,512,3),dtype='float32')\n        y = np.zeros((6,),dtype='float32')\n        \n        row = self.data.iloc[index]\n        if self.mode=='test': \n            offset = 0\n        else:\n            offset = int(row.offset/2)\n            \n        eeg = self.eeg_specs[row.eeg_id]\n        spec = self.specs[row.spec_id]\n        \n        imgs = [spec[offset:offset+300,k*100:(k+1)*100].T for k in [0,2,1,3]] # to match kaggle with eeg\n        img = np.stack(imgs,axis=-1)\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        img = np.nan_to_num(img, nan=0.0)    \n            \n        mn = img.flatten().min()\n        mx = img.flatten().max()\n        ep = 1e-5\n        img = 255 * (img - mn) / (mx - mn + ep)\n        \n        X[0_0+56:100+56,:256,0] = img[:,22:-22,0] # LL_k\n        X[100+56:200+56,:256,0] = img[:,22:-22,2] # RL_k\n        X[0_0+56:100+56,:256,1] = img[:,22:-22,1] # LP_k\n        X[100+56:200+56,:256,1] = img[:,22:-22,3] # RP_k\n        X[0_0+56:100+56,:256,2] = img[:,22:-22,2] # RL_k\n        X[100+56:200+56,:256,2] = img[:,22:-22,1] # LP_k\n        \n        X[0_0+56:100+56,256:,0] = img[:,22:-22,0] # LL_k\n        X[100+56:200+56,256:,0] = img[:,22:-22,2] # RL_k\n        X[0_0+56:100+56,256:,1] = img[:,22:-22,1] # LP_k\n        X[100+56:200+56,256:,1] = img[:,22:-22,3] # RP_K\n        \n        # EEG\n        img = eeg\n        mn = img.flatten().min()\n        mx = img.flatten().max()\n        ep = 1e-5\n        img = 255 * (img - mn) / (mx - mn + ep)\n        X[200+56:300+56,:256,0] = img[:,22:-22,0] # LL_e\n        X[300+56:400+56,:256,0] = img[:,22:-22,2] # RL_e\n        X[200+56:300+56,:256,1] = img[:,22:-22,1] # LP_e\n        X[300+56:400+56,:256,1] = img[:,22:-22,3] # RP_e\n        X[200+56:300+56,:256,2] = img[:,22:-22,2] # RL_e\n        X[300+56:400+56,:256,2] = img[:,22:-22,1] # LP_e\n        \n        X[200+56:300+56,256:,0] = img[:,22:-22,0] # LL_e\n        X[300+56:400+56,256:,0] = img[:,22:-22,2] # RL_e\n        X[200+56:300+56,256:,1] = img[:,22:-22,1] # LP_e\n        X[300+56:400+56,256:,1] = img[:,22:-22,3] # RP_e\n\n        if self.mode!='test':\n            y[:] = row[TARGETS]\n        \n        return X,y\n    \n    def generate_specs(self, index):\n        X = np.zeros((512,512,3),dtype='float32')\n        y = np.zeros((6,),dtype='float32')\n        \n        row = self.data.iloc[index]\n        if self.mode=='test': \n            offset = 0\n        else:\n            offset = int(row.offset/2)\n            \n        if self.data_type == 'eeg':\n            img = self.eeg_specs[row.eeg_id]\n        elif self.data_type == 'kaggle':\n            spec = self.specs[row.spec_id]\n            imgs = [spec[offset:offset+300,k*100:(k+1)*100].T for k in [0,2,1,3]] # to match kaggle with eeg\n            img = np.stack(imgs,axis=-1)\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            img = np.nan_to_num(img, nan=0.0)    \n            \n        mn = img.flatten().min()\n        mx = img.flatten().max()\n        ep = 1e-5\n        img = 255 * (img - mn) / (mx - mn + ep)\n        \n        X[0_0+56:100+56,:256,0] = img[:,22:-22,0]\n        X[100+56:200+56,:256,0] = img[:,22:-22,2]\n        X[0_0+56:100+56,:256,1] = img[:,22:-22,1]\n        X[100+56:200+56,:256,1] = img[:,22:-22,3]\n        X[0_0+56:100+56,:256,2] = img[:,22:-22,2]\n        X[100+56:200+56,:256,2] = img[:,22:-22,1]\n        \n        X[0_0+56:100+56,256:,0] = img[:,22:-22,0]\n        X[100+56:200+56,256:,0] = img[:,22:-22,1]\n        X[0_0+56:100+56,256:,1] = img[:,22:-22,2]\n        X[100+56:200+56,256:,1] = img[:,22:-22,3]\n        \n        X[200+56:300+56,:256,0] = img[:,22:-22,0]\n        X[300+56:400+56,:256,0] = img[:,22:-22,1]\n        X[200+56:300+56,:256,1] = img[:,22:-22,2]\n        X[300+56:400+56,:256,1] = img[:,22:-22,3]\n        X[200+56:300+56,:256,2] = img[:,22:-22,3]\n        X[300+56:400+56,:256,2] = img[:,22:-22,2]\n        \n        X[200+56:300+56,256:,0] = img[:,22:-22,0]\n        X[300+56:400+56,256:,0] = img[:,22:-22,2]\n        X[200+56:300+56,256:,1] = img[:,22:-22,1]\n        X[300+56:400+56,256:,1] = img[:,22:-22,3]\n        \n        if self.mode!='test':\n            y[:] = row[TARGETS]\n        \n        return X,y\n    \n    def generate_raw(self,index):\n        X = np.zeros((10_000,8),dtype='float32')\n        y = np.zeros((6,),dtype='float32')\n        \n        row = self.data.iloc[index]\n        eeg = self.raw_eegs[row.eeg_id]\n            \n        # FEATURE ENGINEER\n        X[:,0] = eeg[:,FEAT2IDX['Fp1']] - eeg[:,FEAT2IDX['T3']]\n        X[:,1] = eeg[:,FEAT2IDX['T3']] - eeg[:,FEAT2IDX['O1']]\n            \n        X[:,2] = eeg[:,FEAT2IDX['Fp1']] - eeg[:,FEAT2IDX['C3']]\n        X[:,3] = eeg[:,FEAT2IDX['C3']] - eeg[:,FEAT2IDX['O1']]\n            \n        X[:,4] = eeg[:,FEAT2IDX['Fp2']] - eeg[:,FEAT2IDX['C4']]\n        X[:,5] = eeg[:,FEAT2IDX['C4']] - eeg[:,FEAT2IDX['O2']]\n            \n        X[:,6] = eeg[:,FEAT2IDX['Fp2']] - eeg[:,FEAT2IDX['T4']]\n        X[:,7] = eeg[:,FEAT2IDX['T4']] - eeg[:,FEAT2IDX['O2']]\n            \n        # STANDARDIZE\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 = self.butter_lowpass_filter(X)\n        # Downsample\n        X = X[::5,:]\n        \n        if self.mode!='test':\n            y[:] = row[TARGETS]\n                \n        return X,y\n        \n    def butter_lowpass_filter(self, data, cutoff_freq=20, sampling_rate=200, order=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\n    \n    def resize(self, img,size):\n        composition = albu.Compose([\n                albu.Resize(size[0],size[1])\n            ])\n        return composition(image=img)['image']\n            \n    def augmentation(self, img):\n        composition = albu.Compose([\n                albu.HorizontalFlip(p=0.4)\n            ])\n        return composition(image=img)['image']","metadata":{"papermill":{"duration":1.920135,"end_time":"2024-02-28T19:48:31.788415","exception":false,"start_time":"2024-02-28T19:48:29.86828","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-13T18:44:06.081660Z","iopub.execute_input":"2024-03-13T18:44:06.081981Z","iopub.status.idle":"2024-03-13T18:44:09.288432Z","shell.execute_reply.started":"2024-03-13T18:44:06.081942Z","shell.execute_reply":"2024-03-13T18:44:09.287656Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class VAEDataGenerator(DataGenerator):\n    def __getitem__(self,index):\n        x,y=super().__getitem__(index)\n        return (x)/256,(x)/256","metadata":{"execution":{"iopub.status.busy":"2024-03-13T18:44:09.289485Z","iopub.execute_input":"2024-03-13T18:44:09.290032Z","iopub.status.idle":"2024-03-13T18:44:09.294979Z","shell.execute_reply.started":"2024-03-13T18:44:09.290004Z","shell.execute_reply":"2024-03-13T18:44:09.294086Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# DISPLAY DATA GENERATOR\nBelow we display example data generator spectrogram images and raw EEG signals.","metadata":{"papermill":{"duration":0.008418,"end_time":"2024-02-28T19:48:31.805954","exception":false,"start_time":"2024-02-28T19:48:31.797536","status":"completed"},"tags":[]}},{"cell_type":"code","source":"if not submission and DATA_TYPE!='raw':\n    gen = DataGenerator(train, augment=False, specs=spectrograms, eeg_specs=all_eegs, data_type=DATA_TYPE)\n    for x,y in gen:\n        break\n    plt.imshow(x[:,:,0])\n    plt.title(f'Target = {y.round(1)}',size=12)\n    plt.yticks([])\n    plt.ylabel('Frequencies (Hz)',size=12)\n    plt.xlabel('Time (sec)',size=12)\n    plt.show()\n    plt.imshow(x[:,:,1])\n    plt.title(f'Target = {y.round(1)}',size=12)\n    plt.yticks([])\n    plt.ylabel('Frequencies (Hz)',size=12)\n    plt.xlabel('Time (sec)',size=12)\n    plt.show()\n    plt.imshow(x[:,:,2])\n    plt.title(f'Target = {y.round(1)}',size=12)\n    plt.yticks([])\n    plt.ylabel('Frequencies (Hz)',size=12)\n    plt.xlabel('Time (sec)',size=12)\n    plt.show()\n    \nif not submission and DATA_TYPE=='raw':\n    gen = DataGenerator(train, raw_eegs=all_raw_eegs, data_type=DATA_TYPE)\n    for x,y in gen:\n        plt.figure(figsize=(20,4))\n        offset = 0\n        for j in range(x.shape[-1]):\n            if j!=0: offset -= x[:,j].min()\n            plt.plot(range(2_000),x[:,j]+offset,label=f'feature {j+1}')\n            offset += x[:,j].max()\n        plt.legend()\n        plt.show()\n        break","metadata":{"papermill":{"duration":0.019565,"end_time":"2024-02-28T19:48:31.833755","exception":false,"start_time":"2024-02-28T19:48:31.81419","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-13T18:44:09.296122Z","iopub.execute_input":"2024-03-13T18:44:09.296410Z","iopub.status.idle":"2024-03-13T18:44:10.129469Z","shell.execute_reply.started":"2024-03-13T18:44:09.296385Z","shell.execute_reply":"2024-03-13T18:44:10.128566Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Autoencoder","metadata":{}},{"cell_type":"code","source":"\"\"\"def halved_glorot_uniform(shape, dtype=None):\n    initializer = tf.keras.initializers.GlorotUniform()\n    weights = initializer(shape, dtype)\n    return weights / 2.0\ntf.keras.layers.Dense.Conv2D = halved_glorot_uniform\"\"\"","metadata":{"execution":{"iopub.status.busy":"2024-03-13T18:44:10.130710Z","iopub.execute_input":"2024-03-13T18:44:10.131047Z","iopub.status.idle":"2024-03-13T18:44:10.137655Z","shell.execute_reply.started":"2024-03-13T18:44:10.131015Z","shell.execute_reply":"2024-03-13T18:44:10.136695Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import tensorflow as tf\n\nclass ResNetBlock(tf.keras.layers.Layer):\n    def __init__(self, in_channels, kernel_size, modify=False, bn=True):\n        super(ResNetBlock, self).__init__()\n        self.modify = modify\n        if modify == 'downsample':\n            self.conv1 = tf.keras.layers.Conv2D(in_channels*2, kernel_size, strides=2, padding='same', use_bias=False, activation=tf.nn.relu, kernel_regularizer=tf.keras.regularizers.l2(0.005))\n            self.conv2 = tf.keras.layers.Conv2D(in_channels*2, kernel_size, padding='same', use_bias=False, activation=tf.nn.relu, kernel_regularizer=tf.keras.regularizers.l2(0.005))\n            if bn:\n                self.bn1 = tf.keras.layers.BatchNormalization()\n                self.bn2 = tf.keras.layers.BatchNormalization()\n            else:\n                self.bn1 = tf.keras.layers.Layer()\n                self.bn2 = tf.keras.layers.Layer()\n        elif modify == 'upsample':\n            self.conv1 = tf.keras.layers.Conv2DTranspose(in_channels//2, kernel_size, strides=2, padding='same', output_padding=1, use_bias=False, activation=tf.nn.relu, kernel_regularizer=tf.keras.regularizers.l2(0.005))\n            self.conv2 = tf.keras.layers.Conv2D(in_channels//2, kernel_size, padding='same', use_bias=False, activation=tf.nn.relu, kernel_regularizer=tf.keras.regularizers.l2(0.005))\n            self.bn1 = tf.keras.layers.BatchNormalization()\n            self.bn2 = tf.keras.layers.BatchNormalization()\n        else:\n            self.conv1 = tf.keras.layers.Conv2D(in_channels, kernel_size, padding='same', activation=tf.nn.relu, kernel_regularizer=tf.keras.regularizers.l2(0.005))\n            self.conv2 = tf.keras.layers.Conv2D(in_channels, kernel_size, padding='same', activation=tf.nn.relu, kernel_regularizer=tf.keras.regularizers.l2(0.005))\n            self.bn1 = tf.keras.layers.BatchNormalization()\n            self.bn2 = tf.keras.layers.BatchNormalization()\n        self.act = tf.keras.layers.ReLU()\n        if modify == 'downsample':\n            self.proj = tf.keras.layers.Conv2D(in_channels*2, kernel_size, strides=2, padding='same', activation=tf.nn.relu, kernel_regularizer=tf.keras.regularizers.l2(0.005))\n        if modify == 'upsample':\n            self.proj = tf.keras.layers.Conv2DTranspose(in_channels//2, kernel_size, strides=2, padding='same', output_padding=1, activation=tf.nn.relu, kernel_regularizer=tf.keras.regularizers.l2(0.005))\n\n    def call(self, x):\n        out = self.conv1(x)\n        out = self.bn1(out)\n        out = self.act(out)\n        out = self.conv2(out)\n        out = self.bn2(out)\n        if self.modify:\n            x = self.proj(x)\n        out = x + out\n        out = self.act(out)\n        return out\n","metadata":{"execution":{"iopub.status.busy":"2024-03-13T18:44:10.139436Z","iopub.execute_input":"2024-03-13T18:44:10.139866Z","iopub.status.idle":"2024-03-13T18:44:10.158071Z","shell.execute_reply.started":"2024-03-13T18:44:10.139826Z","shell.execute_reply":"2024-03-13T18:44:10.157308Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\ntrain, test = train_test_split(train, train_size=0.80)","metadata":{"execution":{"iopub.status.busy":"2024-03-13T18:44:10.159713Z","iopub.execute_input":"2024-03-13T18:44:10.160067Z","iopub.status.idle":"2024-03-13T18:44:10.185478Z","shell.execute_reply.started":"2024-03-13T18:44:10.160036Z","shell.execute_reply":"2024-03-13T18:44:10.184538Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"x_train = VAEDataGenerator(train, augment=False, specs=spectrograms, eeg_specs=all_eegs, raw_eegs=all_raw_eegs)\nx_train = tf.data.Dataset.from_generator(generator=x_train, \n                                               output_signature=(tf.TensorSpec(shape=(512,512,3), dtype=tf.float32),\n                                                                 tf.TensorSpec(shape=(512,512,3), dtype=tf.float32))).batch(16).prefetch(tf.data.AUTOTUNE)\nx_test = VAEDataGenerator(test, augment=False, specs=spectrograms, eeg_specs=all_eegs, raw_eegs=all_raw_eegs)\nx_test= tf.data.Dataset.from_generator(generator=x_test, \n                                               output_signature=(tf.TensorSpec(shape=(512,512,3), dtype=tf.float32),\n                                                                 tf.TensorSpec(shape=(512,512,3), dtype=tf.float32))).batch(16).prefetch(tf.data.AUTOTUNE)","metadata":{"execution":{"iopub.status.busy":"2024-03-13T18:44:10.187249Z","iopub.execute_input":"2024-03-13T18:44:10.188011Z","iopub.status.idle":"2024-03-13T18:44:10.309161Z","shell.execute_reply.started":"2024-03-13T18:44:10.187977Z","shell.execute_reply":"2024-03-13T18:44:10.308418Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for x,y in x_train:\n    print(x.shape)\n    print(y.shape)\n    print(np.max(x))\n    print(np.min(x))\n    break","metadata":{"execution":{"iopub.status.busy":"2024-03-13T18:44:10.312210Z","iopub.execute_input":"2024-03-13T18:44:10.312499Z","iopub.status.idle":"2024-03-13T18:44:10.680311Z","shell.execute_reply.started":"2024-03-13T18:44:10.312475Z","shell.execute_reply":"2024-03-13T18:44:10.679363Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"latent_space_dim = 512*3","metadata":{"execution":{"iopub.status.busy":"2024-03-13T18:44:10.681648Z","iopub.execute_input":"2024-03-13T18:44:10.681972Z","iopub.status.idle":"2024-03-13T18:44:10.686048Z","shell.execute_reply.started":"2024-03-13T18:44:10.681946Z","shell.execute_reply":"2024-03-13T18:44:10.685124Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"inputs = tf.keras.layers.Input(shape=(512,512,3))\nconv = tf.keras.layers.Conv2D(3, 7, 1, padding='same')(inputs)\neres1 = ResNetBlock(16, 3, modify='downsample')(conv)\neres2 = ResNetBlock(32, 3, modify='downsample')(eres1)\neres3 = ResNetBlock(64, 3, modify='downsample')(eres2)\neres4 = ResNetBlock(128, 3, modify='downsample')(eres3)\neres5 = ResNetBlock(256, 3, modify='downsample')(eres4)\neres6 = ResNetBlock(512, 3, modify='downsample')(eres5)\n\nshape_before_flatten = tensorflow.keras.backend.int_shape(eres6)[1:]\nencoder_flatten = tensorflow.keras.layers.Flatten()(eres6)\n\nz_mean_l  = tensorflow.keras.layers.Dense(units=latent_space_dim, name=\"encoder_mu\")\nz_mean = z_mean_l(encoder_flatten) \nz_log_var_l  = tensorflow.keras.layers.Dense(units=latent_space_dim, name=\"encoder_log_variance\", kernel_initializer='zeros', kernel_regularizer=tf.keras.regularizers.l2(3))\nz_log_var = z_log_var_l(encoder_flatten)\n\n#encoder_mu_log_variance_model = tensorflow.keras.models.Model(enc_input_layer, (encoder_mu, encoder_log_variance), name=\"encoder_mu_log_variance_model\")\n\n@tf.function\ndef sampling(args):\n    z_mean_, z_log_var_ = args\n    #tf.print(z_mean_)\n    #tf.print(z_log_var_)\n    #tf.print(tf.reduce_max(z_log_var_))\n    epsilon = K.random_normal(shape=(K.shape(z_mean_)[0], latent_space_dim))\n    #tf.print(epsilon)\n    result = z_mean_ + K.exp(z_log_var_ / 2) * epsilon\n    #tf.print(result)\n    return tf.debugging.check_numerics(result, \"NaN detected in sampling\")\n\n# Reparameterization trick\nz = tf.keras.layers.Lambda(sampling)([z_mean, z_log_var])\n\nencoder = tf.keras.Model(inputs, [z_mean, z_log_var, z], name='encoder')","metadata":{"execution":{"iopub.status.busy":"2024-03-13T18:44:10.687113Z","iopub.execute_input":"2024-03-13T18:44:10.687380Z","iopub.status.idle":"2024-03-13T18:44:11.437043Z","shell.execute_reply.started":"2024-03-13T18:44:10.687350Z","shell.execute_reply":"2024-03-13T18:44:11.436107Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dec_input_layer = tf.keras.layers.Input(shape=(latent_space_dim))\ndecoder_dense_layer1 = tensorflow.keras.layers.Dense(units=np.prod(shape_before_flatten), name=\"decoder_dense_1\")(dec_input_layer)\ndecoder_reshape = tensorflow.keras.layers.Reshape(target_shape=shape_before_flatten)(decoder_dense_layer1)\ndres1 = ResNetBlock(1024, 3, modify='upsample')(decoder_reshape)\ndres2 = ResNetBlock(512, 3, modify='upsample')(dres1)\ndres3 = ResNetBlock(256, 3, modify='upsample')(dres2)\ndres4 = ResNetBlock(128, 3, modify='upsample')(dres3)\ndres5 = ResNetBlock(64, 3, modify='upsample')(dres4)\ndres6 = ResNetBlock(32, 3, modify='upsample')(dres5)\ndconv = tf.keras.layers.Conv2D(3, 3, 1, padding='same', activation = tf.nn.sigmoid)(dres6)\n\ndecoder = tf.keras.models.Model(dec_input_layer, dconv, name=\"decoder_model\")","metadata":{"execution":{"iopub.status.busy":"2024-03-13T18:44:11.438292Z","iopub.execute_input":"2024-03-13T18:44:11.438626Z","iopub.status.idle":"2024-03-13T18:44:11.867900Z","shell.execute_reply.started":"2024-03-13T18:44:11.438598Z","shell.execute_reply":"2024-03-13T18:44:11.867022Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#VAE:\n\n#vae_input = tf.keras.layers.Input(shape=(512, 512, 3), name=\"VAE_input\")\nvae_encoder_output = encoder(inputs)\noutputs = decoder(vae_encoder_output[2])\nvae = tf.keras.models.Model(inputs, outputs, name=\"VAE\")","metadata":{"execution":{"iopub.status.busy":"2024-03-13T18:44:11.869051Z","iopub.execute_input":"2024-03-13T18:44:11.869363Z","iopub.status.idle":"2024-03-13T18:44:12.284929Z","shell.execute_reply.started":"2024-03-13T18:44:11.869330Z","shell.execute_reply":"2024-03-13T18:44:12.283768Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"encoder.summary()","metadata":{"execution":{"iopub.status.busy":"2024-03-13T18:44:12.286574Z","iopub.execute_input":"2024-03-13T18:44:12.287461Z","iopub.status.idle":"2024-03-13T18:44:12.338744Z","shell.execute_reply.started":"2024-03-13T18:44:12.287419Z","shell.execute_reply":"2024-03-13T18:44:12.337910Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"decoder.summary()","metadata":{"execution":{"iopub.status.busy":"2024-03-13T18:44:12.339905Z","iopub.execute_input":"2024-03-13T18:44:12.340171Z","iopub.status.idle":"2024-03-13T18:44:12.381951Z","shell.execute_reply.started":"2024-03-13T18:44:12.340147Z","shell.execute_reply":"2024-03-13T18:44:12.380892Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"vae.summary()","metadata":{"execution":{"iopub.status.busy":"2024-03-13T18:44:12.383307Z","iopub.execute_input":"2024-03-13T18:44:12.383704Z","iopub.status.idle":"2024-03-13T18:44:12.417411Z","shell.execute_reply.started":"2024-03-13T18:44:12.383669Z","shell.execute_reply":"2024-03-13T18:44:12.416575Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tensorflow.keras.losses import mse\nreconstruction_loss = mse(K.flatten(inputs), K.flatten(outputs))\nreconstruction_loss *= 512*512*3\nkl_loss = -0.5 * K.sum(1 + z_log_var - K.square(z_mean) - K.exp(z_log_var), axis=1)\nB = 1\nvae_loss = K.mean(B * reconstruction_loss + kl_loss)\nvae.add_loss(vae_loss)\nvae.add_metric(kl_loss, name=\"kl_loss\")\nvae.add_metric(reconstruction_loss, name=\"reconstruction_loss\")\nvae.compile(optimizer=tf.keras.optimizers.Adam(lr=10e-6, global_clipnorm=10e-6, clipvalue=10e-6, weight_decay=1))\n\n#vae.fit(x_train, epochs=500, batch_size=batch_size, validation_data=(x_test, None))","metadata":{"execution":{"iopub.status.busy":"2024-03-13T18:44:12.418697Z","iopub.execute_input":"2024-03-13T18:44:12.418996Z","iopub.status.idle":"2024-03-13T18:44:12.549859Z","shell.execute_reply.started":"2024-03-13T18:44:12.418970Z","shell.execute_reply":"2024-03-13T18:44:12.549084Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"np.array([next(iter(x_train))[0][0].numpy()]).shape","metadata":{"execution":{"iopub.status.busy":"2024-03-13T18:44:12.550996Z","iopub.execute_input":"2024-03-13T18:44:12.551637Z","iopub.status.idle":"2024-03-13T18:44:12.828105Z","shell.execute_reply.started":"2024-03-13T18:44:12.551602Z","shell.execute_reply":"2024-03-13T18:44:12.827076Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tqdm import tqdm\nfrom tensorflow.keras.callbacks import LambdaCallback\ndef plot_callback(epoch, logs):\n    x = np.array([next(iter(x_train))[0][0].numpy()])\n    fig, axes = plt.subplots(1, 2, figsize=(10, 5))\n\n    # Plot x\n    axes[0].imshow(x[0][:,:,0])\n    axes[0].set_yticks([])\n    axes[0].set_ylabel('Frequencies (Hz)', size=12)\n    axes[0].set_xlabel('Time (sec)', size=12)\n    # Plot out\n    out = vae.predict(x)\n    axes[1].imshow(out[0][:,:,0])\n    axes[1].set_yticks([])\n    axes[1].set_ylabel('Frequencies (Hz)', size=12)\n    axes[1].set_xlabel('Time (sec)', size=12)\n\n    plt.show()\n    \nclass PrintValidationLoss(tf.keras.callbacks.Callback):\n    def on_epoch_end(self, epoch, logs=None):\n        val_loss = logs.get('val_loss')\n        val_kl_loss = logs.get('val_kl_loss')\n        val_reconstruction_loss = logs.get(\"val_reconstruction_loss\")\n        print(f'Validation Loss: {val_loss} - kl_loss: {val_kl_loss} - reconstruction_loss: {val_reconstruction_loss}')\n# Define the LambdaCallback\nplot_callback_lambda = LambdaCallback(on_epoch_end=plot_callback)\n#m_checkpoint = tf.keras.callbacks.ModelCheckpoint(filepath='/kaggle/working', save_weights_only = True, period=5)\n\n# Fit the model\nHistory = vae.fit(x_train, epochs=10, batch_size=64, callbacks=[plot_callback_lambda,PrintValidationLoss()])","metadata":{"execution":{"iopub.status.busy":"2024-03-13T18:44:12.829259Z","iopub.execute_input":"2024-03-13T18:44:12.829615Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"vae.save_weights('final_output')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}