{"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":"none","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"},{"sourceId":7392733,"sourceType":"datasetVersion","datasetId":4297749},{"sourceId":7515232,"sourceType":"datasetVersion","datasetId":4377463},{"sourceId":7637636,"sourceType":"datasetVersion","datasetId":4450985},{"sourceId":7815460,"sourceType":"datasetVersion","datasetId":4578406},{"sourceId":7846054,"sourceType":"datasetVersion","datasetId":4600476}],"dockerImageVersionId":30665,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# <span style='color:#B22222'>|</span> HMS: <span style='color:#B22222'>Harmful Brain Activity Classification</span><span style='color:#ABABAB'> [Inference]</span>\n<!-- ## Ensemble Model of WaveNet, EfficientNetB0 and EEGNet -->\n\nThis notebook is an ensemble model of an EEGNet, WaveNet and EfficientNet. For more explanation on the particular implementations of one model, we refer to the respective training notebooks:\n- The EEGNet model can be found [here](https://www.kaggle.com/code/jasperpieterse/hms-eegnet-training).  \n- The WaveNet model is the baseline model from Moth, which can be found [here](https://www.kaggle.com/code/alejopaullier/hms-wavenet-pytorch-train).\n- The EfficientNet is the baseline model from Moth, which can be found [here](https://www.kaggle.com/code/alejopaullier/hms-efficientnetb0-pytorch-train)\n\n### <b><span style='color: #B22222 '>Table of Contents</span></b> <a class='anchor' id='top'></a>\n<div style=\" background-color: #ecf0f1 ; padding: 13px 13px; border-radius: 8px; color: white\">\n<li><a href=\"#import_libraries\">General Setup</a></li>\n<li><a href=\"#EfficientNet\">EfficientNet</a></li>\n<li><a href=\"#WaveNet\">WaveNet</a></li>\n<li><a href=\"#EEGNet\">EEGNet</a></li>    \n<li><a href=\"#submission\">Ensemble Submission</a></li>\n    \n</div>","metadata":{}},{"cell_type":"markdown","source":"# <span style='color:#B22222'>|</span> General Setup <a class='anchor' id='setup'></a> [↑](#top) \n\n\nSetup before running all three seperate models. This includes importing all libraries, loading the test data, loading of the test EEGs and the test EEG spectrograms.\n\n### Imports\n\nImports all required libraries for EEGNet, WaveNet and EfficientNet, checks for available GPUs and creates a class paths that contains file paths for all three models.","metadata":{}},{"cell_type":"code","source":"# Imports\nimport sys\nimport warnings\nimport albumentations as A\nimport gc\nimport librosa\nimport matplotlib.pyplot as plt\nimport math\nimport multiprocessing\nimport numpy as np\nimport os\nimport pandas as pd\nimport pywt\nimport random\nimport time\nimport timm\nimport torch\nimport torch.nn as nn\nimport torchvision.transforms as transforms\nimport torchvision.io \nimport torch.multiprocessing as mp\nimport torch.optim.lr_scheduler as scheduler \nimport torchmetrics\nimport pytorch_lightning as pl\n\nsys.path.append('/kaggle/input/kaggle-kl-div')\nfrom kaggle_kl_div import score\nfrom scipy.signal import butter, lfilter\nfrom albumentations.pytorch import ToTensorV2\nfrom glob import glob\nfrom torch.utils.data import DataLoader, Dataset\nfrom tqdm import tqdm\nfrom typing import Dict, List\n\nos.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)')\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    TEST_SPECTROGRAMS = \"/kaggle/input/hms-harmful-brain-activity-classification/test_spectrograms/\"\n    TRAIN_CSV    = \"/kaggle/input/hms-harmful-brain-activity-classification/train.csv\"\n#     WEIGHTS_8CH  = \"/kaggle/input/eegnet-8-channel-trained-model/\" \n    WEIGHTS_EEGNET = \"/kaggle/input/the-very-best-most-fine-weights\" ","metadata":{"execution":{"iopub.status.busy":"2024-03-15T16:57:15.768998Z","iopub.execute_input":"2024-03-15T16:57:15.769679Z","iopub.status.idle":"2024-03-15T16:57:29.084087Z","shell.execute_reply.started":"2024-03-15T16:57:15.769636Z","shell.execute_reply":"2024-03-15T16:57:29.083150Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Functions to create our own spectrograms","metadata":{}},{"cell_type":"code","source":"USE_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\ndef maddest(d, axis: int = None):\n    \"\"\"\n    Denoise function.\n    \"\"\"\n    return np.mean(np.absolute(d - np.mean(d, axis)), axis)\n\ndef denoise(x: np.ndarray, wavelet: str = 'haar', level: int = 1): \n    coeff = pywt.wavedec(x, wavelet, mode=\"per\") # multilevel 1D Discrete Wavelet Transform of data.\n    sigma = (1/0.6745) * maddest(coeff[-level])\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    output = pywt.waverec(coeff, wavelet, mode='per')\n    return output\n\ndef spectrogram_from_eeg(parquet_path, display=False):\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:\n        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    return img\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    \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(20)\ncolumns_to_average = ['seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']","metadata":{"execution":{"iopub.status.busy":"2024-03-15T16:57:29.085925Z","iopub.execute_input":"2024-03-15T16:57:29.086462Z","iopub.status.idle":"2024-03-15T16:57:29.112940Z","shell.execute_reply.started":"2024-03-15T16:57:29.086434Z","shell.execute_reply":"2024-03-15T16:57:29.112010Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Functions to load the EEG data","metadata":{}},{"cell_type":"code","source":"EEG_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\ndef eeg_plot(data):\n    \"\"\"\n    Plots the EEG data of the first channel Fp1 for patient with eeg_id\n    :param data: EEG signal.\n    \"\"\"\n    plt.plot(data[:,0])\n    plt.title(f'EEG {eeg_id} - Channel {EEG_FEATURES[0]}')\n    plt.show()\n\ndef eeg_from_parquet(parquet_path: str,display=False) -> 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    if display:\n        eeg_plot(data)\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']]","metadata":{"execution":{"iopub.status.busy":"2024-03-15T16:57:29.114071Z","iopub.execute_input":"2024-03-15T16:57:29.114402Z","iopub.status.idle":"2024-03-15T16:57:29.127119Z","shell.execute_reply.started":"2024-03-15T16:57:29.114362Z","shell.execute_reply":"2024-03-15T16:57:29.126130Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <span style='color:#B22222'>|</span> Load the test data","metadata":{}},{"cell_type":"code","source":"test_df = pd.read_csv(paths.TEST_CSV)\nprint(f\"Test dataframe shape is: {test_df.shape}\")\ntest_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-03-15T16:57:29.129247Z","iopub.execute_input":"2024-03-15T16:57:29.129522Z","iopub.status.idle":"2024-03-15T16:57:29.160445Z","shell.execute_reply.started":"2024-03-15T16:57:29.129500Z","shell.execute_reply":"2024-03-15T16:57:29.159410Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <span style='color:#B22222'>|</span> Load the test EEG spectrograms </span>\n\nLoad the test EEG spectrograms provided by Kaggle. These spectrograms are used for the EfficientNet inference.","metadata":{}},{"cell_type":"code","source":"%%time\n\npaths_spectrograms = glob(paths.TEST_SPECTROGRAMS + \"*.parquet\")\nprint(f'There are {len(paths_spectrograms)} EEG parquets')\nall_spectrograms = {}\n\nfor file_path in tqdm(paths_spectrograms):\n    aux = pd.read_parquet(file_path)\n    name = int(file_path.split(\"/\")[-1].split('.')[0])\n    all_spectrograms[name] = aux.iloc[:,1:].values\n    del aux","metadata":{"execution":{"iopub.status.busy":"2024-03-15T16:57:29.161769Z","iopub.execute_input":"2024-03-15T16:57:29.162252Z","iopub.status.idle":"2024-03-15T16:57:29.622159Z","shell.execute_reply.started":"2024-03-15T16:57:29.162221Z","shell.execute_reply":"2024-03-15T16:57:29.621070Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <span style='color:#B22222'>|</span> Load the test EEGs & create spectrograms \n\nLoad the test EEGs from Kaggle to create our own spectrograms out of it. For each of the four sides LL = left lateral, RL = right lateral, LP = left parasagittal, RP = right parasagittal a spectrogram is created and the first pair of spectrograms is visualized. These spectrograms are used for the EfficientNet inference.","metadata":{}},{"cell_type":"code","source":"%%time\n\npaths_eegs = glob(paths.TEST_EEGS + \"*.parquet\")\nprint(f'There are {len(paths_eegs)} EEGs')\nall_self_made_spectrograms = {}\ncounter = 0\n\nfor file_path in tqdm(paths_eegs):\n    eeg_id = file_path.split(\"/\")[-1].split(\".\")[0]\n    eeg_spectrogram = spectrogram_from_eeg(file_path,counter < 1)\n    all_self_made_spectrograms[int(eeg_id)] = eeg_spectrogram\n    counter += 1","metadata":{"execution":{"iopub.status.busy":"2024-03-15T16:57:29.623612Z","iopub.execute_input":"2024-03-15T16:57:29.624011Z","iopub.status.idle":"2024-03-15T16:57:41.265765Z","shell.execute_reply.started":"2024-03-15T16:57:29.623975Z","shell.execute_reply":"2024-03-15T16:57:41.264797Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <span style='color:#B22222'>|</span> Load the EEG data \n\nRead the EEG parquets and load the EEG data. These EEGs are used for the WaveNet inference.","metadata":{}},{"cell_type":"code","source":"%%time\n\nall_eegs = {}\ncounter = 0\n\nfor file_path in tqdm(paths_eegs):\n    eeg_id = file_path.split(\"/\")[-1].split(\".\")[0]\n    eeg_data = eeg_from_parquet(file_path,counter < 1)\n    all_eegs[eeg_id] = eeg_data\n    counter += 1","metadata":{"execution":{"iopub.status.busy":"2024-03-15T16:57:41.267057Z","iopub.execute_input":"2024-03-15T16:57:41.268173Z","iopub.status.idle":"2024-03-15T16:57:41.514132Z","shell.execute_reply.started":"2024-03-15T16:57:41.268135Z","shell.execute_reply":"2024-03-15T16:57:41.513147Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <span style='color:#520099'>|</span> Model 1: <span style='color:#520099'> EfficientNet <a class='anchor' id='EfficientNet'></a> [↑](#top) \n\n## <span style='color:#520099'>|</span> Configuration \n\nSet the configuration for EfficientNet and load in the model weights.","metadata":{}},{"cell_type":"code","source":"class config:\n    BATCH_SIZE = 64\n    MODEL = \"tf_efficientnet_b0\"\n    NUM_WORKERS = 0 # multiprocessing.cpu_count()\n    PRINT_FREQ = 20\n    SEED = 20\n    VISUALIZE = False\n    \n# Load in the model weights for the EfficientNet \nmodel_weights = [x for x in glob(\"/kaggle/input/hms-efficientnetb0-5-folds/*.pth\")]\nmodel_weights","metadata":{"execution":{"iopub.status.busy":"2024-03-15T16:57:41.515294Z","iopub.execute_input":"2024-03-15T16:57:41.515598Z","iopub.status.idle":"2024-03-15T16:57:41.530902Z","shell.execute_reply.started":"2024-03-15T16:57:41.515573Z","shell.execute_reply":"2024-03-15T16:57:41.530133Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## <span style='color:#520099'>|</span> Dataset \n\nCreate a custom `Dataset` to load data for the EfficientNet inference. The function `__data_generation` constructs a 8 features image out of the four spectrograms provided by Kaggle and four spectrograms made by ourselves. This generates an image of size (128,256,8) which is the output of the Dataloader.","metadata":{}},{"cell_type":"code","source":"class CustomDatasetEfficientNet(Dataset):\n    def __init__(\n        self, df: pd.DataFrame, config,\n        augment: bool = False, mode: str = 'train',\n        specs: Dict[int, np.ndarray] = all_spectrograms,\n        eeg_specs: Dict[int, np.ndarray] = all_self_made_spectrograms\n    ): \n        self.df = df\n        self.config = config\n        self.batch_size = self.config.BATCH_SIZE\n        self.augment = augment\n        self.mode = mode\n        self.spectrograms = all_spectrograms\n        self.eeg_spectrograms = all_self_made_spectrograms\n        \n    def __len__(self):\n        \"\"\"\n        Denotes the number of batches per epoch.\n        \"\"\"\n        return len(self.df)\n        \n    def __getitem__(self, index):\n        \"\"\"\n        Generate one batch of data.\n        \"\"\"\n        X, y = self.__data_generation(index)\n        if self.augment:\n            X = self.__transform(X)\n        return torch.tensor(X, dtype=torch.float32), torch.tensor(y, dtype=torch.float32)\n                        \n    def __data_generation(self, index):\n        \"\"\"\n        Generates data containing batch_size samples.\n        \"\"\"\n        X = np.zeros((128, 256, 8), dtype='float32')\n        y = np.zeros(6, dtype='float32')\n        img = np.ones((128,256), dtype='float32')\n        row = self.df.iloc[index]\n        if self.mode=='test': \n            r = 0\n        else: \n            r = int((row['min'] + row['max']) // 4)\n            \n        for region in range(4):\n            img = self.spectrograms[row.spectrogram_id][r:r+300, region*100:(region+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            # Standarize per image\n            ep = 1e-6\n            mu = np.nanmean(img.flatten())\n            std = np.nanstd(img.flatten())\n            img = (img-mu)/(std+ep)\n            img = np.nan_to_num(img, nan=0.0)\n            X[14:-14, :, region] = img[:, 22:-22] / 2.0\n            img = self.eeg_spectrograms[row.eeg_id]\n            X[:, :, 4:] = img\n                \n            if self.mode != 'test':\n                y = row[label_cols].values.astype(np.float32)\n            \n        return X, y\n    \n    def __transform(self, img):\n        transforms = A.Compose([\n            A.HorizontalFlip(p=0.5),\n        ])\n        return transforms(image=img)['image']","metadata":{"execution":{"iopub.status.busy":"2024-03-15T16:57:41.532216Z","iopub.execute_input":"2024-03-15T16:57:41.532606Z","iopub.status.idle":"2024-03-15T16:57:41.550254Z","shell.execute_reply.started":"2024-03-15T16:57:41.532574Z","shell.execute_reply":"2024-03-15T16:57:41.549354Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## <span style='color:#520099'>|</span> Dataloader\n\nGenerate the dataloader spectrogram images.","metadata":{}},{"cell_type":"code","source":"test_dataset = CustomDatasetEfficientNet(test_df, config, mode=\"test\")\ntest_loader = DataLoader(\n    test_dataset,\n    batch_size=config.BATCH_SIZE,\n    shuffle=False,\n    num_workers=config.NUM_WORKERS, pin_memory=True, drop_last=False\n)\nX, y = test_dataset[0]\nprint(f\"X shape: {X.shape}\")\nprint(f\"y shape: {y.shape}\")","metadata":{"execution":{"iopub.status.busy":"2024-03-15T16:57:41.554459Z","iopub.execute_input":"2024-03-15T16:57:41.554756Z","iopub.status.idle":"2024-03-15T16:57:41.597766Z","shell.execute_reply.started":"2024-03-15T16:57:41.554731Z","shell.execute_reply":"2024-03-15T16:57:41.596777Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## <span style='color:#520099'>|</span> Model \n\nThe spectrogram images of size (128,256,8) are flattened to the size of (512,512,3) to be feed into the EfficientNet model.","metadata":{}},{"cell_type":"code","source":"class CustomModel(nn.Module):\n    def __init__(self, config, num_classes: int = 6):\n        super(CustomModel, self).__init__()\n        self.USE_KAGGLE_SPECTROGRAMS = True\n        self.USE_EEG_SPECTROGRAMS = True\n        self.model = timm.create_model(\n            config.MODEL,\n            pretrained=False\n        )\n        self.features = nn.Sequential(*list(self.model.children())[:-2])\n        self.custom_layers = nn.Sequential(\n            nn.AdaptiveAvgPool2d(1),\n            nn.Flatten(),\n            nn.Linear(self.model.num_features, num_classes)\n        )\n\n    def __reshape_input(self, x):\n        \"\"\"\n        Reshapes input (128, 256, 8) -> (512, 512, 3) monotone image.\n        \"\"\" \n        # === Get Kaggle spectrograms ===\n        spectrograms = [x[:, :, :, i:i+1] for i in range(4)]\n        spectrograms = torch.cat(spectrograms, dim=1)\n        \n        # === Get self made spectrograms ===\n        eegs = [x[:, :, :, i:i+1] for i in range(4,8)]\n        eegs = torch.cat(eegs, dim=1)\n        \n        # === Reshape (512,512,3) ===\n        if self.USE_KAGGLE_SPECTROGRAMS & self.USE_EEG_SPECTROGRAMS:\n            x = torch.cat([spectrograms, eegs], dim=2)\n        elif self.USE_EEG_SPECTROGRAMS:\n            x = eegs\n        else:\n            x = spectrograms\n            \n        x = torch.cat([x,x,x], dim=3)\n        x = x.permute(0, 3, 1, 2)\n        return x\n    \n    def forward(self, x):\n        x = self.__reshape_input(x)\n        x = self.features(x)\n        x = self.custom_layers(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2024-03-15T16:57:41.598939Z","iopub.execute_input":"2024-03-15T16:57:41.599253Z","iopub.status.idle":"2024-03-15T16:57:41.610942Z","shell.execute_reply.started":"2024-03-15T16:57:41.599227Z","shell.execute_reply":"2024-03-15T16:57:41.610016Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## <span style='color:#520099'>|</span> Inference function ","metadata":{}},{"cell_type":"code","source":"def inference_function_efficientnet(test_loader, model, device):\n    model.eval()\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, (X, y) in enumerate(tqdm_test_loader):\n            X = X.to(device)\n            y = y.to(device)\n            batch_size = y.size(0)\n            with torch.no_grad():\n                y_preds = model(X)\n            y_preds = softmax(y_preds)\n            preds.append(y_preds.to('cpu').numpy()) \n                \n    prediction_dict[\"predictions\"] = np.concatenate(preds) \n    return prediction_dict","metadata":{"execution":{"iopub.status.busy":"2024-03-15T16:57:41.611922Z","iopub.execute_input":"2024-03-15T16:57:41.612219Z","iopub.status.idle":"2024-03-15T16:57:41.624797Z","shell.execute_reply.started":"2024-03-15T16:57:41.612187Z","shell.execute_reply":"2024-03-15T16:57:41.624029Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## <span style='color:#520099'>|</span> Inference","metadata":{}},{"cell_type":"code","source":"predictions_efficientnet = []\n\nfor model_weight in model_weights:\n    test_dataset = CustomDatasetEfficientNet(test_df, config, mode=\"test\", augment=False)\n    train_loader = DataLoader(\n        test_dataset,\n        batch_size=config.BATCH_SIZE,\n        shuffle=False,\n        num_workers=config.NUM_WORKERS,\n        pin_memory=True, drop_last=False\n    )\n    model = CustomModel(config)\n    checkpoint = torch.load(model_weight)\n    model.load_state_dict(checkpoint[\"model\"])\n    model.to(device)\n    prediction_dict_efficientnet = inference_function_efficientnet(test_loader, model, device)\n    predictions_efficientnet.append(prediction_dict_efficientnet[\"predictions\"])\n    torch.cuda.empty_cache()\n    gc.collect()\n    \npredictions_efficientnet = np.array(predictions_efficientnet)\npredictions_efficientnet = np.mean(predictions_efficientnet, axis=0)","metadata":{"execution":{"iopub.status.busy":"2024-03-15T16:57:41.625849Z","iopub.execute_input":"2024-03-15T16:57:41.626136Z","iopub.status.idle":"2024-03-15T16:57:47.280806Z","shell.execute_reply.started":"2024-03-15T16:57:41.626092Z","shell.execute_reply":"2024-03-15T16:57:47.279909Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## <span style='color:#520099'>|</span> Print prediction","metadata":{}},{"cell_type":"code","source":"efficientnet = pd.DataFrame({'eeg_id': test_df.eeg_id.values})\nefficientnet[columns_to_average] = predictions_efficientnet\nefficientnet[:1].head()","metadata":{"execution":{"iopub.status.busy":"2024-03-15T16:57:47.281879Z","iopub.execute_input":"2024-03-15T16:57:47.282196Z","iopub.status.idle":"2024-03-15T16:57:47.298083Z","shell.execute_reply.started":"2024-03-15T16:57:47.282171Z","shell.execute_reply":"2024-03-15T16:57:47.297152Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Clean up\n\nDelete data that will be no longer used to make sure that we keep enough disk space. ","metadata":{}},{"cell_type":"code","source":"del efficientnet\ndel paths_eegs\ndel paths_spectrograms\ndel all_spectrograms\ndel all_self_made_spectrograms","metadata":{"execution":{"iopub.status.busy":"2024-03-15T16:57:47.299559Z","iopub.execute_input":"2024-03-15T16:57:47.299896Z","iopub.status.idle":"2024-03-15T16:57:47.314511Z","shell.execute_reply.started":"2024-03-15T16:57:47.299862Z","shell.execute_reply":"2024-03-15T16:57:47.313445Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <span style='color:#50C878'>|</span> Model 2: <span style='color:#50C878'> WaveNet <a class='anchor' id='WaveNet'></a> [↑](#top) \n\n## <span style='color:#50C878'>|</span> Configuration \n\nSet the configuration for WaveNet and load in the model weights.","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# Load in the model weights of the WaveNet\nmodel_weights = [x for x in glob(\"/kaggle/input/hms-wavenet/*.pth\")]\nmodel_weights ","metadata":{"execution":{"iopub.status.busy":"2024-03-15T16:57:47.315717Z","iopub.execute_input":"2024-03-15T16:57:47.316032Z","iopub.status.idle":"2024-03-15T16:57:47.332801Z","shell.execute_reply.started":"2024-03-15T16:57:47.316007Z","shell.execute_reply":"2024-03-15T16:57:47.331839Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### <span style='color:#50C878'>|</span> Butter Low-Pass Filter \n\nFunction to apply butter low-pass filtering.","metadata":{}},{"cell_type":"code","source":"def 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-03-15T16:57:47.333904Z","iopub.execute_input":"2024-03-15T16:57:47.334265Z","iopub.status.idle":"2024-03-15T16:57:47.340710Z","shell.execute_reply.started":"2024-03-15T16:57:47.334231Z","shell.execute_reply":"2024-03-15T16:57:47.339846Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## <span style='color:#50C878'>|</span> Dataset\n\nCreate a custom `Dataset` to load data for the WaveNet inference. The `__data_generation` subtracts EEG signals from neighbouring channels. The difference in EEG signals between those channels is taken as output for the Dataloader.","metadata":{}},{"cell_type":"code","source":"class CustomDatasetWavenet(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[str(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-03-15T16:57:47.341985Z","iopub.execute_input":"2024-03-15T16:57:47.342402Z","iopub.status.idle":"2024-03-15T16:57:47.356588Z","shell.execute_reply.started":"2024-03-15T16:57:47.342375Z","shell.execute_reply":"2024-03-15T16:57:47.355643Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## <span style='color:#50C878'>|</span> Dataloader \n\nGenerate the dataloader.","metadata":{}},{"cell_type":"code","source":"test_dataset = CustomDatasetWavenet(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, pin_memory=True,drop_last=False\n)\nX = test_dataset[0]\nprint(f\"X shape: {X['X'].shape}\")","metadata":{"execution":{"iopub.status.busy":"2024-03-15T16:57:47.357727Z","iopub.execute_input":"2024-03-15T16:57:47.358049Z","iopub.status.idle":"2024-03-15T16:57:47.377550Z","shell.execute_reply.started":"2024-03-15T16:57:47.358020Z","shell.execute_reply":"2024-03-15T16:57:47.376483Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## <span style='color:#50C878'>|</span> 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-03-15T16:57:47.381220Z","iopub.execute_input":"2024-03-15T16:57:47.381897Z","iopub.status.idle":"2024-03-15T16:57:47.435314Z","shell.execute_reply.started":"2024-03-15T16:57:47.381868Z","shell.execute_reply":"2024-03-15T16:57:47.434298Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## <span style='color:#50C878'>|</span> Inference Function </span>","metadata":{}},{"cell_type":"code","source":"def inference_function_wavenet(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-03-15T16:57:47.436340Z","iopub.execute_input":"2024-03-15T16:57:47.436626Z","iopub.status.idle":"2024-03-15T16:57:47.443725Z","shell.execute_reply.started":"2024-03-15T16:57:47.436602Z","shell.execute_reply":"2024-03-15T16:57:47.442693Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## <span style='color:#50C878'>|</span> Inference","metadata":{}},{"cell_type":"code","source":"predictions_wavenet = []\n\nfor model_weight in model_weights:\n    test_dataset = CustomDatasetWavenet(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_wavenet = inference_function_wavenet(test_loader, model, device)\n    predictions_wavenet.append(prediction_dict_wavenet[\"predictions\"])\n    torch.cuda.empty_cache()\n    gc.collect()\n    \npredictions_wavenet = np.array(predictions_wavenet)\npredictions_wavenet = np.mean(predictions_wavenet, axis=0)","metadata":{"execution":{"iopub.status.busy":"2024-03-15T16:57:47.445033Z","iopub.execute_input":"2024-03-15T16:57:47.445359Z","iopub.status.idle":"2024-03-15T16:57:49.616361Z","shell.execute_reply.started":"2024-03-15T16:57:47.445319Z","shell.execute_reply":"2024-03-15T16:57:49.615475Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## <span style='color:#50C878'>|</span> Print prediction","metadata":{}},{"cell_type":"code","source":"wavenet = pd.DataFrame({'eeg_id': test_df.eeg_id.values})\nwavenet[columns_to_average] = predictions_wavenet\nwavenet[:1].head()","metadata":{"execution":{"iopub.status.busy":"2024-03-15T16:57:49.617448Z","iopub.execute_input":"2024-03-15T16:57:49.617701Z","iopub.status.idle":"2024-03-15T16:57:49.633471Z","shell.execute_reply.started":"2024-03-15T16:57:49.617679Z","shell.execute_reply":"2024-03-15T16:57:49.632474Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del wavenet","metadata":{"execution":{"iopub.status.busy":"2024-03-15T16:57:49.634823Z","iopub.execute_input":"2024-03-15T16:57:49.635119Z","iopub.status.idle":"2024-03-15T16:57:49.642258Z","shell.execute_reply.started":"2024-03-15T16:57:49.635081Z","shell.execute_reply":"2024-03-15T16:57:49.641386Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <span style='color:#F1A424'>|</span> Model 3: <span style='color:#F1A424'> EEGNet <a class='anchor' id='EEGNet'></a> [↑](#top) \n\n## <span style='color:#F1A424'>|</span> Configuration \n\nSet the configuration for WaveNet and load in the model weights.","metadata":{}},{"cell_type":"code","source":"class Config:\n    #Data\n    CUSTOM_FEATURES =['Fp1','T3','C3','O1','Fp2','C4','T4','O2'] # List containing what EEG features to use to create data. Leave empty to use all2\n    NUM_CLASSES = 6 \n    FEATURE_ENGINEERING = True\n    \n    #Training\n    NUM_WORKERS = 2         \n    PRECISION   = 32        \n    BATCH_SIZE  = 32\n    EPOCHS      = 30 \n    PATIENCE    = 20   \n    P_DROPOUT   = 0.0\n    SEED        = 2024\n    LR          = 8e-3\n    FOLDS       = 5 \n    \n    #Model parameters\n    KERNELS = [3,5,7,9]\n    FIXED_KERNEL_SIZE = 5 #Size of the kernel used in the fixed blocks\n    NUM_FEATURE_MAPS = 24 #Number of feature maps (channels) outputed for each parallel convolution (num_channels = num_kernels * num_feature_maps)\n    RESNET_BLOCKS = 9     #Amount of ResNet blocks to be used\n    if CUSTOM_FEATURES:\n        NUM_CHANNELS = len(CUSTOM_FEATURES)\n    else:  \n        NUM_CHANNELS = 20     # Amount of channels used for CNN\n\n\n    #Augmentation\n    PRETRAINED = False    #Whether to use a pretrained EEG net        \n    WEIGHT_DECAY = 0.01\n    USE_MIXUP = False   # Whether to use mixup augmentation\n    MIXUP_ALPHA = 0.1   # Alpha parameter for mixup\n\npl.seed_everything(Config.SEED, workers=True)\n\nwarnings.filterwarnings('ignore')\nprint(\"Using float 32 for intermediate calculations\")\ntorch.set_float32_matmul_precision('high')","metadata":{"execution":{"iopub.status.busy":"2024-03-15T16:57:49.643597Z","iopub.execute_input":"2024-03-15T16:57:49.643922Z","iopub.status.idle":"2024-03-15T16:57:49.655481Z","shell.execute_reply.started":"2024-03-15T16:57:49.643896Z","shell.execute_reply":"2024-03-15T16:57:49.654528Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## <span style='color:#F1A424'>|</span> Pre-processing functions\n\nFunctions of possible augmentations to apply to our data. These will be used in creating the dataset.","metadata":{}},{"cell_type":"code","source":"def quantize_data(data, classes):\n    \"\"\"Quantizes data using Mu-law encoding.\n    Args:\n        classes (int): The number of quantization levels.\n    \"\"\"\n    mu_x = mu_law_encoding(data, classes)\n#     bins = np.linspace(-1, 1, classes)  # Create equally spaced bins\n#     quantized = np.digitize(mu_x, bins) - 1  # Assign data to bins\n    return mu_x\n\ndef mu_law_encoding(data, mu):\n    \"\"\"Performs Mu-law encoding on a NumPy array.\"\"\"\n    mu_x = np.sign(data) * np.log(1 + mu * np.abs(data)) / np.log(mu + 1)\n    return mu_x\n\ndef mu_law_expansion(data, mu):\n    \"\"\"Performs Mu-law expansion (decoding) on a NumPy array.\"\"\"\n    s = np.sign(data) * (np.exp(np.abs(data) * np.log(mu + 1)) - 1) / mu\n    return s","metadata":{"execution":{"iopub.status.busy":"2024-03-15T16:57:49.656871Z","iopub.execute_input":"2024-03-15T16:57:49.657161Z","iopub.status.idle":"2024-03-15T16:57:49.667026Z","shell.execute_reply.started":"2024-03-15T16:57:49.657128Z","shell.execute_reply":"2024-03-15T16:57:49.666151Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## <span style='color:#F1A424'>|</span> Dataset\n\nCreate a custom `Dataset` to load data. Here we can easily implement custom loading mechanisms, data pre-processing, and augmentation techniques. Also allows us to use the Pytorch Dataloader, streamlining data loading and batching.","metadata":{}},{"cell_type":"code","source":"class EEGDataset(torch.utils.data.Dataset):\n    def __init__(self, data, eegs=None, augmentations=None, test=False): \n        self.data = data  \n        self.eegs = eegs \n        self.augmentations = augmentations  \n        self.test = test  # Flag to indicate test mode\n\n    def __len__(self):\n        return len(self.data)  # Return the number of samples in the dataset\n\n    def __getitem__(self, index):\n        # Get a single data sample and its label (if not in test mode)\n        row = self.data.iloc[index]       \n        data = self.eegs[str(row.eeg_id)]  # Load EEG data based on an ID\n\n        if Config.FEATURE_ENGINEERING:\n            if len(Config.CUSTOM_FEATURES) == 8:\n\n                X = np.zeros((10_000, Config.NUM_CHANNELS), dtype='float32')\n                    \n                # === Feature engineering ===\n                #LL Chain\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                #LP Chain\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                #RP Chain\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                #RR Chain\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            elif not Config.CUSTOM_FEATURES:\n                Config.NUM_CHANNELS = 16 #update number of channels\n                X = np.zeros((10_000, Config.NUM_CHANNELS), dtype='float16') #Full data cannot be represented with float32 anyway\n\n                #LL Chain\n                X[:,0] = data[:,feature_to_index['Fp1']] - data[:,feature_to_index['F7']]\n                X[:,1] = data[:,feature_to_index['F7']]  - data[:,feature_to_index['T3']]\n                X[:,2] = data[:,feature_to_index['T3']] - data[:,feature_to_index['T5']]\n                X[:,3] = data[:,feature_to_index['T5']]  - data[:,feature_to_index['O1']]\n\n                #LP Chain\n                X[:,4] = data[:,feature_to_index['Fp1']] - data[:,feature_to_index['F3']]\n                X[:,5] = data[:,feature_to_index['C4']]  - data[:,feature_to_index['C3']]\n                X[:,6] = data[:,feature_to_index['Fp2']] - data[:,feature_to_index['P3']]\n                X[:,7] = data[:,feature_to_index['T4']]  - data[:,feature_to_index['O1']]\n\n                #RP Chain\n                X[:,8] = data[:,feature_to_index['Fp2']] - data[:,feature_to_index['F4']]\n                X[:,9] = data[:,feature_to_index['F4']]  - data[:,feature_to_index['C4']]\n                X[:,10] = data[:,feature_to_index['C4']] - data[:,feature_to_index['P4']]\n                X[:,11] = data[:,feature_to_index['P4']]  - data[:,feature_to_index['O2']]\n\n                #RR Chain\n                X[:,12] = data[:,feature_to_index['Fp2']] - data[:,feature_to_index['F8']]\n                X[:,13] = data[:,feature_to_index['F8']]  - data[:,feature_to_index['T4']]\n                X[:,14] = data[:,feature_to_index['T4']] - data[:,feature_to_index['T6']]\n                X[:,15] = data[:,feature_to_index['T6']]  - data[:,feature_to_index['O2']]\n                \n        else:\n\n            X = data\n\n        # === Standarize ===\n        X = np.clip(X,-1024,1024)    \n        X = np.nan_to_num(X, nan=0) / 32.0  \n        \n        # === Preprocess ===\n        X = butter_lowpass_filter(X)  # Apply a low-pass filter \n        X = quantize_data(X, 1)       # Apply quantization\n        samples = torch.from_numpy(X).float()  # Convert to a PyTorch tensor\n        samples = samples.permute(1, 0)  # Adjust tensor shape to get (time_steps, features)\n\n        if not self.test:\n            label = row[columns_to_average]         # Get label from metadata\n            label = torch.tensor(label).float()  \n            return samples, label        # Return preprocessed data and label\n        else:\n            return samples               # Return preprocessed data only (for testing)","metadata":{"execution":{"iopub.status.busy":"2024-03-15T16:57:49.668269Z","iopub.execute_input":"2024-03-15T16:57:49.668547Z","iopub.status.idle":"2024-03-15T16:57:49.694095Z","shell.execute_reply.started":"2024-03-15T16:57:49.668523Z","shell.execute_reply":"2024-03-15T16:57:49.693038Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## <span style='color:#F1A424'>|</span> Dataloader ","metadata":{}},{"cell_type":"code","source":"def get_test_dls(df_test):\n    ds_test = EEGDataset(\n        df_test, \n        eegs=all_eegs,\n        augmentations = None,\n        test = True\n    )\n    dl_test = DataLoader(ds_test, batch_size=Config.BATCH_SIZE , shuffle=False, num_workers = Config.NUM_WORKERS)    \n    return dl_test, ds_test","metadata":{"execution":{"iopub.status.busy":"2024-03-15T16:57:49.700039Z","iopub.execute_input":"2024-03-15T16:57:49.700348Z","iopub.status.idle":"2024-03-15T16:57:49.706360Z","shell.execute_reply.started":"2024-03-15T16:57:49.700324Z","shell.execute_reply":"2024-03-15T16:57:49.705444Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## <span style='color:#F1A424'>|</span> Model ","metadata":{}},{"cell_type":"code","source":"class ResNet_1D_Block(nn.Module):\n    \"\"\"\n    Implements a ResNet-inspired residual block for 1D convolutional networks.\n\n    Args:\n        in_channels (int): Number of input channels.\n        out_channels (int): Number of output channels.\n        kernel_size (int): Kernel size for the convolutions.\n        stride (int): Stride for the convolutions.\n        padding (int): Padding for the convolutions.\n        downsampling (bool): Whether to apply downsampling.\n    \"\"\"\n\n    def __init__(self, in_channels, out_channels, kernel_size, stride, padding, downsampling):\n        super(ResNet_1D_Block, self).__init__()\n        self.bn1 = nn.BatchNorm1d(num_features=in_channels)\n        self.relu = nn.ReLU(inplace=False)                               \n        self.dropout = nn.Dropout(p=Config.P_DROPOUT, inplace=False)    \n\n        self.conv1 = nn.Conv1d(in_channels=in_channels, out_channels=out_channels, kernel_size=kernel_size,\n                               stride=stride, padding=padding, bias=False) \n        self.bn2 = nn.BatchNorm1d(num_features=out_channels)                \n\n        self.conv2 = nn.Conv1d(in_channels=out_channels, out_channels=out_channels, kernel_size=kernel_size,\n                               stride=stride, padding=padding, bias=False)  \n        self.maxpool = nn.MaxPool1d(kernel_size=2, stride=2, padding=0)   \n        self.downsampling = downsampling\n        self.downsample_layer = nn.MaxPool1d(kernel_size=2, stride=2, padding=0) \n\n    def forward(self, x):\n        \"Basically applies a BN, Relu, Dropout and 1D Convolution twice and then maxpools\"\n        identity = x  # Store input for residual connection\n\n        out = self.bn1(x)\n        out = self.relu(out)\n        out = self.dropout(out)\n        out = self.conv1(out)\n        out = self.bn2(out)\n        out = self.relu(out)\n        out = self.dropout(out)\n        out = self.conv2(out)\n\n        out = self.maxpool(out)  \n        if self.downsampling:\n            identity = self.downsample_layer(x)  # Apply downsampling if needed\n\n        out += identity  # Add the orginal input (residual connection)\n        return out\n\n    \nclass EEGNet(nn.Module):\n    \"\"\"\n    EEGNet architecture: A convolutional neural network with ResNet-style blocks designed for EEG signal classification.\n\n    Args:\n        kernels (list): A list of kernel sizes for the initial parallel convolutions.\n        in_channels (int, optional): Number of input channels (e.g., number of electrodes). \n                                    Defaults to 20.\n        fixed_kernel_size (int, optional): Kernel size for subsequent shared convolutions. \n                                           Defaults to 17.\n        NUM_CLASSES (int, optional): Number of output classes for classification. \n                                     Defaults to 6.\n    \"\"\"\n\n    def __init__(self, kernels, num_feature_maps = Config.NUM_FEATURE_MAPS, in_channels=Config.NUM_CHANNELS, \n                       fixed_kernel_size=Config.FIXED_KERNEL_SIZE, num_classes=6):\n        super(EEGNet, self).__init__()\n        self.kernels = kernels\n        self.planes = num_feature_maps  # Initial number of feature maps (=planes) outputed for each kernel (and used in 1D conv and ResNet)\n        self.in_channels = in_channels\n        \n        #----DEFINE LAYERS TO BE USED----\n        self.parallel_conv = nn.ModuleList()  # Parallel convolutions with varying kernel sizes\n        for kernel_size in kernels:\n            conv = nn.Conv1d(in_channels, self.planes, kernel_size, stride=1, padding=0, bias=False)\n            self.parallel_conv.append(conv)\n\n        # Batch normalizations, ReLU activation\n        self.bn1 = nn.BatchNorm1d(self.planes)\n        self.relu = nn.ReLU(inplace=False)\n\n        # Shared convolution and ResNet blocks\n        self.conv1 = nn.Conv1d(self.planes, self.planes, fixed_kernel_size, stride=2, padding=2, bias=False)\n        self.block = self._make_resnet_layer(fixed_kernel_size, stride=1, padding=fixed_kernel_size//2)\n\n        # Downsampling Layers \n        self.bn2 = nn.BatchNorm1d(self.planes)\n        self.avgpool = nn.AvgPool1d(6, stride=6, padding=2)\n        \n        # Gated Recurrent Unit (GRU) for processing sequential patterns in the EEG data\n        self.rnn = nn.GRU(self.in_channels, hidden_size=128, num_layers=1, bidirectional=True)\n        \n        #Final FC layer\n        self.fc = nn.Linear(424, num_classes) \n\n    def _make_resnet_layer(self, kernel_size, stride, blocks=Config.RESNET_BLOCKS, padding=0):\n        \"\"\"Creates a sequence of ResNet blocks.\"\"\"\n        layers = []\n        for _ in range(blocks):\n            downsampling = nn.MaxPool1d(2, stride=2, padding=0)  # Downsampling if needed\n            layers.append(ResNet_1D_Block(self.planes, self.planes, kernel_size, stride, padding, downsampling))\n        return nn.Sequential(*layers)\n\n    def forward(self, x):\n        \"\"\"Forward pass through the network. First separate parallel convolutions \n           which are then concetaned to go through 1D conv and ResNet block\"\"\"\n        # Parallel convolutions, list containing output of different kernels applied to input\n        out_sep = [conv(x) for conv in self.parallel_conv]  \n\n        # Concatenate outputs\n        out = torch.cat(out_sep, dim=2) \n        \n        # 1D Convolution\n        out = self.bn1(out)\n        out = self.relu(out)\n        out = self.conv1(out)\n        \n        # ResNet Block [Has BN and ReLu build in]\n        out = self.block(out)\n        \n        # Avgerage Pooling\n        out = self.bn2(out)\n        out = self.relu(out)\n        out = self.avgpool(out)\n\n        # Flatten for dense layers\n        out = out.reshape(out.shape[0], -1)  \n\n        # Pass through GRU for sequence representations\n        rnn_out, _ = self.rnn(x.permute(0, 2, 1))   \n        new_rnn_h = rnn_out[:, -1, :]  \n        new_out = torch.cat([out, new_rnn_h], dim=1)  # Combine features\n        result = self.fc(new_out)  # Final classification layer\n\n        return result","metadata":{"execution":{"iopub.status.busy":"2024-03-15T16:57:49.707786Z","iopub.execute_input":"2024-03-15T16:57:49.708068Z","iopub.status.idle":"2024-03-15T16:57:49.732976Z","shell.execute_reply.started":"2024-03-15T16:57:49.708046Z","shell.execute_reply":"2024-03-15T16:57:49.731987Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## <span style='color:#F1A424'>|</span> Optimizer & Scheduler ","metadata":{}},{"cell_type":"code","source":"def get_optimizer(lr, params):\n    model_optimizer = torch.optim.Adam(\n            filter(lambda p: p.requires_grad, params), \n            lr=lr,\n            weight_decay=Config.weight_decay\n        )\n    interval = \"epoch\"\n    \n    lr_scheduler = CosineAnnealingWarmRestarts(\n                            model_optimizer, \n                            T_0=Config.epochs, \n                            T_mult=1, \n                            eta_min=1e-5, \n                            last_epoch=-1\n                        )\n\n    return {\n        \"optimizer\": model_optimizer, \n        \"lr_scheduler\": {\n            \"scheduler\": lr_scheduler,\n            \"interval\": interval,\n            \"monitor\": \"val_loss\",\n            \"frequency\": 1\n        }\n    }","metadata":{"execution":{"iopub.status.busy":"2024-03-15T16:57:49.734320Z","iopub.execute_input":"2024-03-15T16:57:49.734659Z","iopub.status.idle":"2024-03-15T16:57:49.745689Z","shell.execute_reply.started":"2024-03-15T16:57:49.734629Z","shell.execute_reply":"2024-03-15T16:57:49.744804Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## <span style='color:#F1A424'>|</span> Loss function ","metadata":{}},{"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.log_softmax(y,  dim=1)\n        loss = super().forward(y, t)\n\n        return loss","metadata":{"execution":{"iopub.status.busy":"2024-03-15T16:57:49.746772Z","iopub.execute_input":"2024-03-15T16:57:49.747066Z","iopub.status.idle":"2024-03-15T16:57:49.758549Z","shell.execute_reply.started":"2024-03-15T16:57:49.747043Z","shell.execute_reply":"2024-03-15T16:57:49.757714Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## <span style='color:#F1A424'>|</span> PyTorch Lightning Model ","metadata":{}},{"cell_type":"code","source":"class EEGModel(pl.LightningModule):\n    \"\"\"\n    Pytorch Lightning module for EEG classification.\n    \"\"\"\n\n    def __init__(self, num_classes = Config.NUM_CLASSES, pretrained = Config.PRETRAINED, fold = -1):\n        super().__init__()\n        self.num_classes = num_classes\n        self.fold = fold\n\n        # Create the EEGNet backbone with specified kernel sizes\n        self.backbone = EEGNet(kernels=Config.KERNELS, in_channels=Config.NUM_CHANNELS, \n                               fixed_kernel_size=Config.FIXED_KERNEL_SIZE, num_classes=Config.NUM_CLASSES)\n\n        # Loss function for multi-class classification\n        self.loss_function = KLDivLossWithLogits() \n\n        self.validation_step_outputs = []  # Storage for validation results\n        self.lin = nn.Softmax(dim=1)       # Softmax for probability outputs \n        self.best_score = 1000.0           # Track best validation score\n\n    def forward(self, images):\n        \"\"\"Forward pass through the EEGNet model.\"\"\"\n        logits = self.backbone(images)  # Extract features and get predictions\n        return logits\n\n    def configure_optimizers(self):\n        \"\"\"Set up the optimizer (using configuration parameters).\"\"\"\n        return get_optimizer(lr=Config.LR, params=self.parameters())\n\n\n    def train_with_mixup(self, X, y):\n        X, y_a, y_b, lam = mixup_data(X, y, alpha=Config.MIXUP_ALPHA)\n        y_pred = self(X)\n        loss_mixup = mixup_criterion(KLDivLossWithLogits(), y_pred, y_a, y_b, lam)\n        return loss_mixup\n\n    def training_step(self, batch, batch_idx):\n        image, target = batch        \n        if Config.USE_MIXUP:\n            loss = self.train_with_mixup(image, target)\n        else:\n            y_pred = self(image)\n            loss = self.loss_function(y_pred,target)\n\n        self.log(\"train_loss\", loss, on_step=True, on_epoch=True, prog_bar=True)\n        return loss        \n\n    def validation_step(self, batch, batch_idx):\n        image, target = batch \n        y_pred = self(image)\n        val_loss = self.loss_function(y_pred, target)\n        self.log(\"val_loss\", val_loss, on_step=True, on_epoch=True, logger=True, prog_bar=True, sync_dist=True)\n        self.validation_step_outputs.append({\"val_loss\": val_loss, \"logits\": y_pred, \"targets\": target})\n\n        return {\"val_loss\": val_loss, \"logits\": y_pred, \"targets\": target}\n\n    def train_dataloader(self):\n        return self._train_dataloader \n    \n    def validation_dataloader(self):\n        return self._validation_dataloader\n    \n    def on_validation_epoch_end(self):\n        outputs = self.validation_step_outputs\n        # print(len(outputs))\n        avg_loss = torch.stack([x['val_loss'] for x in outputs]).mean()\n        output_val = nn.Softmax(dim=1)(torch.cat([x['logits'] for x in outputs],dim=0)).cpu().detach().numpy()\n        target_val = torch.cat([x['targets'] for x in outputs],dim=0).cpu().detach().numpy()\n        self.validation_step_outputs = []\n\n        val_df = pd.DataFrame(target_val, columns = list(columns_to_average))\n        pred_df = pd.DataFrame(output_val, columns = list(columns_to_average))\n\n        val_df['id'] = [f'id_{i}' for i in range(len(val_df))] \n        pred_df['id'] = [f'id_{i}' for i in range(len(pred_df))] \n\n\n        avg_score = score(val_df, pred_df, row_id_column_name = 'id')\n\n        if avg_score < self.best_score:\n            print(f'Fold {self.fold}: Epoch {self.current_epoch} validation loss {avg_loss}')\n            print(f'Fold {self.fold}: Epoch {self.current_epoch} validation KDL score {avg_score}')\n            self.best_score = avg_score\n            # val_df.to_csv(f'{paths.OUTPUT_DIR}/val_df_f{self.fold}.csv',index=False)\n            # pred_df.to_csv(f'{paths.OUTPUT_DIR}/pred_df_f{self.fold}.csv',index=False)\n        \n        return {'val_loss': avg_loss,'val_cmap':avg_score}","metadata":{"execution":{"iopub.status.busy":"2024-03-15T16:57:49.759749Z","iopub.execute_input":"2024-03-15T16:57:49.760465Z","iopub.status.idle":"2024-03-15T16:57:49.780185Z","shell.execute_reply.started":"2024-03-15T16:57:49.760441Z","shell.execute_reply":"2024-03-15T16:57:49.779360Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## <span style='color:#F1A424'>|</span> Inference functions\n\nWe will use a seperate prediction function to generate predictions of the trained model, outside the Pytorch Lightning environment","metadata":{}},{"cell_type":"code","source":"def inference_function_eegnet(data_loader, model):      \n    model.to('cuda')\n    model.eval()    \n    predictions = []\n    for batch in tqdm(data_loader):\n\n        with torch.no_grad():\n            x = batch\n            x = x.cuda()\n            # inputs = {key:val.reshape(val.shape[0], -1).to(config.device) for key,val in batch.items()}\n            outputs = model(x)\n            outputs = nn.Softmax(dim=1)(outputs)\n        predictions.extend(outputs.detach().cpu().numpy())\n    predictions = np.vstack(predictions)\n    return predictions\n\ndef run_inference(fold_id, Config):\n    logger = None\n    pred_cols = [f'pred_{t}' for t in columns_to_average]\n    df_test = test_df.copy()\n    dl_test, ds_test = get_test_dls(df_test)\n    \n\n#     print(f\"Running inference model '{paths.WEIGHTS_EEGNET}/eegnet_best_loss_fold{fold_id}.ckpt'..\")\n    \n    model = EEGModel.load_from_checkpoint(f'{paths.WEIGHTS_EEGNET}/eegnet_best_loss_fold{fold_id}.ckpt',map_location='cuda:0',\n                                          train_dataloader=None,validation_dataloader=None,config=Config)\n    \n    preds = inference_function_eegnet(dl_test, model)  \n    gc.collect()\n    return preds","metadata":{"execution":{"iopub.status.busy":"2024-03-15T16:57:49.781244Z","iopub.execute_input":"2024-03-15T16:57:49.781564Z","iopub.status.idle":"2024-03-15T16:57:49.794417Z","shell.execute_reply.started":"2024-03-15T16:57:49.781533Z","shell.execute_reply":"2024-03-15T16:57:49.793573Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## <span style='color:#F1A424'>|</span> Inference","metadata":{}},{"cell_type":"code","source":"predictions_eegnet = []\n\nfor en,f in enumerate(range(Config.FOLDS)):\n    preds = run_inference(f, Config)\n    predictions_eegnet.append(preds)\n    \npredictions_eegnet = np.array(predictions_eegnet)\npredictions_eegnet = np.mean(predictions_eegnet, axis=0)","metadata":{"execution":{"iopub.status.busy":"2024-03-15T16:57:49.795407Z","iopub.execute_input":"2024-03-15T16:57:49.795667Z","iopub.status.idle":"2024-03-15T16:57:53.147481Z","shell.execute_reply.started":"2024-03-15T16:57:49.795646Z","shell.execute_reply":"2024-03-15T16:57:53.146576Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## <span style='color:#F1A424'>|</span> Print prediction ","metadata":{}},{"cell_type":"code","source":"eegnet = pd.DataFrame({'eeg_id': test_df.eeg_id.values})\neegnet[columns_to_average] = predictions_eegnet\neegnet.head()","metadata":{"execution":{"iopub.status.busy":"2024-03-15T16:57:53.148955Z","iopub.execute_input":"2024-03-15T16:57:53.149899Z","iopub.status.idle":"2024-03-15T16:57:53.166960Z","shell.execute_reply.started":"2024-03-15T16:57:53.149861Z","shell.execute_reply":"2024-03-15T16:57:53.165891Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del eegnet","metadata":{"execution":{"iopub.status.busy":"2024-03-15T16:57:53.168200Z","iopub.execute_input":"2024-03-15T16:57:53.168463Z","iopub.status.idle":"2024-03-15T16:57:53.177088Z","shell.execute_reply.started":"2024-03-15T16:57:53.168441Z","shell.execute_reply":"2024-03-15T16:57:53.176079Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# <span style='color:#B22222'>|</span> Ensemble submission <a class='anchor' id='submission'></a> [↑](#top) ","metadata":{}},{"cell_type":"code","source":"# Create the submission file and read in the the eeg_id\nsubmission = pd.DataFrame({'eeg_id': test_df.eeg_id.values})\n\n# Blend the vote chategories\ncolumns_to_average = ['seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']\n\n#Set weights\nWEIGHT_EEGNET = 0.18\nWEIGHT_WAVENET = 0.12\nWEIGHT_EFFICIENTNET = 0.7\n\n# Add the vote values from both Wavenet and EfficientNetB0\nfor i,column in enumerate(columns_to_average):\n    submission[column] = (WEIGHT_EFFICIENTNET*predictions_efficientnet[:,i] + WEIGHT_WAVENET*predictions_wavenet[:,i]+WEIGHT_EEGNET*predictions_eegnet[:,i]) # \n#     submission[column] = (0.55*predictions_efficientnet[:,i] + 0.45*predictions_wavenet[:,i]) # Do we want equal weighting or not??\n    \n# Santify check to confirm predictions sum to one\nsubmission.iloc[:,-6:].sum(axis=1)","metadata":{"execution":{"iopub.status.busy":"2024-03-15T16:57:53.178273Z","iopub.execute_input":"2024-03-15T16:57:53.178534Z","iopub.status.idle":"2024-03-15T16:57:53.195726Z","shell.execute_reply.started":"2024-03-15T16:57:53.178512Z","shell.execute_reply":"2024-03-15T16:57:53.194662Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission.to_csv('submission.csv',index=False)\nprint(f'Submission shape: {submission.shape}')\nsubmission.head()","metadata":{"execution":{"iopub.status.busy":"2024-03-15T16:57:53.196924Z","iopub.execute_input":"2024-03-15T16:57:53.197213Z","iopub.status.idle":"2024-03-15T16:57:53.216959Z","shell.execute_reply.started":"2024-03-15T16:57:53.197190Z","shell.execute_reply":"2024-03-15T16:57:53.216095Z"},"trusted":true},"execution_count":null,"outputs":[]}]}