{"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":32724,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":27393}],"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 = 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(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 = []\nmodel_weights = [\"//kaggle/input/hms_multi_modals/pytorch/vit_base/1/hms_multi_modals.pth\"]\nmodel_types = [\"vit_base\"]\ndevice = \"cuda:0\"\nfor 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        \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_5 = {}\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\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\ntorch.cuda.empty_cache()\ngc.collect()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_rows = []\nfor eeg_id in result_5.keys():\n    r = result_5[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":[]}]}