{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"},{"sourceId":7923411,"sourceType":"datasetVersion","datasetId":4656488},{"sourceId":7932598,"sourceType":"datasetVersion","datasetId":4662851},{"sourceId":7951276,"sourceType":"datasetVersion","datasetId":4676075},{"sourceId":8030827,"sourceType":"datasetVersion","datasetId":4733500},{"sourceId":8030921,"sourceType":"datasetVersion","datasetId":4733568},{"sourceId":8030928,"sourceType":"datasetVersion","datasetId":4733574},{"sourceId":8032112,"sourceType":"datasetVersion","datasetId":4734441},{"sourceId":8039079,"sourceType":"datasetVersion","datasetId":4739493},{"sourceId":8039551,"sourceType":"datasetVersion","datasetId":4739820},{"sourceId":8057011,"sourceType":"datasetVersion","datasetId":4752154}],"dockerImageVersionId":30674,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport random\nimport numpy as np\nimport pandas as pd\nimport tqdm\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.nn import Linear\nfrom torch.utils.data import Dataset, DataLoader\nfrom fastai.vision.all import *\nfrom skimage.transform import resize\nimport matplotlib.pyplot as plt\nfrom sklearn.model_selection import KFold\nimport gc\nimport shutil\n\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-04-08T17:40:11.275422Z","iopub.execute_input":"2024-04-08T17:40:11.276081Z","iopub.status.idle":"2024-04-08T17:40:23.834266Z","shell.execute_reply.started":"2024-04-08T17:40:11.276041Z","shell.execute_reply":"2024-04-08T17:40:23.833447Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"PATH = {\n    'test':'/kaggle/input/hms-harmful-brain-activity-classification/test_',\n    'train':'/kaggle/input/hms-harmful-brain-activity-classification/train_'\n}\nBS = 512\nmodels = {\n    'eeg':[\n        \"/kaggle/input/eegs-cv-kaggle/EEGS_CNN_0\",\n        \"/kaggle/input/eegs-cv-kaggle/EEGS_CNN_1\",\n        \"/kaggle/input/eegs-cv-kaggle/EEGS_CNN_2\",\n        \"/kaggle/input/eegs-cv-kaggle/EEGS_CNN_3\",\n        \"/kaggle/input/eegs-cv-kaggle/EEGS_CNN_4\"\n    ],\n    'spec':[\n        \"/kaggle/input/specs-cv-kaggle/SPECS_CNN_0\",\n        \"/kaggle/input/specs-cv-kaggle/SPECS_CNN_1\",\n        \"/kaggle/input/specs-cv-kaggle/SPECS_CNN_2\",\n        \"/kaggle/input/specs-cv-kaggle/SPECS_CNN_3\",\n        \"/kaggle/input/specs-cv-kaggle/SPECS_CNN_4\"\n    ],\n    'spec_fixed':[\n        \"/kaggle/input/specs-cv-kaggle-fixed/SPECS_CNN_0\",\n        \"/kaggle/input/specs-cv-kaggle-fixed/SPECS_CNN_1\",\n        \"/kaggle/input/specs-cv-kaggle-fixed/SPECS_CNN_2\",\n        \"/kaggle/input/specs-cv-kaggle-fixed/SPECS_CNN_3\",\n        \"/kaggle/input/specs-cv-kaggle-fixed/SPECS_CNN_4\"\n    ],\n    'custom':[\n        \"/kaggle/input/custom-cv-kaggle/custom_SPECS_CNN_0\",\n        \"/kaggle/input/custom-cv-kaggle/custom_SPECS_CNN_1\",\n        \"/kaggle/input/custom-cv-kaggle/custom_SPECS_CNN_2\",\n        \"/kaggle/input/custom-cv-kaggle/custom_SPECS_CNN_3\",\n        \"/kaggle/input/custom-cv-kaggle/custom_SPECS_CNN_4\"\n    ],\n    'custom_fixed':[\n        \"/kaggle/input/custom-cv-kaggle-fixed/custom_SPECS_CNN_0\",\n        \"/kaggle/input/custom-cv-kaggle-fixed/custom_SPECS_CNN_1\",\n        \"/kaggle/input/custom-cv-kaggle-fixed/custom_SPECS_CNN_2\",\n        \"/kaggle/input/custom-cv-kaggle-fixed/custom_SPECS_CNN_3\",\n        \"/kaggle/input/custom-cv-kaggle-fixed/custom_SPECS_CNN_4\"\n    ],\n    'HMS':[\n        \"/kaggle/input/hms-xtra/HMSmodel_0\",\n        \"/kaggle/input/hms-xtra/HMSmodel_1\",\n        \"/kaggle/input/hms-xtra/HMSmodel_2\",\n        \"/kaggle/input/hms-xtra/HMSmodel_3\",\n        \"/kaggle/input/hms-xtra/HMSmodel_4\"\n    ]\n}\nDEBUG = False","metadata":{"execution":{"iopub.status.busy":"2024-04-08T17:40:23.836385Z","iopub.execute_input":"2024-04-08T17:40:23.837044Z","iopub.status.idle":"2024-04-08T17:40:23.844702Z","shell.execute_reply.started":"2024-04-08T17:40:23.837011Z","shell.execute_reply":"2024-04-08T17:40:23.843765Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Last Bullet Train\nTH = 10# Percentile of pseudolabels to take\npseudo_votes = 5# The votes assigned to pseudo_labels\nbatch_size = 128\nmin_votes = 0\nLR = 1e-3\nEPOCHS = 5\nstart_p  = .5\nend_p = 1\nstart_min_v = 1\nend_min_v = 1\nlabel_aug = False\nN_FOLDS = 5\nFOLDS = [0, 1, 2, 3, 4]\nHMS_DROP = .9\nCV = 'eeg_id'# Different eegs even of the same patient can be quite different.\nGB = 'votation'# The string TAG of sample normalized votates, 'expert_consensus' is another reasonable choice","metadata":{"execution":{"iopub.status.busy":"2024-04-08T17:40:23.845829Z","iopub.execute_input":"2024-04-08T17:40:23.846095Z","iopub.status.idle":"2024-04-08T17:40:23.862426Z","shell.execute_reply.started":"2024-04-08T17:40:23.846071Z","shell.execute_reply":"2024-04-08T17:40:23.861736Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def normalize(spec,epsilon=1e-6,NATURAL=False):\n    if NATURAL:\n#       https://www.kaggle.com/code/cdeotte/efficientnetb0-starter-lb-0-43?scriptVersionId=159911317\n        spec = np.clip(spec,np.exp(-4),np.exp(8))\n        spec = np.log(spec)\n    else:\n#       https://www.kaggle.com/code/rafaelzimmermann1/hms-spectrogram-creation-using-gpu\n        spec = np.clip(spec,np.exp(-4),np.exp(6))\n        spec = np.log10(spec)\n\n    mask = ~np.isnan(spec)\n    mean = np.mean(spec[mask])\n    std = np.std(spec[mask])\n    spec[mask] = spec[mask] - mean\n    if std > 0: spec[mask] /= (std + epsilon)\n    \n    return spec","metadata":{"execution":{"iopub.status.busy":"2024-04-08T17:40:23.864240Z","iopub.execute_input":"2024-04-08T17:40:23.864512Z","iopub.status.idle":"2024-04-08T17:40:23.876473Z","shell.execute_reply.started":"2024-04-08T17:40:23.864488Z","shell.execute_reply":"2024-04-08T17:40:23.875566Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from scipy.signal import butter, lfilter\ndef butter_lowpass_filter(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","metadata":{"execution":{"iopub.status.busy":"2024-04-08T17:40:23.877602Z","iopub.execute_input":"2024-04-08T17:40:23.877972Z","iopub.status.idle":"2024-04-08T17:40:23.936509Z","shell.execute_reply.started":"2024-04-08T17:40:23.877938Z","shell.execute_reply":"2024-04-08T17:40:23.935608Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class HMSmodel(nn.Module):\n\n    def __init__(self,models):\n        super().__init__()\n        self.models = torch.nn.ModuleList(models)\n        self.FC = nn.Linear(1280*len(models),6).to(device)\n\n        for model in self.models:\n            for p in model.parameters():\n                p.requires_grad = False\n        \n        weights = []\n        bias = 0\n        for model in self.models:\n            model.features[8][0].weight.requires_grad = True\n            model.classifier[0] = nn.Dropout(HMS_DROP)\n            weights.append(model.classifier[1].weight)\n            bias += model.classifier[1].bias\n            model.classifier[1] = nn.Identity()\n            \n        self.FC.weight = nn.Parameter(torch.cat(weights,1).to(device))\n        self.FC.bias = nn.Parameter(bias.to(device))\n\n    def forward(self,X):\n        eeg,spec,custom_spec = X\n        X = torch.cat([self.models[i](X[i]) for i in range(len(self.models))],1)\n        OUT = self.FC(X)\n        return OUT","metadata":{"execution":{"iopub.status.busy":"2024-04-08T17:40:23.939499Z","iopub.execute_input":"2024-04-08T17:40:23.939928Z","iopub.status.idle":"2024-04-08T17:40:23.949046Z","shell.execute_reply.started":"2024-04-08T17:40:23.939903Z","shell.execute_reply":"2024-04-08T17:40:23.948199Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**INFERENCE**","metadata":{}},{"cell_type":"code","source":"submission = pd.read_csv('/kaggle/input/hms-harmful-brain-activity-classification/sample_submission.csv')\nsubmission.head()","metadata":{"execution":{"iopub.status.busy":"2024-04-08T17:40:23.950004Z","iopub.execute_input":"2024-04-08T17:40:23.950500Z","iopub.status.idle":"2024-04-08T17:40:23.983885Z","shell.execute_reply.started":"2024-04-08T17:40:23.950475Z","shell.execute_reply":"2024-04-08T17:40:23.983033Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"votes = [c for c in submission.columns if '_vote' in c]\nvotes","metadata":{"execution":{"iopub.status.busy":"2024-04-08T17:40:23.985208Z","iopub.execute_input":"2024-04-08T17:40:23.985458Z","iopub.status.idle":"2024-04-08T17:40:23.991516Z","shell.execute_reply.started":"2024-04-08T17:40:23.985436Z","shell.execute_reply":"2024-04-08T17:40:23.990497Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test = pd.read_csv('/kaggle/input/hms-harmful-brain-activity-classification/test.csv')\ntest.head()","metadata":{"execution":{"iopub.status.busy":"2024-04-08T17:40:23.992758Z","iopub.execute_input":"2024-04-08T17:40:23.993025Z","iopub.status.idle":"2024-04-08T17:40:24.009069Z","shell.execute_reply.started":"2024-04-08T17:40:23.993003Z","shell.execute_reply":"2024-04-08T17:40:24.008407Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if DEBUG:\n    test = pd.read_csv('/kaggle/input/hms-harmful-brain-activity-classification/train.csv')[:1000]\n    PATH['test'] = PATH['train']\n    test.head()","metadata":{"execution":{"iopub.status.busy":"2024-04-08T17:40:24.012314Z","iopub.execute_input":"2024-04-08T17:40:24.012574Z","iopub.status.idle":"2024-04-08T17:40:24.245507Z","shell.execute_reply.started":"2024-04-08T17:40:24.012551Z","shell.execute_reply":"2024-04-08T17:40:24.244707Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"i = np.random.randint(len(test))\nrow = test[i:i+1]\nrow","metadata":{"execution":{"iopub.status.busy":"2024-04-08T17:40:24.246621Z","iopub.execute_input":"2024-04-08T17:40:24.246930Z","iopub.status.idle":"2024-04-08T17:40:24.261269Z","shell.execute_reply.started":"2024-04-08T17:40:24.246903Z","shell.execute_reply":"2024-04-08T17:40:24.260267Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"spec = pd.read_parquet(PATH['test']+'spectrograms/'+str(row.spectrogram_id.values[0])+'.parquet')\nLL = [c for c in spec.columns if 'LL' in c]\nRL = [c for c in spec.columns if 'RL' in c]\nLP = [c for c in spec.columns if 'LP' in c]\nRP = [c for c in spec.columns if 'RP' in c]","metadata":{"execution":{"iopub.status.busy":"2024-04-08T17:40:24.262700Z","iopub.execute_input":"2024-04-08T17:40:24.263114Z","iopub.status.idle":"2024-04-08T17:40:24.494938Z","shell.execute_reply.started":"2024-04-08T17:40:24.263081Z","shell.execute_reply":"2024-04-08T17:40:24.494093Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**CUSTOM**","metadata":{}},{"cell_type":"code","source":"# Sorry Rapids, not this time by https://www.kaggle.com/sergiosaharovskiy\n# https://www.kaggle.com/competitions/hms-harmful-brain-activity-classification/discussion/487110\n# HOTFIX for cupy.signal.filtfilt with bandpass_coefficients cannot find the reason...\n# It produces different results than scipy.signal.filtfilt with the same filter coefficients.\n# signal_filtered = filtfilt(*bandpass_coefficients, signal_filtered)\n\n\nimport cupy as cp\nimport numpy as np\nimport pandas as pd\nimport tqdm\nimport os\nfrom cupyx.scipy.ndimage import gaussian_filter\nfrom cupyx.scipy.signal import filtfilt, iirnotch\nfrom cupyx.scipy.signal import spectrogram as cupyx_spectrogram\nfrom scipy.signal import filtfilt as scipy_filtfilt, butter as scipy_butter\nfrom skimage.transform import rescale, resize, downscale_local_mean\n\ndef create_spectrogram_with_cupy(eeg_data, eeg_id,\n                                 low_cut_freq=0.7, high_cut_freq=20, order_band=5,\n                                 nperseg=1500, noverlap=1483, nfft=2750,\n                                 sigma_gaussian=0.7,\n                                 mean_montage_names=4):\n    electrode_pair_name_locations = {'LL': ['Fp1', 'F7', 'T3', 'T5', 'O1'],\n                                     'RL': ['Fp2', 'F8', 'T4', 'T6', 'O2'],\n                                     'LP': ['Fp1', 'F3', 'C3', 'P3', 'O1'],\n                                     'RP': ['Fp2', 'F4', 'C4', 'P4', 'O2']}\n\n    # Filter specifications\n    nyquist_freq = 0.5 * 200\n    low_cut_freq_normalized = low_cut_freq / nyquist_freq\n    high_cut_freq_normalized = high_cut_freq / nyquist_freq\n\n    # Bandpass and notch filter\n    # bandpass_coefficients = butter(order_band, [low_cut_freq_normalized, high_cut_freq_normalized], btype='band')\n    notch_coefficients = iirnotch(w0=60, Q=30, fs=200)\n    sci_bandpass_coefficients = scipy_butter(order_band, [low_cut_freq_normalized, high_cut_freq_normalized],\n                                             btype='band')\n    \n    spec_size = len(eeg_data)\n\n    # Spectrogram parameters\n    fs = 200\n\n    processed_eeg = {}\n\n    for i, (electrode_pair_name, electrode_locs) in enumerate(electrode_pair_name_locations.items()):\n        processed_eeg[electrode_pair_name] = np.zeros(spec_size)\n\n        for j in range(4):\n            # Compute differential signals\n            signal = cp.array(eeg_data[electrode_locs[j]].values - eeg_data[electrode_locs[j + 1]].values)\n\n            # Handles NaNs \n            mean_signal = cp.nanmean(signal)\n            signal = cp.nan_to_num(signal, nan=mean_signal) if cp.isnan(signal).mean() < 1 else cp.zeros_like(signal)\n\n            # Filters bandpass and notch\n            signal_filtered = filtfilt(*notch_coefficients, signal)\n            signal_filtered = scipy_filtfilt(*sci_bandpass_coefficients, signal_filtered.get())  # HOTFIX\n\n            # GPU-accelerated spectrogram computation\n            frequencies, times, Sxx = cupyx_spectrogram(signal_filtered, fs, nperseg=nperseg, noverlap=noverlap,\n                                                        nfft=nfft)\n            # Filters frequency range \n            valid_freq = (frequencies >= 0.59) & (frequencies <= 20)\n            Sxx_filtered = Sxx[valid_freq, :]\n\n            # Logarithmic transformation and normalization using Cupy\n            spectrogram_slice = cp.clip(Sxx_filtered, cp.exp(-4), cp.exp(6))\n            spectrogram_slice = cp.log10(spectrogram_slice)\n\n            normalization_epsilon = 1e-6\n            mean = spectrogram_slice.mean(axis=(0, 1), keepdims=True)\n            std = spectrogram_slice.std(axis=(0, 1), keepdims=True)\n            spectrogram_slice = (spectrogram_slice - mean) / (std + normalization_epsilon)\n\n            try:\n                spectrogram[:, :, i] += spectrogram_slice\n            except:\n#               Initialize spectrogram container\n                h,w = spectrogram_slice.shape\n                spectrogram = cp.zeros((h, w, 4), dtype='float32')\n                spectrogram[:, :, i] = spectrogram_slice\n            \n            processed_eeg[f'{electrode_locs[j]}_{electrode_locs[j + 1]}'] = signal.get()\n            processed_eeg[electrode_pair_name] += signal.get()\n\n        # AVERAGES THE 4 MONTAGE DIFFERENCES\n        if mean_montage_names > 0:\n            spectrogram[:, :, i] /= mean_montage_names\n\n    # Applies Gaussian filter and retrieves the spectrogram as a NumPy array using cupy.ndarray.get()\n    spec_numpy = gaussian_filter(spectrogram, sigma=sigma_gaussian).get() if sigma_gaussian > 0 else spectrogram.get()\n\n    # Filter EKG signal\n    ekg_signal_filtered = filtfilt(*notch_coefficients, cp.array(eeg_data[\"EKG\"].values))\n    processed_eeg['EKG'] = scipy_filtfilt(*sci_bandpass_coefficients, ekg_signal_filtered.get())  # HOTFIX\n    return spec_numpy, processed_eeg","metadata":{"execution":{"iopub.status.busy":"2024-04-08T17:40:24.496403Z","iopub.execute_input":"2024-04-08T17:40:24.496702Z","iopub.status.idle":"2024-04-08T17:40:26.371361Z","shell.execute_reply.started":"2024-04-08T17:40:24.496676Z","shell.execute_reply":"2024-04-08T17:40:26.370388Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.mkdir('/kaggle/working/custom')\nfor eeg_id in tqdm.tqdm(test.eeg_id.unique()):\n    eeg = pd.read_parquet(PATH['test']+'eegs/'+str(eeg_id)+'.parquet')\n    mask = np.isnan(eeg.values).max(1)\n    spec,_ = create_spectrogram_with_cupy(eeg, eeg_id,\n                                 low_cut_freq=0.7, high_cut_freq=20, order_band=5,\n                                 nperseg= 500, noverlap= 200,\n                                 nfft=1024,\n                                 sigma_gaussian=0.7,\n                                 mean_montage_names=4)\n    _,w,_ = spec.shape\n    mask = resize(mask,(w,1))\n    spec[:,mask[:,0],:] = 0\n    np.save('/kaggle/working/custom/'+str(eeg_id), spec)\n    del eeg,mask,spec","metadata":{"execution":{"iopub.status.busy":"2024-04-08T17:40:26.372794Z","iopub.execute_input":"2024-04-08T17:40:26.373453Z","iopub.status.idle":"2024-04-08T17:42:58.682968Z","shell.execute_reply.started":"2024-04-08T17:40:26.373418Z","shell.execute_reply":"2024-04-08T17:42:58.681709Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**INFERENCE**","metadata":{}},{"cell_type":"code","source":"class HMS_DS(torch.utils.data.Dataset):\n    '''\n    '''  \n    def __init__(self, df):\n        self.START = (10000 - 2048)//2\n        self.data = df\n        self.eeg = np.zeros((8,2048),dtype=np.float32)\n        self.spec = np.zeros((1,256,256))\n   \n    def __len__(self):\n        return len(self.data)\n        \n    def __getitem__(self, idx):\n#       https://www.kaggle.com/competitions/hms-harmful-brain-activity-classification/discussion/468010\n        row = self.data.iloc[idx]\n#==============================================================================================\n#       EEG\n#==============================================================================================\n        eeg = pd.read_parquet(PATH['test']+'eegs/'+str(row.eeg_id)+'.parquet')\n        eeg = eeg.iloc[self.START:self.START+2048]\n    \n        mask = ~np.isnan(eeg['Fp1'])\n        self.eeg[:,:] = 0\n        \n        self.eeg[0][mask] = eeg['Fp1'][mask] - eeg['T3'][mask]\n        self.eeg[1][mask] = eeg['T3'][mask] - eeg['O1'][mask]\n\n        self.eeg[2][mask] = eeg['Fp1'][mask] - eeg['C3'][mask]\n        self.eeg[3][mask] = eeg['C3'][mask] - eeg['O1'][mask]\n\n        self.eeg[4][mask] = eeg['Fp2'][mask] - eeg['C4'][mask]\n        self.eeg[5][mask] = eeg['C4'][mask] - eeg['O2'][mask]\n\n        self.eeg[6][mask] = eeg['Fp2'][mask] - eeg['T4'][mask]\n        self.eeg[7][mask] = eeg['T4'][mask] - eeg['O2'][mask]\n\n#       === Standarize ===\n        self.eeg = np.clip(self.eeg,-1024, 1024)/32.0\n\n#       === Butter Low-pass Filter ===\n        self.eeg = butter_lowpass_filter(self.eeg)\n\n        eeg = torch.from_numpy(self.eeg).float().to(device)\n#==============================================================================================\n#       SPEC\n#==============================================================================================\n        spec = pd.read_parquet(PATH['test']+'spectrograms/'+str(row.spectrogram_id)+'.parquet')\n        \n        spec = spec[LL + RL + LP + RP]\n        \n        t = 22\n        spec = spec[LL + RL + LP + RP][t:t+256]\n        \n        self.spec[0,:,:64] = normalize(resize(spec[LP].values,(256,64)))\n        self.spec[0,:,64:128] = normalize(resize(spec[LL].values,(256,64)))\n        self.spec[0,:,128:-64] = normalize(resize(spec[RP].values,(256,64)))\n        self.spec[0,:,-64:] = normalize(resize(spec[RL].values,(256,64)))\n        \n        mask = np.isnan(self.spec)\n        self.spec[mask] = 0\n\n        spec = torch.from_numpy(self.spec).float().to(device)\n#==============================================================================================\n#       CUSTOM SPECS\n#==============================================================================================\n        custom_spec = np.load('/kaggle/working/custom/'+str(row.eeg_id)+'.npy')\n        custom_spec = np.concatenate((custom_spec[:,:,3],\n                                      custom_spec[:,:,2],\n                                      custom_spec[:,:,1],\n                                      custom_spec[:,:,0]))\n        custom_spec = resize(custom_spec,(256,300))\n        \n        t = 22\n        self.spec[0] = custom_spec[:,t:t+256]\n                    \n        custom_spec = torch.from_numpy(self.spec).float().to(device)\n\n        return int(row.eeg_id),row.spectrogram_id,eeg,spec,custom_spec","metadata":{"execution":{"iopub.status.busy":"2024-04-08T17:42:58.689235Z","iopub.execute_input":"2024-04-08T17:42:58.689916Z","iopub.status.idle":"2024-04-08T17:42:58.731978Z","shell.execute_reply.started":"2024-04-08T17:42:58.689870Z","shell.execute_reply":"2024-04-08T17:42:58.730767Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**First Prediction**","metadata":{}},{"cell_type":"code","source":"if len(test) > 1:\n    submission = pd.DataFrame()\n    \n    ds = HMS_DS(test)\n    dl = DataLoader(\n        ds,\n        batch_size=BS\n    )\n    \n    eeg_ensemble = []\n    for model in models['eeg']:\n        eeg_ensemble.append(torch.load(model).eval())\n\n    spec_ensemble = []\n    for model in models['spec']:\n        spec_ensemble.append(torch.load(model).eval())\n        \n    custom_ensemble = []\n    for model in models['custom']:\n        custom_ensemble.append(torch.load(model).eval())\n        \n    HMS_ensemble = []\n    for model in models['HMS']:\n        HMS_ensemble.append(\n            torch.load(model).eval()\n        )\n        \n    with torch.no_grad():\n        PREDS = torch.zeros((BS,20,6),device=device)\n        for eeg_id,spectrogram_id,eegs,specs,custom_specs in dl:\n            N = len(eegs)\n            \n            i = 0\n            for model in eeg_ensemble:\n                PREDS[:N,i] = torch.softmax(model(eegs),-1)\n                i += 1\n\n            for model in spec_ensemble:\n                PREDS[:N,i] = torch.softmax(model(specs),-1)\n                i += 1\n\n            for model in custom_ensemble:\n                PREDS[:N,i] = torch.softmax(model(custom_specs),-1)\n                i += 1\n               \n            for model in HMS_ensemble:\n                PREDS[:N,i] = torch.softmax(model([eegs,specs,custom_specs]),-1)\n                i += 1\n                \n            mu = torch.softmax(PREDS[:N].mean(1),-1).cpu()\n            sigma = PREDS[:N].std(1).cpu()\n            sigma = (mu*sigma).sum(-1)/(mu.sum(-1))\n            \n            df = {'eeg_id':eeg_id,'spectrogram_id':spectrogram_id}\n            for i in range(6):\n                df[votes[i]] = mu[:,i].cpu()\n                \n            df['sigma'] = sigma\n            \n            submission = pd.concat([submission,pd.DataFrame(df)],ignore_index=True)\n            \n    TH = np.percentile(submission['sigma'],TH)\n#   The predictions with sigma under threshold will be taken as confident pseudolabels\n#   with a total weigth of pseudo_votes votes\n    pseudo_labels = submission[submission['sigma'] < TH].copy()\n    pseudo_labels[votes] *= pseudo_votes#10\n    print(pseudo_labels.head())\n\nprint(submission.head())","metadata":{"execution":{"iopub.status.busy":"2024-04-08T17:42:58.734270Z","iopub.execute_input":"2024-04-08T17:42:58.735128Z","iopub.status.idle":"2024-04-08T17:44:35.450487Z","shell.execute_reply.started":"2024-04-08T17:42:58.735092Z","shell.execute_reply":"2024-04-08T17:44:35.449463Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**LAST BULLET TRAIN**","metadata":{}},{"cell_type":"code","source":"df = pd.read_csv('/kaggle/input/hms-harmful-brain-activity-classification/train.csv')\ndf['origin'] = 'train'\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2024-04-08T17:44:35.475857Z","iopub.execute_input":"2024-04-08T17:44:35.476125Z","iopub.status.idle":"2024-04-08T17:44:35.670231Z","shell.execute_reply.started":"2024-04-08T17:44:35.476101Z","shell.execute_reply":"2024-04-08T17:44:35.669203Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seed_everything(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\n\nseed_everything(2024)","metadata":{"execution":{"iopub.status.busy":"2024-04-08T17:44:35.671472Z","iopub.execute_input":"2024-04-08T17:44:35.671809Z","iopub.status.idle":"2024-04-08T17:44:35.677955Z","shell.execute_reply.started":"2024-04-08T17:44:35.671780Z","shell.execute_reply":"2024-04-08T17:44:35.676943Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"m = .55# Define the majority criteria\ndf.expert_consensus = 'T'\nS = \"df.loc[(df.seizure_vote > m*(df.seizure_vote + df.lpd_vote))*(df.seizure_vote > m*(df.seizure_vote + df.gpd_vote))*(df.seizure_vote > m*(df.seizure_vote + df.lrda_vote))*(df.seizure_vote > m*(df.seizure_vote + df.grda_vote))*(df.seizure_vote > m*(df.seizure_vote + df.other_vote)),'expert_consensus'] = 'seizure'\"\n\nexec(S)\nfor a in ['lpd','gpd','lrda','grda','other']:\n    SWAP = S.replace(a,'SWAP')\n    SWAP = SWAP.replace('seizure',a)\n    exec(SWAP.replace('SWAP','seizure'))\n\nconsensus = ['seizure', 'lpd', 'gpd', 'lrda', 'grda', 'other', 'T']\nfor c in consensus:\n    c_df = df[['eeg_id']+votes].loc[df.expert_consensus == c].groupby('eeg_id').mean()\n    values,counts = np.unique(c_df[votes].sum(1).values,return_counts=True)\n    plt.title(c + ': ' + str(len(c_df)))\n    plt.bar(values,counts)\n    plt.show()\n    print(c_df[votes].head())","metadata":{"execution":{"iopub.status.busy":"2024-04-08T17:44:35.679086Z","iopub.execute_input":"2024-04-08T17:44:35.679378Z","iopub.status.idle":"2024-04-08T17:44:39.167360Z","shell.execute_reply.started":"2024-04-08T17:44:35.679350Z","shell.execute_reply":"2024-04-08T17:44:39.166339Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if len(test) > 1:\n    pseudo_labels['expert_consensus'] = 'T'\n    S = S.replace('df','pseudo_labels')\n\n    exec(S)\n    for a in ['lpd','gpd','lrda','grda','other']:\n        SWAP = S.replace(a,'SWAP')\n        SWAP = SWAP.replace('seizure',a)\n        exec(SWAP.replace('SWAP','seizure'))\n\n    for c in consensus:\n        c_df = pseudo_labels[['eeg_id']+votes].loc[pseudo_labels['expert_consensus'] == c].groupby('eeg_id').mean()\n        print(c,':',c_df)\n        \n    pseudo_labels['origin'] = 'test'\n    pseudo_labels[['spectrogram_label_offset_seconds','eeg_label_offset_seconds']] = 0\n    print(pseudo_labels.head())","metadata":{"execution":{"iopub.status.busy":"2024-04-08T17:44:39.168758Z","iopub.execute_input":"2024-04-08T17:44:39.169065Z","iopub.status.idle":"2024-04-08T17:44:39.223237Z","shell.execute_reply.started":"2024-04-08T17:44:39.169038Z","shell.execute_reply":"2024-04-08T17:44:39.222228Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class HMS_TRAIN_DS(torch.utils.data.Dataset):\n    '''\n    '''  \n    def __init__(self, df, W = 1024, VALID = False):\n        self.W = W\n        self.VALID = VALID\n        self.START = (10000 - 2048)//2\n        self.data = np.array(df[['spectrogram_id',\n                                 'eeg_id','expert_consensus',\n                                 'spectrogram_label_offset_seconds',\n                                 'eeg_label_offset_seconds','origin']+votes].groupby(['eeg_id','expert_consensus']), dtype=object)\n        if VALID:\n            self.eeg = np.zeros((8,2048),dtype=np.float32)\n        else:\n            self.eeg = np.zeros((8,1024),dtype=np.float32)\n        self.spec = np.zeros((1,256,256))\n   \n    def __len__(self):\n        return len(self.data)\n        \n    def __getitem__(self, idx):\n#       https://www.kaggle.com/competitions/hms-harmful-brain-activity-classification/discussion/468010\n        df = self.data[idx][1]\n        if self.VALID:\n            row = df.iloc[len(df)//2]\n        else:\n            row = df.iloc[np.random.randint(len(df))]\n#==============================================================================================\n#       EEG\n#==============================================================================================\n#       eeg = pd.read_parquet(PATH['train']+'eegs/'+str(row.eeg_id)+'.parquet')\n        eeg = eegs[row.origin][row.eeg_id]\n        START = 200*int(row.eeg_label_offset_seconds) + self.START \n        if not self.VALID:\n            if self.W < 2048:\n                START += np.random.randint(2048 - self.W)\n                eeg = eeg.iloc[START:START+self.W]         \n        else:\n            eeg = eeg.iloc[START:START+2048]\n    \n        mask = ~np.isnan(eeg['Fp1'])\n        self.eeg[:,:] = 0\n        \n        self.eeg[0][mask] = eeg['Fp1'][mask] - eeg['T3'][mask]\n        self.eeg[1][mask] = eeg['T3'][mask] - eeg['O1'][mask]\n\n        self.eeg[2][mask] = eeg['Fp1'][mask] - eeg['C3'][mask]\n        self.eeg[3][mask] = eeg['C3'][mask] - eeg['O1'][mask]\n\n        self.eeg[4][mask] = eeg['Fp2'][mask] - eeg['C4'][mask]\n        self.eeg[5][mask] = eeg['C4'][mask] - eeg['O2'][mask]\n\n        self.eeg[6][mask] = eeg['Fp2'][mask] - eeg['T4'][mask]\n        self.eeg[7][mask] = eeg['T4'][mask] - eeg['O2'][mask]\n\n#       === Standarize ===\n        self.eeg = np.clip(self.eeg,-1024, 1024)/32.0\n\n#       === Butter Low-pass Filter ===\n        self.eeg = butter_lowpass_filter(self.eeg)\n\n        eeg = torch.from_numpy(self.eeg).float().to(device)\n#==============================================================================================\n#       SPEC\n#==============================================================================================\n        spec = pd.read_parquet(PATH[row['origin']]+'spectrograms/'+str(row.spectrogram_id)+'.parquet')\n\n        spec_offset = int( row.spectrogram_label_offset_seconds )\n        spec = spec[LL + RL + LP + RP].loc[(spec.time>=spec_offset)\n                     &(spec.time<spec_offset+600)]\n        \n        t = 22\n        if not self.VALID:\n#           Decentralize the event\n            t += int(np.clip(np.random.normal(0,11),-22,22))\n        spec = spec[LL + RL + LP + RP][t:t+256]\n        \n        self.spec[0,:,:64] = normalize(resize(spec[LP].values,(256,64)))\n        self.spec[0,:,64:128] = normalize(resize(spec[LL].values,(256,64)))\n        self.spec[0,:,128:-64] = normalize(resize(spec[RP].values,(256,64)))\n        self.spec[0,:,-64:] = normalize(resize(spec[RL].values,(256,64)))\n        \n        mask = np.isnan(self.spec)\n        if not self.VALID:\n#           Contrast augmentation\n            self.spec *= np.random.normal(1,.0001)\n#           Brightness augmentation\n            self.spec += np.random.normal(0,.0001)\n        self.spec[mask] = 0\n        if not self.VALID:\n            if np.random.rand(1) > mask.sum()/(2*256*256):\n                w = int(np.clip(np.random.normal(256/5,256/20),0,512/5))\n                t = (256 - w)/2\n                t = int(np.clip(np.random.normal(0,t/2),-2*t,2*t))\n            \n                if t < 0:\n                    self.spec[:,t-w:t] = 0\n                else:\n                    self.spec[:,t:t+w] = 0\n\n        spec = torch.from_numpy(self.spec).float().to(device)\n#==============================================================================================\n#       CUSTOM SPECS\n#==============================================================================================\n        if row['origin'] == 'train':\n            custom_spec = np.load('/kaggle/input/custom-specs/custom/'+str(row.eeg_id)+'.npy')\n        else:\n            custom_spec = np.load('/kaggle/working/custom/'+str(row.eeg_id)+'.npy')\n        spec_offset = int( 2*row.eeg_label_offset_seconds/3 )\n        custom_spec = custom_spec[:,spec_offset:spec_offset+32,:]\n        custom_spec = np.concatenate((custom_spec[:,:,3],\n                                      custom_spec[:,:,2],\n                                      custom_spec[:,:,1],\n                                      custom_spec[:,:,0]))\n        custom_spec = resize(custom_spec,(256,300))\n        \n        t = 22\n        if not self.VALID:\n#           Decentralize the event\n            t += int(np.clip(np.random.normal(0,11),-22,22))\n        \n        self.spec[0] = custom_spec[:,t:t+256]\n\n        mask = np.isnan(self.spec)\n        if not self.VALID:\n#           Contrast augmentation\n            self.spec *= np.random.normal(1,.0001)\n#           Brightness augmentation\n            self.spec += np.random.normal(0,.0001)\n        self.spec[mask] = 0\n        if not self.VALID:\n            if np.random.rand(1) > mask.sum()/(2*256*256):\n                w = int(np.clip(np.random.normal(256/5,256/20),0,512/5))\n                t = (256 - w)/2\n                t = int(np.clip(np.random.normal(0,t/2),-2*t,2*t))\n            \n                if t < 0:\n                    self.spec[:,:,t-w:t] = 0\n                else:\n                    self.spec[:,:,t:t+w] = 0\n                    \n        custom_spec = torch.from_numpy(self.spec).float().to(device)\n#==============================================================================================\n#       LABELS\n#==============================================================================================\n        labels = row[votes].values\n        if label_aug and not self.VALID:\n            v = [0]*int(labels[0])+[1]*int(labels[1])+[2]*int(labels[2])+[3]*int(labels[3])+[4]*int(labels[4])+[5]*int(labels[5])\n            v = np.array(random.sample(v, max([1,len(v) - int(abs(np.random.normal(0,len(v)/8)))])))\n            labels[0] = np.sum(v == 0)\n            labels[1] = np.sum(v == 1)\n            labels[2] = np.sum(v == 2)\n            labels[3] = np.sum(v == 3)\n            labels[4] = np.sum(v == 4)\n            labels[5] = np.sum(v == 5)\n        \n        labels = torch.from_numpy(labels.astype(np.float32)).to(device)\n\n        return [eeg,spec,custom_spec],labels","metadata":{"execution":{"iopub.status.busy":"2024-04-08T17:44:39.244821Z","iopub.execute_input":"2024-04-08T17:44:39.245158Z","iopub.status.idle":"2024-04-08T17:44:39.285185Z","shell.execute_reply.started":"2024-04-08T17:44:39.245131Z","shell.execute_reply":"2024-04-08T17:44:39.284229Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SemisupervisedKLDiv(nn.KLDivLoss):\n#   Labels with less total votes are less confident\n#   Smoothing those labels with semisupervised votes\n#   will reduce his loss so gradient will be more\n#   affected by samples with higher total votes\n    def __init__(\n            self,\n            min_v=1# The minium votes to be accepted\n        ):\n        super().__init__(reduce=False)\n        self.min_v = torch.tensor(start_min_v,dtype=torch.float).to(device)\n        self.p = torch.tensor(start_p,dtype=torch.float).to(device)\n        self.step = 0\n\n    def __steps__(self,steps):\n        self.steps = steps\n\n    def set_min_v(self):\n        self.step += 1\n        if self.step <= self.steps: self.min_v = nt(start_min_v,end_min_v,self.step,self.steps)\n\n    def set_p(self):\n        self.step += 1\n        if self.step <= self.steps: self.p = nt(start_p,end_p,self.step,self.steps)\n\n    def forward(\n            self,\n            y,# Raw predictions\n            t # Raw targets\n        ):\n    #   How many votes do we have?\n        v = t.sum(-1,keepdim=True)\n    #   Where votes < min_v add model votes\n        mask = v[:,0] < self.min_v\n        t[mask] += (self.min_v - v[mask])*(self.p*torch.softmax(y[mask].detach(),-1) + (1 - self.p)*torch.softmax(t[mask],-1))\n        v[mask] = self.min_v\n        t /= v\n        y = nn.functional.log_softmax(y,  dim=1)\n        loss = super().forward(y, t).sum(-1,keepdim=True)\n        loss = loss*v\n        loss = loss.sum()/v.sum()\n        return loss","metadata":{"execution":{"iopub.status.busy":"2024-04-08T17:44:39.286318Z","iopub.execute_input":"2024-04-08T17:44:39.286605Z","iopub.status.idle":"2024-04-08T17:44:39.299670Z","shell.execute_reply.started":"2024-04-08T17:44:39.286581Z","shell.execute_reply":"2024-04-08T17:44:39.298768Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# CosineAnnealingVotes\ndef nt(nmin,nmax,tcur,tmax):\n    return nmax - .5*(nmax-nmin)*(1+np.cos(tcur*np.pi/tmax))","metadata":{"execution":{"iopub.status.busy":"2024-04-08T17:44:39.300700Z","iopub.execute_input":"2024-04-08T17:44:39.301001Z","iopub.status.idle":"2024-04-08T17:44:39.313125Z","shell.execute_reply.started":"2024-04-08T17:44:39.300978Z","shell.execute_reply":"2024-04-08T17:44:39.312181Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def cb(self):\n    learn.loss_func.set_min_v()\n    learn.loss_func.set_p()\nmin_v_cb = Callback(before_step=cb)","metadata":{"execution":{"iopub.status.busy":"2024-04-08T17:44:39.314173Z","iopub.execute_input":"2024-04-08T17:44:39.314400Z","iopub.status.idle":"2024-04-08T17:44:39.322688Z","shell.execute_reply.started":"2024-04-08T17:44:39.314379Z","shell.execute_reply":"2024-04-08T17:44:39.321846Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"seed_everything(2024)\nsplits = {}\nfor c in consensus:\n    ID = df.loc[df.expert_consensus == c][CV].unique()\n    splits[c] = []\n    for t,v in KFold(N_FOLDS).split(ID):\n        splits[c].append([ID[t],ID[v]])","metadata":{"execution":{"iopub.status.busy":"2024-04-08T17:44:39.323665Z","iopub.execute_input":"2024-04-08T17:44:39.323924Z","iopub.status.idle":"2024-04-08T17:44:39.467644Z","shell.execute_reply.started":"2024-04-08T17:44:39.323902Z","shell.execute_reply":"2024-04-08T17:44:39.466768Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"seed_everything(2024)\npseudo_labels_splits = {}\nfor c in consensus:\n    ID = pseudo_labels.loc[pseudo_labels['expert_consensus'] == c][CV].unique()\n    if len(ID) > N_FOLDS:\n        pseudo_labels_splits[c] = []\n        for t,v in KFold(N_FOLDS).split(ID):\n            pseudo_labels_splits[c].append([ID[t],ID[v]])","metadata":{"execution":{"iopub.status.busy":"2024-04-08T17:44:39.474007Z","iopub.execute_input":"2024-04-08T17:44:39.474318Z","iopub.status.idle":"2024-04-08T17:44:39.488284Z","shell.execute_reply.started":"2024-04-08T17:44:39.474291Z","shell.execute_reply":"2024-04-08T17:44:39.487014Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if len(test) > 1:\n    eegs = {'train':{}, 'test':{}}\n    for eeg_id in tqdm.tqdm(df[df[votes].sum(1) >= 10].eeg_id.unique()):\n        eegs['train'][eeg_id] = pd.read_parquet(PATH['train']+'eegs/'+str(eeg_id)+'.parquet')\n    \n    for eeg_id in tqdm.tqdm(pseudo_labels.eeg_id.unique()):\n        eegs['test'][eeg_id] = pd.read_parquet(PATH['test']+'eegs/'+str(eeg_id)+'.parquet')\n    \n    for f in FOLDS:\n        seed_everything(2024)\n        print('FOLD: ',f)\n        t = pd.DataFrame()\n        v = pd.DataFrame()\n        for c in splits:\n            c_df = df.loc[df.expert_consensus == c]\n            t = pd.concat([t,c_df[c_df[CV].isin(splits[c][f][0])]])\n            v = pd.concat([v,c_df[c_df[CV].isin(splits[c][f][1])]])\n#       2 stage training            \n        t = t[t[votes].sum(1) >= 10]\n        v = v[v[votes].sum(1) >= 10]\n        \n        for c in pseudo_labels_splits:\n            c_df = pseudo_labels.loc[pseudo_labels['expert_consensus'] == c]\n            t = pd.concat([t,c_df[c_df[CV].isin(pseudo_labels_splits[c][f][0])]])\n            v = pd.concat([v,c_df[c_df[CV].isin(pseudo_labels_splits[c][f][1])]])\n\n        tds = HMS_TRAIN_DS(t)\n        vds = HMS_TRAIN_DS(v,VALID=True)\n\n        tdl = DataLoader(\n            tds,\n            batch_size=batch_size,\n            shuffle=True,\n            drop_last=True\n        )\n        vdl = DataLoader(\n            vds,\n            batch_size=2*batch_size\n        )\n        dls = DataLoaders(tdl,vdl)\n\n        model = torch.load(models['HMS'][f])\n\n        learn = Learner(\n            dls,\n            model,\n            lr=LR,\n            loss_func=SemisupervisedKLDiv(),\n            cbs=[\n                GradientClip(3.0),\n                SaveModelCallback(),\n                ShowGraphCallback(),\n                min_v_cb\n            ]\n        )\n        learn.loss_func.__steps__((learn.dls.dataset.__len__()//batch_size)*EPOCHS)\n        learn.fit_one_cycle(EPOCHS)\n\n        torch.save(model,'HMSmodel_'+str(f))\n        del tdl,vdl,tds,vds,dls,learn,model\n        gc.collect()","metadata":{"execution":{"iopub.status.busy":"2024-04-08T17:44:39.489891Z","iopub.execute_input":"2024-04-08T17:44:39.490240Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Final Prediction**","metadata":{}},{"cell_type":"code","source":"if len(test) > 1:\n    del submission,PREDS\n    submission = pd.DataFrame()\n    \n    ds = HMS_DS(test)\n    dl = DataLoader(\n        ds,\n        batch_size=BS\n    )\n    \n    Last_Bullet = [\n        '/kaggle/working/HMSmodel_0',\n        '/kaggle/working/HMSmodel_1',\n        '/kaggle/working/HMSmodel_2',\n        '/kaggle/working/HMSmodel_3',\n        '/kaggle/working/HMSmodel_4'\n    ]\n    HMS_ensemble = []\n    for model in Last_Bullet:\n        HMS_ensemble.append(\n            torch.load(model).eval()\n        )\n        \n    with torch.no_grad():\n        PREDS = torch.zeros((BS,6),device=device)\n        for eeg_id,spectrogram_id,eegs,specs,custom_specs in dl:\n            N = len(eegs)\n            PREDS[:N,:] = 0 \n            for model in HMS_ensemble:\n                PREDS[:N] += torch.softmax(model([eegs,specs,custom_specs]),-1)\n                \n            PREDS[:N] /= PREDS[:N].sum(-1,keepdim=True)\n            \n            df = {'eeg_id':eeg_id,'spectrogram_id':spectrogram_id}\n            for i in range(6):\n                df[votes[i]] = PREDS[:N,i].cpu()\n            \n            submission = pd.concat([submission,pd.DataFrame(df)],ignore_index=True)\n            \n    submission = submission[['eeg_id']+votes]\n\nprint(submission.head())\n\nsubmission.to_csv('submission.csv', index=False)\nshutil.rmtree('/kaggle/working/custom')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission[votes].sum(1).unique()","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}