{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"},{"sourceId":7925878,"sourceType":"datasetVersion","datasetId":4585473},{"sourceId":8013715,"sourceType":"datasetVersion","datasetId":4523844},{"sourceId":8023674,"sourceType":"datasetVersion","datasetId":4689474},{"sourceId":8046306,"sourceType":"datasetVersion","datasetId":4659625}],"dockerImageVersionId":30636,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport torch \nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchvision.transforms as transforms\n\nimport random\nimport warnings\nimport os\nwarnings.filterwarnings('ignore')","metadata":{"execution":{"iopub.status.busy":"2024-01-29T00:06:03.769319Z","iopub.execute_input":"2024-01-29T00:06:03.769702Z","iopub.status.idle":"2024-01-29T00:06:03.775508Z","shell.execute_reply.started":"2024-01-29T00:06:03.769677Z","shell.execute_reply":"2024-01-29T00:06:03.774526Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DEBUG=False","metadata":{"execution":{"iopub.status.busy":"2024-01-29T00:06:03.777026Z","iopub.execute_input":"2024-01-29T00:06:03.777315Z","iopub.status.idle":"2024-01-29T00:06:03.792518Z","shell.execute_reply.started":"2024-01-29T00:06:03.777274Z","shell.execute_reply":"2024-01-29T00:06:03.791506Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"NAMES = ['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']]\nimport librosa","metadata":{"execution":{"iopub.status.busy":"2024-01-29T00:06:03.793969Z","iopub.execute_input":"2024-01-29T00:06:03.794288Z","iopub.status.idle":"2024-01-29T00:06:03.805797Z","shell.execute_reply.started":"2024-01-29T00:06:03.794256Z","shell.execute_reply":"2024-01-29T00:06:03.80488Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from scipy.signal import butter, lfilter\n","metadata":{"execution":{"iopub.status.busy":"2024-01-29T00:06:03.806936Z","iopub.execute_input":"2024-01-29T00:06:03.80793Z","iopub.status.idle":"2024-01-29T00:06:03.823273Z","shell.execute_reply.started":"2024-01-29T00:06:03.807895Z","shell.execute_reply":"2024-01-29T00:06:03.822436Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def stft_spec_from_eeg(parquet_path):\n    EEG_LENGTH = 50\n    eeg = pd.read_parquet(parquet_path)\n    \n    time_temp = 0\n    time_start = round(time_temp + (50 - EEG_LENGTH) / 2 * 200) \n    time_stop = round(time_temp + (50 + EEG_LENGTH) / 2 * 200)\n    \n    eeg = eeg.iloc[time_start: time_stop]\n    \n    list_eeg = list()\n    for k in range(4):\n        COLS = FEATS[k]\n        img = np.zeros((128,142,4),dtype='float32')\n        for kk in range(4):\n            eeg_1 = eeg[COLS[kk]]\n            mean_value = eeg_1.mean()\n            eeg_1.fillna(value=mean_value, inplace=True)\n            eeg_1 = eeg_1.values\n            \n            eeg_2 = eeg[COLS[kk+1]]\n            mean_value = eeg_2.mean()\n            eeg_2.fillna(value=mean_value, inplace=True)\n            eeg_2 = eeg_2.values\n            \n            new_eeg = eeg_1 - eeg_2\n            del eeg_1\n            del eeg_2\n            # new_eeg = signal.filtfilt(b, a, new_eeg, axis=0)\n            fs = 200  \n            nperseg = 70\n            noverlap = 0\n            f, t, spec = signal.spectrogram(new_eeg, fs, nperseg=nperseg, noverlap=noverlap, nfft=256)\n            \n            spec = np.abs(spec) \n            spec = np.log1p(spec).astype(\"float32\")\n\n            img[:,:,kk] += spec[:128, :]\n        img = np.concatenate((img[:,:,0], img[:,:,1], img[:,:,2], img[:,:,3]), 1)\n        list_eeg.append(img)\n    img = np.concatenate(list_eeg, 0) \n    img /= 2.0\n    return img","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"NAMES = ['LL','LP','RP','RR']\n\nSFREQ = 200\n\nfilter_range = [0.5, 40]\n\nRAW_FEATS = {'LL': ['Fp1-F7', 'F7-T3', 'T3-T5', 'T5-O1'],\n         'RL': ['Fp2-F8', 'F8-T4', 'T4-T6', 'T6-O2'],\n         'LP': ['Fp1-F3', 'F3-C3', 'C3-P3', 'P3-O1'],\n         'RP': ['Fp2-F4', 'F4-C4', 'C4-P4', 'P4-O2']}\n    \nfrom scipy import signal\n\nb, a = signal.butter(3, np.float32(filter_range)*2/SFREQ, 'bandpass')\n\ndef raw10seeg_from_eeg(parquet_path, eeg_id):\n    EEG_LENGTH = 10\n    raw_eeg = pd.read_parquet(parquet_path)\n    \n    time_temp = 0\n    time_start =  round(time_temp + (50 - EEG_LENGTH) / 2 * 200) \n    time_stop =  round(time_temp + (50 + EEG_LENGTH) / 2 * 200)\n\n    eeg_default = raw_eeg.loc[time_start: (time_stop - 1), :].reset_index(drop=True)\n    list_eeg = list()\n    for region in RAW_FEATS.keys():\n        eeg = np.zeros((len(RAW_FEATS[region]), eeg_default.shape[0]), dtype=np.float32)\n        for chan_i, chan in enumerate(RAW_FEATS[region]):\n            eeg_1 = eeg_default.loc[:, chan.split('-')[0]]\n            mean_value = eeg_1.mean()\n            eeg_1.fillna(value=mean_value, inplace=True)\n            eeg_1 = eeg_1.values\n\n            eeg_2 = eeg_default.loc[:, chan.split('-')[1]]\n            mean_value = eeg_2.mean()\n            eeg_2.fillna(value=mean_value, inplace=True)\n            eeg_2 = eeg_2.values\n\n            new_eeg = eeg_1 - eeg_2\n            new_eeg = signal.filtfilt(b, a, new_eeg, axis=0)\n            new_eeg = np.clip(new_eeg, -1024, 1024)\n            eeg[chan_i, :] = new_eeg\n\n        eeg = np.reshape(eeg, (4, 200, EEG_LENGTH))\n        eeg = np.concatenate([eeg[0,:,:], eeg[1,:,:], eeg[2,:,:], eeg[3,:,:]],1)\n        list_eeg.append(eeg)\n            \n    eeg_c = np.concatenate(list_eeg, 1)\n    eeg_c /= 104\n        \n    time_temp = 0\n    time_start =  round(time_temp + 18 * 200) \n    time_stop =  round(time_temp + 28 * 200)\n\n    eeg_default = raw_eeg.loc[time_start: (time_stop - 1), :].reset_index(drop=True)\n    list_eeg = list()\n    for region in RAW_FEATS.keys():\n        eeg = np.zeros((len(RAW_FEATS[region]), eeg_default.shape[0]), dtype=np.float32)\n        for chan_i, chan in enumerate(RAW_FEATS[region]):\n            eeg_1 = eeg_default.loc[:, chan.split('-')[0]]\n            mean_value = eeg_1.mean()\n            eeg_1.fillna(value=mean_value, inplace=True)\n            eeg_1 = eeg_1.values\n\n            eeg_2 = eeg_default.loc[:, chan.split('-')[1]]\n            mean_value = eeg_2.mean()\n            eeg_2.fillna(value=mean_value, inplace=True)\n            eeg_2 = eeg_2.values\n\n            new_eeg = eeg_1 - eeg_2\n            new_eeg = signal.filtfilt(b, a, new_eeg, axis=0)\n            new_eeg = np.clip(new_eeg, -1024, 1024)\n            eeg[chan_i, :] = new_eeg\n\n        eeg = np.reshape(eeg, (4, 200, EEG_LENGTH))\n        eeg = np.concatenate([eeg[0,:,:], eeg[1,:,:], eeg[2,:,:], eeg[3,:,:]],1)\n        list_eeg.append(eeg)\n            \n    eeg_l = np.concatenate(list_eeg, 1)\n    eeg_l /= 104\n        \n    time_temp = 0\n    time_start =  round(time_temp + 22 * 200) \n    time_stop =  round(time_temp + 32 * 200)\n\n    eeg_default = raw_eeg.loc[time_start: (time_stop - 1), :].reset_index(drop=True)\n    list_eeg = list()\n    for region in RAW_FEATS.keys():\n        eeg = np.zeros((len(RAW_FEATS[region]), eeg_default.shape[0]), dtype=np.float32)\n        for chan_i, chan in enumerate(RAW_FEATS[region]):\n            eeg_1 = eeg_default.loc[:, chan.split('-')[0]]\n            mean_value = eeg_1.mean()\n            eeg_1.fillna(value=mean_value, inplace=True)\n            eeg_1 = eeg_1.values\n\n            eeg_2 = eeg_default.loc[:, chan.split('-')[1]]\n            mean_value = eeg_2.mean()\n            eeg_2.fillna(value=mean_value, inplace=True)\n            eeg_2 = eeg_2.values\n\n            new_eeg = eeg_1 - eeg_2\n            new_eeg = signal.filtfilt(b, a, new_eeg, axis=0)\n            new_eeg = np.clip(new_eeg, -1024, 1024)\n            eeg[chan_i, :] = new_eeg\n\n        eeg = np.reshape(eeg, (4, 200, EEG_LENGTH))\n        eeg = np.concatenate([eeg[0,:,:], eeg[1,:,:], eeg[2,:,:], eeg[3,:,:]],1)\n        list_eeg.append(eeg)\n            \n    eeg_r = np.concatenate(list_eeg, 1)\n    eeg_r /= 104\n        \n    return eeg_l, eeg_c, eeg_r\n\ndef raw50seeg_from_eeg(parquet_path):\n    EEG_LENGTH = 50\n    raw_eeg = pd.read_parquet(parquet_path)\n    time_temp = 0\n    time_start =  round(time_temp + (50 - EEG_LENGTH) / 2 * 200) \n    time_stop =  round(time_temp + (50 + EEG_LENGTH) / 2 * 200)\n\n    eeg_default = raw_eeg.loc[time_start: (time_stop - 1), :].reset_index(drop=True)\n    list_eeg = list()\n    for region in RAW_FEATS.keys():\n\n        eeg = np.zeros((len(RAW_FEATS[region]), eeg_default.shape[0]), dtype=np.float32)\n        for chan_i, chan in enumerate(RAW_FEATS[region]):\n            eeg_1 = eeg_default.loc[:, chan.split('-')[0]]\n            mean_value = eeg_1.mean()\n            eeg_1.fillna(value=mean_value, inplace=True)\n            eeg_1 = eeg_1.values\n            \n            eeg_2 = eeg_default.loc[:, chan.split('-')[1]]\n            mean_value = eeg_2.mean()\n            eeg_2.fillna(value=mean_value, inplace=True)\n            eeg_2 = eeg_2.values\n            \n            new_eeg = eeg_1 - eeg_2\n            del eeg_1\n            del eeg_2\n            new_eeg = signal.filtfilt(b, a, new_eeg, axis=0)\n            new_eeg = np.clip(new_eeg, -1024, 1024).astype(\"float32\")\n            eeg[chan_i, :] = new_eeg\n        \n        eeg = np.reshape(eeg, (4, 200, EEG_LENGTH))\n        eeg = np.concatenate((eeg[0,:,:], eeg[1,:,:], eeg[2,:,:], eeg[3,:,:]), 1)\n        list_eeg.append(eeg)\n\n    eeg = np.concatenate(list_eeg, 1)\n    eeg /= 104\n    \n    return eeg","metadata":{"execution":{"iopub.status.busy":"2024-01-29T00:06:03.824991Z","iopub.execute_input":"2024-01-29T00:06:03.825301Z","iopub.status.idle":"2024-01-29T00:06:03.840047Z","shell.execute_reply.started":"2024-01-29T00:06:03.825272Z","shell.execute_reply":"2024-01-29T00:06:03.839248Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CLASSES = [\"seizure_vote\", \"lpd_vote\", \"gpd_vote\", \"lrda_vote\", \"grda_vote\", \"other_vote\"]\nN_CLASSES = len(CLASSES)\n\nif DEBUG == True:\n    test = pd.read_csv(\"/kaggle/input/hms-harmful-brain-activity-classification/train.csv\")[:40]\n    SPEC_PATH = '/kaggle/input/hms-harmful-brain-activity-classification/train_spectrograms/'\n    EEG_PATH = '/kaggle/input/hms-harmful-brain-activity-classification/train_eegs/'\nelse:\n    test = pd.read_csv(\"/kaggle/input/hms-harmful-brain-activity-classification/test.csv\")\n    SPEC_PATH = '/kaggle/input/hms-harmful-brain-activity-classification/test_spectrograms/'\n    EEG_PATH = '/kaggle/input/hms-harmful-brain-activity-classification/test_eegs/'\n\n# test = test.groupby(\"eeg_id\").head(1).reset_index(drop=True)\nprint(test.shape)\n\nspec_directory_path = 'spec_spectrograms/'\nif not os.path.exists(spec_directory_path):\n    os.makedirs(spec_directory_path)\n    \neeg_directory_path = 'eeg_spectrograms/'\nif not os.path.exists(eeg_directory_path):\n    os.makedirs(eeg_directory_path)\n    \nraw_10s_directory_path = 'eeg_10s_raws/'\nif not os.path.exists(raw_10s_directory_path):\n    os.makedirs(raw_10s_directory_path)\n    \nraw_50s_directory_path = 'eeg_50s_raws/'\nif not os.path.exists(raw_50s_directory_path):\n    os.makedirs(raw_50s_directory_path)\n\nfrom joblib import Parallel, delayed\n    \nEEG_IDS = test.eeg_id.unique()\n\ndef save(row):\n    eeg_id = row[\"eeg_id\"]\n    spec_id = row[\"spectrogram_id\"]\n    spec = pd.read_parquet(f\"{SPEC_PATH}{spec_id}.parquet\")\n    \n    spec_arr = spec.values[:, 1:].T.astype(\"float32\")  # (Hz, Time) = (400, 300)\n    \n    split_spec_arr = spec_arr[:, 0:300]\n    np.save(f'{spec_directory_path}{eeg_id}',split_spec_arr)\n    \n    img_l, img_c, img_r = raw10seeg_from_eeg(f'{EEG_PATH}{eeg_id}.parquet', eeg_id)\n    np.save(f'{raw_10s_directory_path}{eeg_id}_l',img_l)\n    np.save(f'{raw_10s_directory_path}{eeg_id}_c',img_c)\n    np.save(f'{raw_10s_directory_path}{eeg_id}_r',img_r)\n    \n    img = raw50seeg_from_eeg(f'{EEG_PATH}{eeg_id}.parquet')\n    np.save(f'{raw_50s_directory_path}{eeg_id}',img)\n    \n    img = stft_spec_from_eeg(f'{EEG_PATH}{eeg_id}.parquet')\n    np.save(f'{eeg_directory_path}{eeg_id}',img)\n\n_ = Parallel(n_jobs=4)(delayed(save)(row)\n                    for index, row in test.iterrows()\n                )","metadata":{"execution":{"iopub.status.busy":"2024-01-29T00:06:03.887721Z","iopub.execute_input":"2024-01-29T00:06:03.888283Z","iopub.status.idle":"2024-01-29T00:06:04.18401Z","shell.execute_reply.started":"2024-01-29T00:06:03.888259Z","shell.execute_reply":"2024-01-29T00:06:04.183056Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Config:\n    seed=2024\n    num_folds=5","metadata":{"execution":{"iopub.status.busy":"2024-01-29T00:06:04.185926Z","iopub.execute_input":"2024-01-29T00:06:04.18648Z","iopub.status.idle":"2024-01-29T00:06:04.19263Z","shell.execute_reply.started":"2024-01-29T00:06:04.186426Z","shell.execute_reply":"2024-01-29T00:06:04.191355Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seed_everything(seed):\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = True\n    torch.manual_seed(seed)\n    np.random.seed(seed)\n    random.seed(seed)\nseed_everything(Config.seed)","metadata":{"execution":{"iopub.status.busy":"2024-01-29T00:06:04.193972Z","iopub.execute_input":"2024-01-29T00:06:04.194444Z","iopub.status.idle":"2024-01-29T00:06:04.202255Z","shell.execute_reply.started":"2024-01-29T00:06:04.194411Z","shell.execute_reply":"2024-01-29T00:06:04.201029Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import timm\n","metadata":{"execution":{"iopub.status.busy":"2024-01-29T00:06:06.259982Z","iopub.execute_input":"2024-01-29T00:06:06.260317Z","iopub.status.idle":"2024-01-29T00:06:06.276602Z","shell.execute_reply.started":"2024-01-29T00:06:06.260287Z","shell.execute_reply":"2024-01-29T00:06:06.275359Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch.utils.data as data\nimport torchvision","metadata":{"execution":{"iopub.status.busy":"2024-01-29T00:06:06.27779Z","iopub.execute_input":"2024-01-29T00:06:06.278111Z","iopub.status.idle":"2024-01-29T00:06:06.295269Z","shell.execute_reply.started":"2024-01-29T00:06:06.278086Z","shell.execute_reply":"2024-01-29T00:06:06.294346Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import DataLoader\n","metadata":{"execution":{"iopub.status.busy":"2024-01-29T00:06:13.020785Z","iopub.execute_input":"2024-01-29T00:06:13.021086Z","iopub.status.idle":"2024-01-29T00:06:15.061773Z","shell.execute_reply.started":"2024-01-29T00:06:13.02106Z","shell.execute_reply":"2024-01-29T00:06:15.060643Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\n","metadata":{"execution":{"iopub.status.busy":"2024-01-29T00:06:15.063229Z","iopub.execute_input":"2024-01-29T00:06:15.06357Z","iopub.status.idle":"2024-01-29T00:06:15.344476Z","shell.execute_reply.started":"2024-01-29T00:06:15.063541Z","shell.execute_reply":"2024-01-29T00:06:15.343448Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from skimage.transform import resize\nclass ImageFolder(data.Dataset):\n    def __init__(self, df, test_imgsize):\n        super().__init__()\n        df['eeg_id'] = df[\"eeg_id\"]\n        self.spec_data_path = spec_directory_path\n        self.eeg_data_path = eeg_directory_path\n        self.raw_50s_data_path = raw_50s_directory_path\n        self.raw_10s_data_path = raw_10s_directory_path\n        self.df = df.reset_index(drop=True)\n        self.test_imgsize = test_imgsize\n#         self.test_transform = torchvision.transforms.Resize(self.test_imgsize)\n        \n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, index):\n        row = self.df.loc[index]\n        eeg_id = str(row.eeg_id)\n        spec_image_path = os.path.join(self.spec_data_path, eeg_id + \".npy\")\n        eeg_image_path = os.path.join(self.eeg_data_path, eeg_id + \".npy\")\n        raw_50s_image_path = os.path.join(self.raw_50s_data_path, eeg_id + \".npy\")\n        raw_10s_l_image_path = os.path.join(self.raw_10s_data_path, eeg_id + \"_l.npy\")\n        raw_10s_c_image_path = os.path.join(self.raw_10s_data_path, eeg_id + \"_c.npy\")\n        raw_10s_r_image_path = os.path.join(self.raw_10s_data_path, eeg_id + \"_r.npy\")\n        \n        spec_img = np.load(spec_image_path).astype(\"float32\")\n        raw_50s_img = np.load(raw_50s_image_path).astype(\"float32\")\n        raw_10s_l_img = np.load(raw_10s_l_image_path).astype(\"float32\")\n        raw_10s_c_img = np.load(raw_10s_c_image_path).astype(\"float32\")\n        raw_10s_r_img = np.load(raw_10s_r_image_path).astype(\"float32\")\n        eeg_img = np.load(eeg_image_path)\n        \n        eeg_img = resize(eeg_img, self.test_imgsize)\n        spec_img = resize(spec_img, self.test_imgsize)\n        raw_10s_l_img = resize(raw_10s_l_img, self.test_imgsize)\n        raw_10s_c_img = resize(raw_10s_c_img, self.test_imgsize)\n        raw_10s_r_img = resize(raw_10s_r_img, self.test_imgsize)\n        raw_50s_img = resize(raw_50s_img, self.test_imgsize)\n        \n        eeg_img = np.expand_dims(eeg_img, -1)\n        spec_img = np.expand_dims(spec_img, -1)\n        raw_50s_img = np.expand_dims(raw_50s_img, -1)\n        raw_10s_l_img = np.expand_dims(raw_10s_l_img, -1)\n        raw_10s_c_img = np.expand_dims(raw_10s_c_img, -1)\n        raw_10s_r_img = np.expand_dims(raw_10s_r_img, -1)\n\n        eps = 1e-6\n        spec_img = np.clip(spec_img,np.exp(-4),np.exp(8))\n        spec_img = np.log(spec_img)\n        spec_img = np.nan_to_num(spec_img, nan=0.0) \n        \n#         img_mean = eeg_img.mean(axis=(0, 1))\n#         img_std = eeg_img.std(axis=(0, 1))\n#         eeg_img = (eeg_img - img_mean) / (img_std + eps)\n        \n        img_mean = spec_img.mean(axis=(0, 1))\n        img_std = spec_img.std(axis=(0, 1))\n        spec_img = (spec_img - img_mean) / (img_std + eps)\n\n        return spec_img, eeg_img, raw_50s_img, raw_10s_l_img, raw_10s_c_img, raw_10s_r_img, eeg_id","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Net(nn.Module):\n    def __init__(self, back_bone, device_id):\n        super().__init__()\n        self.spec_model = timm.create_model(back_bone, num_classes=6, pretrained=False, in_chans=1)\n        self.eeg_model = timm.create_model(back_bone, num_classes=6, pretrained=False, in_chans=1)\n        self.raw_50s_model = timm.create_model(back_bone, num_classes=6, pretrained=False, in_chans=1)\n        self.raw_10s_model = timm.create_model(back_bone, num_classes=6, pretrained=False, in_chans=1)\n        \n        self.device_id = device_id\n        \n        self.spec_model.fc_norm = nn.Identity()\n        self.spec_model.head_drop = nn.Identity()\n        self.spec_model.head = nn.Identity()\n\n        self.eeg_model.fc_norm = nn.Identity()\n        self.eeg_model.head_drop = nn.Identity()\n        self.eeg_model.head = nn.Identity()\n\n        self.raw_50s_model.fc_norm = nn.Identity()\n        self.raw_50s_model.head_drop = nn.Identity()\n        self.raw_50s_model.head = nn.Identity()\n        \n        self.raw_10s_model.fc_norm = nn.Identity()\n        self.raw_10s_model.head_drop = nn.Identity()\n        self.raw_10s_model.head = nn.Identity()\n        \n        self.head = nn.Linear(384*4, 6)\n        self.head1 = nn.Linear(384, 6)\n        self.head2= nn.Linear(384, 6)\n        self.head3 = nn.Linear(384, 6)\n        self.head4 = nn.Linear(384, 6)\n       \n\n    def forward(self, spec_imgs, eeg_imgs, raw_50s_imgs, raw_10s_imgs):\n        spec_imgs = spec_imgs.transpose(1, 2).transpose(1, 3).contiguous()\n        eeg_imgs = eeg_imgs.transpose(1, 2).transpose(1, 3).contiguous()\n        raw_50s_imgs = raw_50s_imgs.transpose(1, 2).transpose(1, 3).contiguous()\n        raw_10s_imgs = raw_10s_imgs.transpose(1, 2).transpose(1, 3).contiguous()\n       \n        spec_feature = self.spec_model.forward_features(spec_imgs)[:, 0]\n        eeg_feature = self.eeg_model.forward_features(eeg_imgs)[:, 0]\n        raw_50s_feature = self.raw_50s_model.forward_features(raw_50s_imgs)[:, 0]\n        raw_10s_feature = self.raw_10s_model.forward_features(raw_10s_imgs)[:, 0]\n\n        feature = torch.cat((spec_feature, eeg_feature, raw_50s_feature, raw_10s_feature), 1)\n        logits = self.head(feature)\n        logits_1 = self.head1(spec_feature)\n        logits_2 = self.head2(eeg_feature)\n        logits_3 = self.head3(raw_50s_feature)\n        logits_4 = self.head4(raw_10s_feature)\n\n        return logits, logits_1, logits_2, logits_3, logits_4","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# vit_models = []\n# model_weights = [\"/kaggle/input/hms-stage2/fold_0_spec_raw_50_10_bestlb.pth\",\n#                 \"/kaggle/input/hms-stage2/fold_1_spec_raw_50_10_bestlb.pth\",\n#                 \"/kaggle/input/hms-stage2/fold_2_spec_raw_50_10_bestlb.pth\",\n#                 \"/kaggle/input/hms-stage2/fold_3_spec_raw_50_10_bestlb.pth\",\n#                 \"/kaggle/input/hms-stage2/fold_4_spec_raw_50_10_bestlb.pth\"]\n# model_types = [\"vit_small\", \"vit_small\", \"vit_small\", \"vit_small\", \"vit_small\"]\n# device = \"cuda:0\"\n# for i in range(len(model_types)):\n#     if model_types[i] == \"vit_small\":\n#         model = Net(\"vit_small_patch14_reg4_dinov2.lvd142m\", device).to(device)\n#         model.load_state_dict(torch.load(model_weights[i]))\n#         model.eval()\n#         vit_models.append(model)\n        \n# test_data = ImageFolder(test, (518, 518))\n# test_loader = DataLoader(test_data, batch_size=32, \n#                 pin_memory=False, num_workers=4, drop_last=False)\n# result_6 = {}\n# with torch.no_grad():\n#     for batch_idx, (spec_imgs, eeg_imgs, raw_50s_imgs, raw_10s_l_imgs, raw_10s_c_imgs, raw_10s_r_imgs, eeg_ids) in enumerate(test_loader):   \n#         spec_imgs = spec_imgs.to(device).float()\n#         eeg_imgs = eeg_imgs.to(device).float()\n#         raw_50s_imgs = raw_50s_imgs.to(device).float()\n#         raw_10s_l_imgs = raw_10s_l_imgs.to(device).float()\n#         raw_10s_c_imgs = raw_10s_c_imgs.to(device).float()\n#         raw_10s_r_imgs = raw_10s_r_imgs.to(device).float()\n#         ensemble_probs = torch.zeros((spec_imgs.shape[0], 6)).to(device)\n#         for model in vit_models:\n#             model.eval()\n#             logits_l, _, _, _, _ = model(spec_imgs, eeg_imgs, raw_50s_imgs, raw_10s_l_imgs)\n#             logits_c, _, _, _, _ = model(spec_imgs, eeg_imgs, raw_50s_imgs, raw_10s_c_imgs)\n#             logits_r, _, _, _, _ = model(spec_imgs, eeg_imgs, raw_50s_imgs, raw_10s_r_imgs)\n#             probs_l = logits_l.softmax(dim=1)\n#             probs_c = logits_c.softmax(dim=1)\n#             probs_r = logits_r.softmax(dim=1)\n#             for j in range(len(eeg_ids)):\n#                 print(\"eeg id: {}, {}, {}, {}\".format(eeg_ids[j], probs_l[j], probs_c[j], probs_r[j]))\n#             ensemble_probs += (probs_l + probs_c + probs_r) / 3\n#         ensemble_probs /= len(vit_models)\n#         ensemble_probs = ensemble_probs.detach().cpu().numpy()\n#         for j in range(len(eeg_ids)):\n#             eeg_id = eeg_ids[j]\n#             if eeg_id not in result_6.keys():\n#                 result_6[eeg_id] = np.array([0.0, 0.0, 0.0, 0.0, 0.0, 0.0])\n#             result_6[eeg_id] += ensemble_probs[j]\n#             print(\"eeg id: {}, {}\".format(eeg_id, ensemble_probs[j]))","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# for model in vit_models:\n#     del model\n# torch.cuda.empty_cache()\n# gc.collect()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"vit_models = []\nmodel_weights = [\"/kaggle/input/hms-stage2/fold_0_raw_50_10_bestlb_twostage.pth\",\n                \"/kaggle/input/hms-stage2/fold_1_raw_50_10_bestlb_twostage.pth\",\n                \"/kaggle/input/hms-stage2/fold_2_raw_50_10_bestlb_twostage.pth\",\n                \"/kaggle/input/hms-stage2/fold_3_raw_50_10_bestlb_twostage.pth\",\n                \"/kaggle/input/hms-stage2/fold_4_raw_50_10_bestlb_twostage.pth\"]\nmodel_types = [\"vit_small\", \"vit_small\", \"vit_small\", \"vit_small\", \"vit_small\"]\ndevice = \"cuda:0\"\nfor i in range(len(model_types)):\n    if model_types[i] == \"vit_small\":\n        model = Net(\"vit_small_patch14_reg4_dinov2.lvd142m\", device).to(device)\n        model.load_state_dict(torch.load(model_weights[i]))\n        model.eval()\n        vit_models.append(model)\n        \ntest_data = ImageFolder(test, (518, 518))\ntest_loader = DataLoader(test_data, batch_size=32, \n                pin_memory=False, num_workers=4, drop_last=False)\nresult_7 = {}\nwith torch.no_grad():\n    for batch_idx, (spec_imgs, eeg_imgs, raw_50s_imgs, raw_10s_l_imgs, raw_10s_c_imgs, raw_10s_r_imgs, eeg_ids) in enumerate(test_loader):   \n        spec_imgs = spec_imgs.to(device).float()\n        eeg_imgs = eeg_imgs.to(device).float()\n        raw_50s_imgs = raw_50s_imgs.to(device).float()\n        raw_10s_l_imgs = raw_10s_l_imgs.to(device).float()\n        raw_10s_c_imgs = raw_10s_c_imgs.to(device).float()\n        raw_10s_r_imgs = raw_10s_r_imgs.to(device).float()\n        ensemble_probs = torch.zeros((spec_imgs.shape[0], 6)).to(device)\n        for model in vit_models:\n            model.eval()\n            logits_l, _, _, _, _ = model(spec_imgs, eeg_imgs, raw_50s_imgs, raw_10s_l_imgs)\n            logits_c, _, _, _, _ = model(spec_imgs, eeg_imgs, raw_50s_imgs, raw_10s_c_imgs)\n            logits_r, _, _, _, _ = model(spec_imgs, eeg_imgs, raw_50s_imgs, raw_10s_r_imgs)\n            probs_l = logits_l.softmax(dim=1)\n            probs_c = logits_c.softmax(dim=1)\n            probs_r = logits_r.softmax(dim=1)\n            for j in range(len(eeg_ids)):\n                print(\"eeg id: {}, {}, {}, {}\".format(eeg_ids[j], probs_l[j], probs_c[j], probs_r[j]))\n            ensemble_probs += (probs_l + probs_c + probs_r)/3.0\n        ensemble_probs /= len(vit_models)\n        ensemble_probs = ensemble_probs.detach().cpu().numpy()\n        for j in range(len(eeg_ids)):\n            eeg_id = eeg_ids[j]\n            if eeg_id not in result_7.keys():\n                result_7[eeg_id] = np.array([0.0, 0.0, 0.0, 0.0, 0.0, 0.0])\n            result_7[eeg_id] += ensemble_probs[j]\n            print(\"eeg id: {}, {}\".format(eeg_id, ensemble_probs[j]))","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for model in vit_models:\n    del model\ntorch.cuda.empty_cache()\ngc.collect()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# class Net(nn.Module):\n#     def __init__(self, back_bone, device_id):\n#         super().__init__()\n#         self.spec_model = timm.create_model(\"vit_small_patch14_reg4_dinov2.lvd142m\", num_classes=6, pretrained=False, in_chans=1)\n#         self.eeg_model = timm.create_model(\"vit_small_patch14_reg4_dinov2.lvd142m\", num_classes=6, pretrained=False, in_chans=1)\n#         self.raw_50s_model = timm.create_model(back_bone, num_classes=6, pretrained=False, in_chans=1)\n#         self.raw_10s_model = timm.create_model(\"vit_small_patch14_reg4_dinov2.lvd142m\", num_classes=6, pretrained=False, in_chans=1)\n        \n#         self.device_id = device_id\n        \n#         self.spec_model.fc_norm = nn.Identity()\n#         self.spec_model.head_drop = nn.Identity()\n#         self.spec_model.head = nn.Identity()\n\n#         self.eeg_model.fc_norm = nn.Identity()\n#         self.eeg_model.head_drop = nn.Identity()\n#         self.eeg_model.head = nn.Identity()\n\n#         self.raw_50s_model.fc_norm = nn.Identity()\n#         self.raw_50s_model.head_drop = nn.Identity()\n#         self.raw_50s_model.head = nn.Identity()\n        \n#         self.raw_10s_model.fc_norm = nn.Identity()\n#         self.raw_10s_model.head_drop = nn.Identity()\n#         self.raw_10s_model.head = nn.Identity()\n        \n#         self.head = nn.Linear(384*3+1024, 6)\n#         self.head1 = nn.Linear(384, 6)\n#         self.head2= nn.Linear(384, 6)\n#         self.head3 = nn.Linear(1024, 6)\n#         self.head4 = nn.Linear(384, 6)\n       \n\n#     def forward(self, spec_imgs, eeg_imgs, raw_50s_imgs, raw_10s_imgs):\n#         spec_imgs = spec_imgs.transpose(1, 2).transpose(1, 3).contiguous()\n#         eeg_imgs = eeg_imgs.transpose(1, 2).transpose(1, 3).contiguous()\n#         raw_50s_imgs = raw_50s_imgs.transpose(1, 2).transpose(1, 3).contiguous()\n#         raw_10s_imgs = raw_10s_imgs.transpose(1, 2).transpose(1, 3).contiguous()\n       \n#         spec_feature = self.spec_model.forward_features(spec_imgs)[:, 0]\n#         eeg_feature = self.eeg_model.forward_features(eeg_imgs)[:, 0]\n#         raw_50s_feature = self.raw_50s_model.forward_features(raw_50s_imgs)[:, 0]\n#         raw_10s_feature = self.raw_10s_model.forward_features(raw_10s_imgs)[:, 0]\n\n#         feature = torch.cat((spec_feature, eeg_feature, raw_50s_feature, raw_10s_feature), 1)\n#         logits = self.head(feature)\n#         logits_1 = self.head1(spec_feature)\n#         logits_2 = self.head2(eeg_feature)\n#         logits_3 = self.head3(raw_50s_feature)\n#         logits_4 = self.head4(raw_10s_feature)\n\n#         return logits, logits_1, logits_2, logits_3, logits_4","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# vit_models = []\n# model_weights = [\"/kaggle/input/hms-stage2-vitlarge/fold_0_raw_50_10_bestlb_vitlarge.pth\",\n#                 \"/kaggle/input/hms-stage2-vitlarge/fold_1_raw_50_10_bestlb_vitlarge.pth\",\n#                 \"/kaggle/input/hms-stage2-vitlarge/fold_2_raw_50_10_bestlb_vitlarge.pth\",\n#                 \"/kaggle/input/hms-stage2-vitlarge/fold_3_raw_50_10_bestlb_vitlarge.pth\",\n#                 \"/kaggle/input/hms-stage2-vitlarge/fold_4_raw_50_10_bestlb_vitlarge.pth\"]\n# model_types = [\"vit_large\", \"vit_large\", \"vit_large\", \"vit_large\", \"vit_large\"]\n# device = \"cuda:0\"\n# for i in range(len(model_types)):\n#     if model_types[i] == \"vit_large\":\n#         model = Net(\"vit_large_patch14_reg4_dinov2.lvd142m\", device).to(device)\n#         model.load_state_dict(torch.load(model_weights[i]))\n#         model.eval()\n#         vit_models.append(model)\n        \n# test_data = ImageFolder(test, (518, 518))\n# test_loader = DataLoader(test_data, batch_size=16, \n#                 pin_memory=False, num_workers=4, drop_last=False)\n# result_8 = {}\n# with torch.no_grad():\n#     for batch_idx, (spec_imgs, eeg_imgs, raw_50s_imgs, raw_10s_l_imgs, raw_10s_c_imgs, raw_10s_r_imgs, eeg_ids) in enumerate(test_loader):   \n#         spec_imgs = spec_imgs.to(device).float()\n#         eeg_imgs = eeg_imgs.to(device).float()\n#         raw_50s_imgs = raw_50s_imgs.to(device).float()\n# #         raw_10s_l_imgs = raw_10s_l_imgs.to(device).float()\n#         raw_10s_c_imgs = raw_10s_c_imgs.to(device).float()\n# #         raw_10s_r_imgs = raw_10s_r_imgs.to(device).float()\n#         ensemble_probs = torch.zeros((spec_imgs.shape[0], 6)).to(device)\n#         for model in vit_models:\n#             model.eval()\n# #             logits_l, _, _, _, _ = model(spec_imgs, eeg_imgs, raw_50s_imgs, raw_10s_l_imgs)\n#             logits_c, _, _, _, _ = model(spec_imgs, eeg_imgs, raw_50s_imgs, raw_10s_c_imgs)\n# #             logits_r, _, _, _, _ = model(spec_imgs, eeg_imgs, raw_50s_imgs, raw_10s_r_imgs)\n# #             probs_l = logits_l.softmax(dim=1)\n#             probs_c = logits_c.softmax(dim=1)\n# #             probs_r = logits_r.softmax(dim=1)\n# #             for j in range(len(eeg_ids)):\n# #                 print(\"eeg id: {}, {}, {}, {}\".format(eeg_ids[j], probs_l[j], probs_c[j], probs_r[j]))\n#             ensemble_probs += probs_c\n#         ensemble_probs /= len(vit_models)\n#         ensemble_probs = ensemble_probs.detach().cpu().numpy()\n#         for j in range(len(eeg_ids)):\n#             eeg_id = eeg_ids[j]\n#             if eeg_id not in result_8.keys():\n#                 result_8[eeg_id] = np.array([0.0, 0.0, 0.0, 0.0, 0.0, 0.0])\n#             result_8[eeg_id] += ensemble_probs[j]\n#             print(\"eeg id: {}, {}\".format(eeg_id, ensemble_probs[j]))","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# for model in vit_models:\n#     del model\n# torch.cuda.empty_cache()\n# gc.collect()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# class Net(nn.Module):\n#     def __init__(self, back_bone, device_id):\n#         super().__init__()\n#         self.spec_model = timm.create_model(\"vit_small_patch14_reg4_dinov2.lvd142m\", num_classes=6, pretrained=False, in_chans=1)\n#         self.eeg_model = timm.create_model(\"vit_small_patch14_reg4_dinov2.lvd142m\", num_classes=6, pretrained=False, in_chans=1)\n#         self.raw_50s_model = timm.create_model(back_bone, num_classes=6, pretrained=False, in_chans=1)\n#         self.raw_10s_model = timm.create_model(back_bone, num_classes=6, pretrained=False, in_chans=1)\n        \n#         self.device_id = device_id\n        \n#         self.spec_model.fc_norm = nn.Identity()\n#         self.spec_model.head_drop = nn.Identity()\n#         self.spec_model.head = nn.Identity()\n\n#         self.eeg_model.fc_norm = nn.Identity()\n#         self.eeg_model.head_drop = nn.Identity()\n#         self.eeg_model.head = nn.Identity()\n\n#         self.raw_50s_model.fc_norm = nn.Identity()\n#         self.raw_50s_model.head_drop = nn.Identity()\n#         self.raw_50s_model.head = nn.Identity()\n        \n#         self.raw_10s_model.fc_norm = nn.Identity()\n#         self.raw_10s_model.head_drop = nn.Identity()\n#         self.raw_10s_model.head = nn.Identity()\n        \n#         self.head = nn.Linear(384*2+768*2, 6)\n#         self.head1 = nn.Linear(384, 6)\n#         self.head2= nn.Linear(384, 6)\n#         self.head3 = nn.Linear(768, 6)\n#         self.head4 = nn.Linear(768, 6)\n       \n\n#     def forward(self, spec_imgs, eeg_imgs, raw_50s_imgs, raw_10s_imgs):\n#         spec_imgs = spec_imgs.transpose(1, 2).transpose(1, 3).contiguous()\n#         eeg_imgs = eeg_imgs.transpose(1, 2).transpose(1, 3).contiguous()\n#         raw_50s_imgs = raw_50s_imgs.transpose(1, 2).transpose(1, 3).contiguous()\n#         raw_10s_imgs = raw_10s_imgs.transpose(1, 2).transpose(1, 3).contiguous()\n       \n#         spec_feature = self.spec_model.forward_features(spec_imgs)[:, 0]\n#         eeg_feature = self.eeg_model.forward_features(eeg_imgs)[:, 0]\n#         raw_50s_feature = self.raw_50s_model.forward_features(raw_50s_imgs)[:, 0]\n#         raw_10s_feature = self.raw_10s_model.forward_features(raw_10s_imgs)[:, 0]\n\n#         feature = torch.cat((spec_feature, eeg_feature, raw_50s_feature, raw_10s_feature), 1)\n#         logits = self.head(feature)\n#         logits_1 = self.head1(spec_feature)\n#         logits_2 = self.head2(eeg_feature)\n#         logits_3 = self.head3(raw_50s_feature)\n#         logits_4 = self.head4(raw_10s_feature)\n\n#         return logits, logits_1, logits_2, logits_3, logits_4","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# vit_models = []\n# model_weights = [\"/kaggle/input/hms-bestlb-vitbase/fold_0_exp_6_bestlb.pth\",\n#                 \"/kaggle/input/hms-bestlb-vitbase/fold_1_exp_6_bestlb.pth\",\n#                 \"/kaggle/input/hms-bestlb-vitbase/fold_2_exp_6_bestlb.pth\",\n#                 \"/kaggle/input/hms-bestlb-vitbase/fold_3_exp_6_bestlb.pth\",\n#                 \"/kaggle/input/hms-bestlb-vitbase/fold_4_exp_6_bestlb.pth\"]\n# model_types = [\"vit_base\", \"vit_base\", \"vit_base\", \"vit_base\", \"vit_base\"]\n# device = \"cuda:0\"\n# for i in range(len(model_types)):\n#     if model_types[i] == \"vit_base\":\n#         print(model_weights[i])\n#         model = Net(\"vit_base_patch14_reg4_dinov2.lvd142m\", device).to(device)\n#         model.load_state_dict(torch.load(model_weights[i]))\n#         model.eval()\n#         vit_models.append(model)\n        \n# test_data = ImageFolder(test, (518, 518))\n# test_loader = DataLoader(test_data, batch_size=32, \n#                 pin_memory=False, num_workers=4, drop_last=False)\n# result_5 = {}\n# with torch.no_grad():\n#     for batch_idx, (spec_imgs, eeg_imgs, raw_50s_imgs, raw_10s_l_imgs, raw_10s_c_imgs, raw_10s_r_imgs, eeg_ids) in enumerate(test_loader):   \n#         spec_imgs = spec_imgs.to(device).float()\n#         eeg_imgs = eeg_imgs.to(device).float()\n#         raw_50s_imgs = raw_50s_imgs.to(device).float()\n#         raw_10s_l_imgs = raw_10s_l_imgs.to(device).float()\n#         raw_10s_c_imgs = raw_10s_c_imgs.to(device).float()\n#         raw_10s_r_imgs = raw_10s_r_imgs.to(device).float()\n#         ensemble_probs = torch.zeros((spec_imgs.shape[0], 6)).to(device)\n#         for model in vit_models:\n#             model.eval()\n#             logits_l, _, _, _, _ = model(spec_imgs, eeg_imgs, raw_50s_imgs, raw_10s_l_imgs)\n#             logits_c, _, _, _, _ = model(spec_imgs, eeg_imgs, raw_50s_imgs, raw_10s_c_imgs)\n#             logits_r, _, _, _, _ = model(spec_imgs, eeg_imgs, raw_50s_imgs, raw_10s_r_imgs)\n#             probs_l = logits_l.softmax(dim=1)\n#             probs_c = logits_c.softmax(dim=1)\n#             probs_r = logits_r.softmax(dim=1)\n#             for j in range(len(eeg_ids)):\n#                 print(\"eeg id: {}, {}, {}, {}\".format(eeg_ids[j], probs_l[j], probs_c[j], probs_r[j]))\n#             ensemble_probs += (probs_l + probs_c + probs_r)/3\n#         ensemble_probs /= len(vit_models)\n#         ensemble_probs = ensemble_probs.detach().cpu().numpy()\n#         for j in range(len(eeg_ids)):\n#             eeg_id = eeg_ids[j]\n#             if eeg_id not in result_5.keys():\n#                 result_5[eeg_id] = np.array([0.0, 0.0, 0.0, 0.0, 0.0, 0.0])\n#             result_5[eeg_id] += ensemble_probs[j]\n#             print(\"eeg id: {}, {}\".format(eeg_id, ensemble_probs[j]))","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# for model in vit_models:\n#     del model\n# torch.cuda.empty_cache()\n# gc.collect()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# !rm -rf /kaggle/working/eeg_spectrograms/*","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def stft_spec_from_eeg(parquet_path):\n#     EEG_LENGTH = 50\n#     eeg = pd.read_parquet(parquet_path)\n    \n#     time_temp = 0\n#     time_start = round(time_temp + (50 - EEG_LENGTH) / 2 * 200) \n#     time_stop = round(time_temp + (50 + EEG_LENGTH) / 2 * 200)\n    \n#     eeg = eeg.iloc[time_start: time_stop]\n    \n#     list_eeg = list()\n#     for k in range(4):\n#         COLS = FEATS[k]\n#         img = np.zeros((128,256,4),dtype='float32')\n#         for kk in range(4):\n#             eeg_1 = eeg[COLS[kk]]\n#             mean_value = eeg_1.mean()\n#             eeg_1.fillna(value=mean_value, inplace=True)\n#             eeg_1 = eeg_1.values\n            \n#             eeg_2 = eeg[COLS[kk+1]]\n#             mean_value = eeg_2.mean()\n#             eeg_2.fillna(value=mean_value, inplace=True)\n#             eeg_2 = eeg_2.values\n            \n#             new_eeg = eeg_1 - eeg_2\n#             del eeg_1\n#             del eeg_2\n#             # new_eeg = signal.filtfilt(b, a, new_eeg, axis=0)\n#             fs = 200  \n#             nperseg = 39\n#             noverlap = 0\n#             f, t, spec = signal.spectrogram(new_eeg, fs, nperseg=nperseg, noverlap=noverlap, nfft=256)\n#             # print(spec.shape)\n            \n#             spec = np.abs(spec) \n#             spec = np.log1p(spec).astype(\"float32\")\n\n#             img[:,:,kk] += spec[:128, :]\n#         img = np.concatenate((img[:,:,0], img[:,:,1], img[:,:,2], img[:,:,3]), 1)\n#         # print(img.shape)\n#         list_eeg.append(img)\n#     img = np.concatenate(list_eeg, 0)    \n#     # img = np.concatenate((img[:,:,0], img[:,:,1], img[:,:,2], img[:,:,3]), 0)\n#     img /= 2\n#     return img","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def save(row):\n#     eeg_id = row[\"eeg_id\"]\n#     spec_id = row[\"spectrogram_id\"]\n    \n#     img = stft_spec_from_eeg(f'{EEG_PATH}{eeg_id}.parquet')\n#     np.save(f'{eeg_directory_path}{eeg_id}',img)\n\n# _ = Parallel(n_jobs=4)(delayed(save)(row)\n#                     for index, row in test.iterrows()\n#                 )","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# class ImageFolder(data.Dataset):\n#     def __init__(self, df, test_imgsize):\n#         super().__init__()\n#         df['eeg_id'] = df[\"eeg_id\"]\n#         self.spec_data_path = spec_directory_path\n#         self.eeg_data_path = eeg_directory_path\n#         self.raw_50s_data_path = raw_50s_directory_path\n#         self.raw_10s_data_path = raw_10s_directory_path\n#         self.df = df.reset_index(drop=True)\n#         self.test_imgsize = test_imgsize\n# #         self.test_transform = torchvision.transforms.Resize(self.test_imgsize)\n        \n#     def __len__(self):\n#         return len(self.df)\n\n#     def __getitem__(self, index):\n#         row = self.df.loc[index]\n#         eeg_id = str(row.eeg_id)\n#         spec_image_path = os.path.join(self.spec_data_path, eeg_id + \".npy\")\n#         eeg_image_path = os.path.join(self.eeg_data_path, eeg_id + \".npy\")\n#         raw_50s_image_path = os.path.join(self.raw_50s_data_path, eeg_id + \".npy\")\n#         raw_10s_l_image_path = os.path.join(self.raw_10s_data_path, eeg_id + \"_l.npy\")\n#         raw_10s_c_image_path = os.path.join(self.raw_10s_data_path, eeg_id + \"_c.npy\")\n#         raw_10s_r_image_path = os.path.join(self.raw_10s_data_path, eeg_id + \"_r.npy\")\n        \n#         spec_img = np.load(spec_image_path).astype(\"float32\")\n#         raw_50s_img = np.load(raw_50s_image_path).astype(\"float32\")\n#         raw_10s_l_img = np.load(raw_10s_l_image_path).astype(\"float32\")\n#         raw_10s_c_img = np.load(raw_10s_c_image_path).astype(\"float32\")\n#         raw_10s_r_img = np.load(raw_10s_r_image_path).astype(\"float32\")\n#         eeg_img = np.load(eeg_image_path)\n        \n#         spec_img = resize(spec_img, self.test_imgsize)\n#         raw_10s_l_img = resize(raw_10s_l_img, self.test_imgsize)\n#         raw_10s_c_img = resize(raw_10s_c_img, self.test_imgsize)\n#         raw_10s_r_img = resize(raw_10s_r_img, self.test_imgsize)\n#         raw_50s_img = resize(raw_50s_img, self.test_imgsize)\n        \n#         eeg_img = np.expand_dims(eeg_img, -1)\n#         spec_img = np.expand_dims(spec_img, -1)\n#         raw_50s_img = np.expand_dims(raw_50s_img, -1)\n#         raw_10s_l_img = np.expand_dims(raw_10s_l_img, -1)\n#         raw_10s_c_img = np.expand_dims(raw_10s_c_img, -1)\n#         raw_10s_r_img = np.expand_dims(raw_10s_r_img, -1)\n\n#         eps = 1e-6\n#         spec_img = np.clip(spec_img,np.exp(-4),np.exp(8))\n#         spec_img = np.log(spec_img)\n#         spec_img = np.nan_to_num(spec_img, nan=0.0) \n        \n# #         img_mean = eeg_img.mean(axis=(0, 1))\n# #         img_std = eeg_img.std(axis=(0, 1))\n# #         eeg_img = (eeg_img - img_mean) / (img_std + eps)\n        \n#         img_mean = spec_img.mean(axis=(0, 1))\n#         img_std = spec_img.std(axis=(0, 1))\n#         spec_img = (spec_img - img_mean) / (img_std + eps)\n\n#         return spec_img, eeg_img, raw_50s_img, raw_10s_l_img, raw_10s_c_img, raw_10s_r_img, eeg_id\n","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# class Net(nn.Module):\n#     def __init__(self, back_bone, device_id):\n#         super().__init__()\n#         self.spec_model = timm.create_model(back_bone, num_classes=6, pretrained=False, in_chans=1)\n#         self.eeg_model = timm.create_model(back_bone, num_classes=6, pretrained=False, in_chans=1, dynamic_img_pad=True, dynamic_img_size=True)\n#         self.raw_50s_model = timm.create_model(back_bone, num_classes=6, pretrained=False, in_chans=1)\n#         self.raw_10s_model = timm.create_model(back_bone, num_classes=6, pretrained=False, in_chans=1)\n        \n#         self.device_id = device_id\n        \n#         self.spec_model.fc_norm = nn.Identity()\n#         self.spec_model.head_drop = nn.Identity()\n#         self.spec_model.head = nn.Identity()\n\n#         self.eeg_model.fc_norm = nn.Identity()\n#         self.eeg_model.head_drop = nn.Identity()\n#         self.eeg_model.head = nn.Identity()\n\n#         self.raw_50s_model.fc_norm = nn.Identity()\n#         self.raw_50s_model.head_drop = nn.Identity()\n#         self.raw_50s_model.head = nn.Identity()\n        \n#         self.raw_10s_model.fc_norm = nn.Identity()\n#         self.raw_10s_model.head_drop = nn.Identity()\n#         self.raw_10s_model.head = nn.Identity()\n        \n#         self.head = nn.Linear(384*4, 6)\n#         self.head1 = nn.Linear(384, 6)\n#         self.head2= nn.Linear(384, 6)\n#         self.head3 = nn.Linear(384, 6)\n#         self.head4 = nn.Linear(384, 6)\n       \n\n#     def forward(self, spec_imgs, eeg_imgs, raw_50s_imgs, raw_10s_imgs):\n#         spec_imgs = spec_imgs.transpose(1, 2).transpose(1, 3).contiguous()\n#         eeg_imgs = eeg_imgs.transpose(1, 2).transpose(1, 3).contiguous()\n#         raw_50s_imgs = raw_50s_imgs.transpose(1, 2).transpose(1, 3).contiguous()\n#         raw_10s_imgs = raw_10s_imgs.transpose(1, 2).transpose(1, 3).contiguous()\n       \n#         spec_feature = self.spec_model.forward_features(spec_imgs)[:, 0]\n#         eeg_feature = self.eeg_model.forward_features(eeg_imgs)[:, 0]\n#         raw_50s_feature = self.raw_50s_model.forward_features(raw_50s_imgs)[:, 0]\n#         raw_10s_feature = self.raw_10s_model.forward_features(raw_10s_imgs)[:, 0]\n\n#         feature = torch.cat((spec_feature, eeg_feature, raw_50s_feature, raw_10s_feature), 1)\n#         logits = self.head(feature)\n#         logits_1 = self.head1(spec_feature)\n#         logits_2 = self.head2(eeg_feature)\n#         logits_3 = self.head3(raw_50s_feature)\n#         logits_4 = self.head4(raw_10s_feature)\n\n#         return logits, logits_1, logits_2, logits_3, logits_4","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# vit_models = []\n# model_weights = [\"/kaggle/input/hms-stage2/fold_0_exp_5_bestlb.pth\",\n#                 \"/kaggle/input/hms-stage2/fold_1_exp_5_bestlb.pth\",\n#                 \"/kaggle/input/hms-stage2/fold_2_exp_5_bestlb.pth\",\n#                 \"/kaggle/input/hms-stage2/fold_3_exp_5_bestlb.pth\",\n#                 \"/kaggle/input/hms-stage2/fold_4_exp_5_bestlb.pth\"]\n# model_types = [\"vit_small\", \"vit_small\", \"vit_small\", \"vit_small\", \"vit_small\"]\n# device = \"cuda:0\"\n# for i in range(len(model_types)):\n#     if model_types[i] == \"vit_small\":\n#         print(model_weights[i])\n#         model = Net(\"vit_small_patch14_reg4_dinov2.lvd142m\", device).to(device)\n#         model.load_state_dict(torch.load(model_weights[i]))\n#         model.eval()\n#         vit_models.append(model)\n        \n# test_data = ImageFolder(test, (518, 518))\n# test_loader = DataLoader(test_data, batch_size=32, \n#                 pin_memory=False, num_workers=4, drop_last=False)\n# result_4 = {}\n# with torch.no_grad():\n#     for batch_idx, (spec_imgs, eeg_imgs, raw_50s_imgs, raw_10s_l_imgs, raw_10s_c_imgs, raw_10s_r_imgs, eeg_ids) in enumerate(test_loader):   \n#         spec_imgs = spec_imgs.to(device).float()\n#         eeg_imgs = eeg_imgs.to(device).float()\n#         raw_50s_imgs = raw_50s_imgs.to(device).float()\n#         raw_10s_l_imgs = raw_10s_l_imgs.to(device).float()\n#         raw_10s_c_imgs = raw_10s_c_imgs.to(device).float()\n#         raw_10s_r_imgs = raw_10s_r_imgs.to(device).float()\n#         ensemble_probs = torch.zeros((spec_imgs.shape[0], 6)).to(device)\n#         for model in vit_models:\n#             model.eval()\n#             logits_l, _, _, _, _ = model(spec_imgs, eeg_imgs, raw_50s_imgs, raw_10s_l_imgs)\n#             logits_c, _, _, _, _ = model(spec_imgs, eeg_imgs, raw_50s_imgs, raw_10s_c_imgs)\n#             logits_r, _, _, _, _ = model(spec_imgs, eeg_imgs, raw_50s_imgs, raw_10s_r_imgs)\n#             probs_l = logits_l.softmax(dim=1)\n#             probs_c = logits_c.softmax(dim=1)\n#             probs_r = logits_r.softmax(dim=1)\n#             for j in range(len(eeg_ids)):\n#                 print(\"eeg id: {}, {}, {}, {}\".format(eeg_ids[j], probs_l[j], probs_c[j], probs_r[j]))\n#             ensemble_probs += (probs_l + probs_c + probs_r)/3\n#         ensemble_probs /= len(vit_models)\n#         ensemble_probs = ensemble_probs.detach().cpu().numpy()\n#         for j in range(len(eeg_ids)):\n#             eeg_id = eeg_ids[j]\n#             if eeg_id not in result_4.keys():\n#                 result_4[eeg_id] = np.array([0.0, 0.0, 0.0, 0.0, 0.0, 0.0])\n#             result_4[eeg_id] += ensemble_probs[j]\n#             print(\"eeg id: {}, {}\".format(eeg_id, ensemble_probs[j]))","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# for model in vit_models:\n#     del model\n# torch.cuda.empty_cache()\n# gc.collect()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# !rm -rf /kaggle/working/eeg_spectrograms/*\n# !rm -rf /kaggle/working/eeg_50s_raws/*","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def stft_spec_from_eeg(parquet_path):\n#     EEG_LENGTH = 50\n#     eeg = pd.read_parquet(parquet_path)\n#     time_temp = 0\n#     time_start = round(time_temp + (50 - EEG_LENGTH) / 2 * 200) \n#     time_stop = round(time_temp + (50 + EEG_LENGTH) / 2 * 200)\n    \n#     eeg = eeg.iloc[time_start: time_stop]\n    \n#     img = np.zeros((128,256,4),dtype='float32')\n#     for k in range(4):\n#         COLS = FEATS[k]\n#         for kk in range(4):\n#             eeg_1 = eeg[COLS[kk]]\n#             mean_value = eeg_1.mean()\n#             eeg_1.fillna(value=mean_value, inplace=True)\n#             eeg_1 = eeg_1.values\n            \n#             eeg_2 = eeg[COLS[kk+1]]\n#             mean_value = eeg_2.mean()\n#             eeg_2.fillna(value=mean_value, inplace=True)\n#             eeg_2 = eeg_2.values\n            \n#             new_eeg = eeg_1 - eeg_2\n#             del eeg_1\n#             del eeg_2\n#             # new_eeg = signal.filtfilt(b, a, new_eeg, axis=0)\n#             fs = 200  \n#             nperseg = len(new_eeg)//256\n#             noverlap = 0\n#             f, t, spec = signal.stft(new_eeg, fs, nperseg=nperseg, noverlap=noverlap, nfft=256)\n\n#             spec = np.abs(spec) \n#             spec = np.log1p(spec).astype(\"float32\")\n\n#             img[:,:,k] += spec[:128, 1:257]\n#         img[:,:,k] /= 4\n#     img = np.concatenate((img[:,:,0], img[:,:,1], img[:,:,2], img[:,:,3]), 0)\n#     return img\n\n# def raw50seeg_from_eeg(parquet_path, eeg_id):\n#     EEG_LENGTH = 20\n#     raw_eeg = pd.read_parquet(parquet_path)\n#     time_temp = 0\n#     time_start =  round(time_temp + (50 - EEG_LENGTH) / 2 * 200) \n#     time_stop =  round(time_temp + (50 + EEG_LENGTH) / 2 * 200)\n\n#     eeg_default = raw_eeg.loc[time_start: (time_stop - 1), :].reset_index(drop=True)\n#     list_eeg = list()\n#     for region in RAW_FEATS.keys():\n\n#         eeg = np.zeros((len(RAW_FEATS[region]), eeg_default.shape[0]), dtype=np.float32)\n#         for chan_i, chan in enumerate(RAW_FEATS[region]):\n#             eeg_1 = eeg_default.loc[:, chan.split('-')[0]]\n#             mean_value = eeg_1.mean()\n#             eeg_1.fillna(value=mean_value, inplace=True)\n#             eeg_1 = eeg_1.values\n            \n#             eeg_2 = eeg_default.loc[:, chan.split('-')[1]]\n#             mean_value = eeg_2.mean()\n#             eeg_2.fillna(value=mean_value, inplace=True)\n#             eeg_2 = eeg_2.values\n            \n#             new_eeg = eeg_1 - eeg_2\n#             del eeg_1\n#             del eeg_2\n#             new_eeg = signal.filtfilt(b, a, new_eeg, axis=0)\n#             new_eeg = np.clip(new_eeg, -1024, 1024).astype(\"float32\")\n#             eeg[chan_i, :] = new_eeg\n        \n#         eeg = np.reshape(eeg, (4, 200, EEG_LENGTH))\n#         eeg = np.concatenate((eeg[0,:,:], eeg[1,:,:], eeg[2,:,:], eeg[3,:,:]), 1)\n#         list_eeg.append(eeg)\n\n#     eeg = np.concatenate(list_eeg, 1)\n#     eeg /= 104\n    \n#     return eeg","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def save(row):\n#     eeg_id = row[\"eeg_id\"]\n#     spec_id = row[\"spectrogram_id\"]\n    \n#     img = stft_spec_from_eeg(f'{EEG_PATH}{eeg_id}.parquet')\n#     np.save(f'{eeg_directory_path}{eeg_id}',img)\n#     img = raw50seeg_from_eeg(f'{EEG_PATH}{eeg_id}.parquet', eeg_id)\n#     np.save(f'{raw_50s_directory_path}{eeg_id}',img)\n\n\n# _ = Parallel(n_jobs=4)(delayed(save)(row)\n#                     for index, row in test.iterrows()\n#                 )","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# class ImageFolder(data.Dataset):\n#     def __init__(self, df, test_imgsize):\n#         super().__init__()\n#         df['eeg_id'] = df[\"eeg_id\"]\n#         self.spec_data_path = spec_directory_path\n#         self.eeg_data_path = eeg_directory_path\n#         self.raw_50s_data_path = raw_50s_directory_path\n#         self.raw_10s_data_path = raw_10s_directory_path\n#         self.df = df.reset_index(drop=True)\n#         self.test_imgsize = test_imgsize\n# #         self.test_transform = torchvision.transforms.Resize(self.test_imgsize)\n        \n#     def __len__(self):\n#         return len(self.df)\n\n#     def __getitem__(self, index):\n#         row = self.df.loc[index]\n#         eeg_id = str(row.eeg_id)\n#         spec_image_path = os.path.join(self.spec_data_path, eeg_id + \".npy\")\n#         eeg_image_path = os.path.join(self.eeg_data_path, eeg_id + \".npy\")\n#         raw_50s_image_path = os.path.join(self.raw_50s_data_path, eeg_id + \".npy\")\n#         raw_10s_l_image_path = os.path.join(self.raw_10s_data_path, eeg_id + \"_l.npy\")\n#         raw_10s_c_image_path = os.path.join(self.raw_10s_data_path, eeg_id + \"_c.npy\")\n#         raw_10s_r_image_path = os.path.join(self.raw_10s_data_path, eeg_id + \"_r.npy\")\n        \n#         spec_img = np.load(spec_image_path).astype(\"float32\")\n#         raw_50s_img = np.load(raw_50s_image_path).astype(\"float32\")\n#         raw_10s_l_img = np.load(raw_10s_l_image_path).astype(\"float32\")\n#         raw_10s_c_img = np.load(raw_10s_c_image_path).astype(\"float32\")\n#         raw_10s_r_img = np.load(raw_10s_r_image_path).astype(\"float32\")\n#         eeg_img = np.load(eeg_image_path)\n        \n#         eeg_img = resize(eeg_img, self.test_imgsize)\n#         spec_img = resize(spec_img, self.test_imgsize)\n#         raw_10s_l_img = resize(raw_10s_l_img, self.test_imgsize)\n#         raw_10s_c_img = resize(raw_10s_c_img, self.test_imgsize)\n#         raw_10s_r_img = resize(raw_10s_r_img, self.test_imgsize)\n#         raw_50s_img = resize(raw_50s_img, self.test_imgsize)\n        \n#         eeg_img = np.expand_dims(eeg_img, -1)\n#         spec_img = np.expand_dims(spec_img, -1)\n#         raw_50s_img = np.expand_dims(raw_50s_img, -1)\n#         raw_10s_l_img = np.expand_dims(raw_10s_l_img, -1)\n#         raw_10s_c_img = np.expand_dims(raw_10s_c_img, -1)\n#         raw_10s_r_img = np.expand_dims(raw_10s_r_img, -1)\n\n#         eps = 1e-6\n#         spec_img = np.clip(spec_img,np.exp(-4),np.exp(8))\n#         spec_img = np.log(spec_img)\n#         spec_img = np.nan_to_num(spec_img, nan=0.0) \n        \n#         img_mean = eeg_img.mean(axis=(0, 1))\n#         img_std = eeg_img.std(axis=(0, 1))\n#         eeg_img = (eeg_img - img_mean) / (img_std + eps)\n        \n#         img_mean = spec_img.mean(axis=(0, 1))\n#         img_std = spec_img.std(axis=(0, 1))\n#         spec_img = (spec_img - img_mean) / (img_std + eps)\n\n#         return spec_img, eeg_img, raw_50s_img, raw_10s_l_img, raw_10s_c_img, raw_10s_r_img, eeg_id","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# class Net(nn.Module):\n#     def __init__(self, back_bone, device_id):\n#         super().__init__()\n#         self.spec_model = timm.create_model(\"vit_small_patch14_reg4_dinov2.lvd142m\", num_classes=6, pretrained=False, in_chans=1)\n#         self.eeg_model = timm.create_model(\"vit_small_patch14_reg4_dinov2.lvd142m\", num_classes=6, pretrained=False, in_chans=1)\n#         self.raw_50s_model = timm.create_model(back_bone, num_classes=6, pretrained=False, in_chans=1)\n#         self.raw_10s_model = timm.create_model(back_bone, num_classes=6, pretrained=False, in_chans=1)\n        \n#         self.device_id = device_id\n        \n#         self.spec_model.fc_norm = nn.Identity()\n#         self.spec_model.head_drop = nn.Identity()\n#         self.spec_model.head = nn.Identity()\n\n#         self.eeg_model.fc_norm = nn.Identity()\n#         self.eeg_model.head_drop = nn.Identity()\n#         self.eeg_model.head = nn.Identity()\n\n#         self.raw_50s_model.fc_norm = nn.Identity()\n#         self.raw_50s_model.head_drop = nn.Identity()\n#         self.raw_50s_model.head = nn.Identity()\n        \n#         self.raw_10s_model.fc_norm = nn.Identity()\n#         self.raw_10s_model.head_drop = nn.Identity()\n#         self.raw_10s_model.head = nn.Identity()\n        \n#         self.head = nn.Linear(768*2+384*2, 6)\n#         self.head1 = nn.Linear(384, 6)\n#         self.head2= nn.Linear(384, 6)\n#         self.head3 = nn.Linear(768, 6)\n#         self.head4 = nn.Linear(768, 6)\n       \n\n#     def forward(self, spec_imgs, eeg_imgs, raw_50s_imgs, raw_10s_imgs):\n#         spec_imgs = spec_imgs.transpose(1, 2).transpose(1, 3).contiguous()\n#         eeg_imgs = eeg_imgs.transpose(1, 2).transpose(1, 3).contiguous()\n#         raw_50s_imgs = raw_50s_imgs.transpose(1, 2).transpose(1, 3).contiguous()\n#         raw_10s_imgs = raw_10s_imgs.transpose(1, 2).transpose(1, 3).contiguous()\n       \n#         spec_feature = self.spec_model.forward_features(spec_imgs)[:, 0]\n#         eeg_feature = self.eeg_model.forward_features(eeg_imgs)[:, 0]\n#         raw_50s_feature = self.raw_50s_model.forward_features(raw_50s_imgs)[:, 0]\n#         raw_10s_feature = self.raw_10s_model.forward_features(raw_10s_imgs)[:, 0]\n\n#         feature = torch.cat((spec_feature, eeg_feature, raw_50s_feature, raw_10s_feature), 1)\n#         logits = self.head(feature)\n\n#         return logits","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# vit_models = []\n# model_weights = [\"/kaggle/input/hms-bestlb-vitbase/fold_0_exp_7_bestlb.pth\",\n#                 \"/kaggle/input/hms-bestlb-vitbase/fold_1_exp_7_bestlb.pth\",\n#                 \"/kaggle/input/hms-bestlb-vitbase/fold_2_exp_7_bestlb.pth\",\n#                 \"/kaggle/input/hms-bestlb-vitbase/fold_3_exp_7_bestlb.pth\",\n#                 \"/kaggle/input/hms-bestlb-vitbase/fold_4_exp_7_bestlb.pth\"]\n# model_types = [\"vit_base\", \"vit_base\", \"vit_base\", \"vit_base\", \"vit_base\"]\n# device = \"cuda:0\"\n# for i in range(len(model_types)):\n#     if model_types[i] == \"vit_base\":\n#         model = Net(\"vit_base_patch14_reg4_dinov2.lvd142m\", device).to(device)\n#         model.load_state_dict(torch.load(model_weights[i]))\n#         model.eval()\n#         vit_models.append(model)\n        \n# test_data = ImageFolder(test, (518, 518))\n# test_loader = DataLoader(test_data, batch_size=32, \n#                 pin_memory=False, num_workers=4, drop_last=False)\n# result_3 = {}\n# with torch.no_grad():\n#     for batch_idx, (spec_imgs, eeg_imgs, raw_50s_imgs, raw_10s_l_imgs, raw_10s_c_imgs, raw_10s_r_imgs, eeg_ids) in enumerate(test_loader):   \n#         spec_imgs = spec_imgs.to(device).float()\n#         eeg_imgs = eeg_imgs.to(device).float()\n#         raw_50s_imgs = raw_50s_imgs.to(device).float()\n#         raw_10s_l_imgs = raw_10s_l_imgs.to(device).float()\n#         raw_10s_c_imgs = raw_10s_c_imgs.to(device).float()\n#         raw_10s_r_imgs = raw_10s_r_imgs.to(device).float()\n#         ensemble_probs = torch.zeros((spec_imgs.shape[0], 6)).to(device)\n#         for model in vit_models:\n#             model.eval()\n#             logits_l = model(spec_imgs, eeg_imgs, raw_50s_imgs, raw_10s_l_imgs)\n#             logits_c = model(spec_imgs, eeg_imgs, raw_50s_imgs, raw_10s_c_imgs)\n#             logits_r = model(spec_imgs, eeg_imgs, raw_50s_imgs, raw_10s_r_imgs)\n#             probs_l = logits_l.softmax(dim=1)\n#             probs_c = logits_c.softmax(dim=1)\n#             probs_r = logits_r.softmax(dim=1)\n#             for j in range(len(eeg_ids)):\n#                 print(\"eeg id: {}, {}, {}, {}\".format(eeg_ids[j], probs_l[j], probs_c[j], probs_r[j]))\n#             ensemble_probs += (probs_l + probs_c + probs_r)/3\n#         ensemble_probs /= len(vit_models)\n#         ensemble_probs = ensemble_probs.detach().cpu().numpy()\n#         for j in range(len(eeg_ids)):\n#             eeg_id = eeg_ids[j]\n#             if eeg_id not in result_3.keys():\n#                 result_3[eeg_id] = np.array([0.0, 0.0, 0.0, 0.0, 0.0, 0.0])\n#             result_3[eeg_id] += ensemble_probs[j]\n#             print(\"eeg id: {}, {}\".format(eeg_id, ensemble_probs[j]))","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# for model in vit_models:\n#     del model\n# torch.cuda.empty_cache()\n# gc.collect()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# vit_models = []\n# model_weights = [\"/kaggle/input/hms-bestlb-vitbase/fold_0_exp_7_bestlb.pth\",\n#                 \"/kaggle/input/hms-bestlb-vitbase/fold_1_exp_7_bestlb.pth\",\n#                 \"/kaggle/input/hms-bestlb-vitbase/fold_2_exp_7_bestlb.pth\",\n#                 \"/kaggle/input/hms-bestlb-vitbase/fold_3_exp_7_bestlb.pth\",\n#                 \"/kaggle/input/hms-bestlb-vitbase/fold_4_exp_7_bestlb.pth\"]\n# model_types = [\"vit_base\", \"vit_base\", \"vit_base\", \"vit_base\", \"vit_base\"]\n# device = \"cuda:0\"\n# for i in range(len(model_types)):\n#     if model_types[i] == \"vit_base\":\n#         model = Net(\"vit_base_patch14_reg4_dinov2.lvd142m\", device).to(device)\n#         model.load_state_dict(torch.load(model_weights[i]))\n#         model.eval()\n#         vit_models.append(model)\n        \n# test_data = ImageFolder(test, (518, 518))\n# test_loader = DataLoader(test_data, batch_size=32, \n#                 pin_memory=False, num_workers=4, drop_last=False)\n# result_3 = {}\n# with torch.no_grad():\n#     for batch_idx, (spec_imgs, eeg_imgs, raw_50s_imgs, raw_10s_l_imgs, raw_10s_c_imgs, raw_10s_r_imgs, eeg_ids) in enumerate(test_loader):   \n#         spec_imgs = spec_imgs.to(device).float()\n#         eeg_imgs = eeg_imgs.to(device).float()\n#         raw_50s_imgs = raw_50s_imgs.to(device).float()\n#         raw_10s_l_imgs = raw_10s_l_imgs.to(device).float()\n#         raw_10s_c_imgs = raw_10s_c_imgs.to(device).float()\n#         raw_10s_r_imgs = raw_10s_r_imgs.to(device).float()\n#         ensemble_probs = torch.zeros((spec_imgs.shape[0], 6)).to(device)\n#         for model in vit_models:\n#             model.eval()\n#             logits_l = model(spec_imgs, eeg_imgs, raw_50s_imgs, raw_10s_l_imgs)\n#             logits_c = model(spec_imgs, eeg_imgs, raw_50s_imgs, raw_10s_c_imgs)\n#             logits_r = model(spec_imgs, eeg_imgs, raw_50s_imgs, raw_10s_r_imgs)\n#             probs_l = logits_l.softmax(dim=1)\n#             probs_c = logits_c.softmax(dim=1)\n#             probs_r = logits_r.softmax(dim=1)\n#             for j in range(len(eeg_ids)):\n#                 print(\"eeg id: {}, {}, {}, {}\".format(eeg_ids[j], probs_l[j], probs_c[j], probs_r[j]))\n#             ensemble_probs += (probs_l + probs_c + probs_r)/3\n#         ensemble_probs /= len(vit_models)\n#         ensemble_probs = ensemble_probs.detach().cpu().numpy()\n#         for j in range(len(eeg_ids)):\n#             eeg_id = eeg_ids[j]\n#             if eeg_id not in result_3.keys():\n#                 result_3[eeg_id] = np.array([0.0, 0.0, 0.0, 0.0, 0.0, 0.0])\n#             result_3[eeg_id] += ensemble_probs[j]\n#             print(\"eeg id: {}, {}\".format(eeg_id, ensemble_probs[j]))","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# for model in vit_models:\n#     del model\n# torch.cuda.empty_cache()\n# gc.collect()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_rows = []\nfor eeg_id in result_7.keys():\n#     r = (2*result_8[eeg_id] + result_7[eeg_id] + result_3[eeg_id] + result_4[eeg_id] + 2*result_5[eeg_id])/7\n    r = result_7[eeg_id]\n    row = [eeg_id, r[0], r[1], r[2], r[3], r[4], r[5]]\n    sub_rows.append(row)\nsub_rows = np.array(sub_rows)\ndf = pd.DataFrame(sub_rows, columns=[\"eeg_id\", \"seizure_vote\", \"lpd_vote\", \"gpd_vote\", \"lrda_vote\", \"grda_vote\", \"other_vote\"])\ndf.to_csv(\"submission.csv\",index=False)\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2024-01-29T00:06:32.285676Z","iopub.execute_input":"2024-01-29T00:06:32.285992Z","iopub.status.idle":"2024-01-29T00:06:32.307971Z","shell.execute_reply.started":"2024-01-29T00:06:32.285963Z","shell.execute_reply":"2024-01-29T00:06:32.30696Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if DEBUG == False:\n    !rm -rf /kaggle/working/spec_spectrograms/\n    !rm -rf /kaggle/working/eeg_spectrograms/\n    !rm -rf /kaggle/working/eeg_50s_raws/\n    !rm -rf /kaggle/working/eeg_10s_raws/\n    !rm -rf /kaggle/working/squeezeformer","metadata":{"execution":{"iopub.status.busy":"2024-01-29T00:06:32.31128Z","iopub.execute_input":"2024-01-29T00:06:32.311607Z","iopub.status.idle":"2024-01-29T00:06:37.255263Z","shell.execute_reply.started":"2024-01-29T00:06:32.311581Z","shell.execute_reply":"2024-01-29T00:06:37.25385Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}