{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"},{"sourceId":7928297,"sourceType":"datasetVersion","datasetId":4659732,"isSourceIdPinned":true},{"sourceId":7970005,"sourceType":"datasetVersion","datasetId":4689524,"isSourceIdPinned":true},{"sourceId":7990788,"sourceType":"datasetVersion","datasetId":4515496,"isSourceIdPinned":true},{"sourceId":8007490,"sourceType":"datasetVersion","datasetId":4476729},{"sourceId":8015911,"sourceType":"datasetVersion","datasetId":4469088,"isSourceIdPinned":true},{"sourceId":8036762,"sourceType":"datasetVersion","datasetId":4564661,"isSourceIdPinned":true},{"sourceId":8041644,"sourceType":"datasetVersion","datasetId":4679791,"isSourceIdPinned":true},{"sourceId":8041978,"sourceType":"datasetVersion","datasetId":4716292},{"sourceId":7919067,"sourceType":"datasetVersion","datasetId":4543613,"isSourceIdPinned":true}],"dockerImageVersionId":30648,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"#import\n\nimport random\nimport cv2\nimport json\nimport numpy as np\nimport copy\nimport pandas as pd\nimport torch\nimport gc\n\nimport albumentations as A\nimport os\nimport librosa\nimport pickle\nimport timm\nfrom tqdm import tqdm\n\nimport mne\n\nimport torch\nimport torchaudio\nimport torch.nn as nn\nfrom torch.utils.data import DataLoader\nfrom scipy.signal import butter, lfilter","metadata":{"execution":{"iopub.status.busy":"2024-04-07T11:49:27.452425Z","iopub.execute_input":"2024-04-07T11:49:27.453385Z","iopub.status.idle":"2024-04-07T11:49:27.459125Z","shell.execute_reply.started":"2024-04-07T11:49:27.453347Z","shell.execute_reply":"2024-04-07T11:49:27.458081Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# !cp /kaggle/input/torchvideo/r50.py .\n!cp /kaggle/input/torchvideo/x3d.py .\n!cp /kaggle/input/torchvideo/hgnet.py .\n# from r50 import create_r2plus1d\nfrom x3d import create_x3d\nfrom hgnet import hgnetv2_b5","metadata":{"execution":{"iopub.status.busy":"2024-04-07T11:49:28.742038Z","iopub.execute_input":"2024-04-07T11:49:28.742408Z","iopub.status.idle":"2024-04-07T11:49:30.714109Z","shell.execute_reply.started":"2024-04-07T11:49:28.742378Z","shell.execute_reply":"2024-04-07T11:49:30.712659Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CFG={\n    'batch_size':32,\n    'num_worker':4,\n    'data':'/kaggle/input/hms-harmful-brain-activity-classification/test.csv',\n    'weights_spec':'/kaggle/input/hms-kaggle-spec',\n    'weights_one_image':'/kaggle/input/hms-oneimage',\n    'weights_x3d':'/kaggle/input/hms-x3d',\n    \n    \n    'weights_eeg_raw':'/kaggle/input/hms-eeg-raw',\n    \n    'weights_doublehead':'/kaggle/input/hms-double-head',\n    'weights_doubleheadbutterfilter':'/kaggle/input/hms-doublehead-butter-filter',\n    'weights_raw_eeg_butter_filter':'/kaggle/input/hms-raw-eeg-butter-filter',\n    'weights_hgnet':'/kaggle/input/hms-hgnet',\n    'flip':True\n}","metadata":{"execution":{"iopub.status.busy":"2024-04-07T11:49:31.662911Z","iopub.execute_input":"2024-04-07T11:49:31.663277Z","iopub.status.idle":"2024-04-07T11:49:31.669180Z","shell.execute_reply.started":"2024-04-07T11:49:31.663250Z","shell.execute_reply":"2024-04-07T11:49:31.668242Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CFG['weights_spec']=[os.path.join(CFG['weights_spec'],x) for x in sorted(os.listdir(CFG['weights_spec']))]\nCFG['weights_x3d']=[os.path.join(CFG['weights_x3d'],x) for x in sorted(os.listdir(CFG['weights_x3d']))]\nCFG['weights_eeg_raw']=[os.path.join(CFG['weights_eeg_raw'],x) for x in sorted(os.listdir(CFG['weights_eeg_raw']))]\nCFG['weights_raw_eeg_butter_filter']=[os.path.join(CFG['weights_raw_eeg_butter_filter'],x) for x in sorted(os.listdir(CFG['weights_raw_eeg_butter_filter']))]\nCFG['weights_doubleheadbutterfilter']=[os.path.join(CFG['weights_doubleheadbutterfilter'],x) for x in sorted(os.listdir(CFG['weights_doubleheadbutterfilter']))]\nCFG['weights_hgnet']=[os.path.join(CFG['weights_hgnet'],x) for x in sorted(os.listdir(CFG['weights_hgnet']))]\nCFG['weights_one_image']=[os.path.join(CFG['weights_one_image'],x) for x in sorted(os.listdir(CFG['weights_one_image']))]\n\nCFG","metadata":{"execution":{"iopub.status.busy":"2024-04-07T11:49:32.953474Z","iopub.execute_input":"2024-04-07T11:49:32.954270Z","iopub.status.idle":"2024-04-07T11:49:32.971943Z","shell.execute_reply.started":"2024-04-07T11:49:32.954239Z","shell.execute_reply":"2024-04-07T11:49:32.971083Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#dataiter\nclass AlaskaDataIter():\n    def __init__(self, df,\n                 training_flag=False,shuffle=False,\n                 use_spec=False,\n                 use_eeg=False,\n                 use_mix=False,\n                 ll=0,rr=20,\n                 flip=False,\n                 use_mne_filter=True,\n                 use_18_lead=False):\n        \n        self.flip_eeg=flip\n        self.ll=ll\n        self.rr=rr\n        self.use_18_lead=use_18_lead\n        print(self.ll,self.rr, 'with mne filter:', use_mne_filter,'use 18 lead:',use_18_lead)\n        \n        \n        self.training_flag = training_flag\n        self.shuffle = shuffle\n\n        self.raw_data_set_size = None     ##decided by self.parse_file\n\n\n        self.df=df\n        \n\n        TARS = {'Seizure': 0, 'LPD': 1, 'GPD': 2, 'LRDA': 3, 'GRDA': 4, 'Other': 5}\n        self.TARS2 = {x: y for y, x in TARS.items()}\n\n\n        self.eeg_nms=['Fp1', 'F3', 'C3', 'P3', 'F7', 'T3', 'T5', 'O1', 'Fz', 'Cz', 'Pz',\n       'Fp2', 'F4', 'C4', 'P4', 'F8', 'T4', 'T6', 'O2', 'EKG']\n\n        self.LL = ['Fp1', 'F7', 'T3', 'T5', 'O1']\n\n        self.RL = ['Fp2', 'F8', 'T4', 'T6', 'O2']\n\n        self.LP = ['Fp1', 'F3', 'C3', 'P3', 'O1']\n\n        self.RP = ['Fp2', 'F4', 'C4', 'P4', 'O2']\n\n        self.mid = ['Fz', 'Cz', 'Pz']\n        self.leads_dict = {value: index for index, value in enumerate(self.eeg_nms)}\n\n        self.use_eeg = use_eeg\n        self.use_spec = use_spec\n        self.use_mix = use_mix\n        self.use_mne_filter=use_mne_filter\n        \n        \n        \n    def __getitem__(self, item):\n\n        return self.single_map_func(self.df.iloc[item], self.training_flag)\n\n    def __len__(self):\n\n        return len(self.df)\n\n    def brain_lead(self, waves):\n        waves = copy.deepcopy(waves)\n        brain_leads = [self.LL, self.RL, self.LP, self.RP]\n\n        leads = []\n\n        for chain in brain_leads:\n            for i in range(len(chain) - 1):\n                tmp_lead = waves[self.leads_dict[chain[i]]] - waves[self.leads_dict[chain[i + 1]]]\n                leads.append(tmp_lead)\n\n        data = np.concatenate([leads], axis=0)\n        \n        return data\n    def mirror_spec(self, data):\n\n        # index_choice = [[0, 1, 3, 2], [0, 1, 2, 3], [1, 0, 2, 3], [1, 0, 3, 2]]\n        indx = [1, 0, 3, 2]\n        return data[..., indx]\n\n    def mirror_eeg(self, data):\n\n        indx1 = [0, 1, 2, 3, 4, 5, 6, 7]\n        indx2 = [11, 12, 13, 14, 15, 16, 17, 18]\n\n        data[indx1, ...], data[indx2, ...] = data[indx2, ...], data[indx1, ...]\n\n        return data\n    def butter_bandpass(self,lowcut, highcut, fs, order=5):\n        return butter(order, [lowcut, highcut], fs=fs, btype=\"band\")\n\n    def butter_bandpass_filter(self,data, lowcut, highcut, fs, order=5):\n        b, a = self.butter_bandpass(lowcut, highcut, fs, order=order)\n        y = lfilter(b, a, data)\n        return y\n    def get_eeg(self, dp, is_training,flip=False):\n        eeg_path = '/kaggle/input/hms-harmful-brain-activity-classification/test_eegs/%s.parquet' % (dp['eeg_id'])\n        eeg = pd.read_parquet(eeg_path)\n\n        \n        offset = 0\n        eeg = eeg.iloc[int(offset * 200):int(offset * 200) + 10000]\n\n        waves = eeg.values\n\n        waves = np.transpose(waves, axes=[1, 0])\n\n        for i in range(waves.shape[0]):\n            m = np.nanmean(waves[i])\n            if np.isnan(waves[i]).mean() < 1:\n                waves[i] = np.nan_to_num(waves[i], nan=m)\n            else:\n                waves[i] = 0\n        \n        if flip:\n            waves=self.mirror_eeg(waves)\n        waves = self.brain_lead(waves)\n        waves = np.array(waves, dtype=np.float64)\n\n        waves = np.clip(waves, -1024, 1024)\n        if self.use_mne_filter:\n            waves = mne.filter.filter_data(waves, 200, self.ll, self.rr, verbose=False)\n        else:\n            waves = self.butter_bandpass_filter(waves,0.5,20,200,2)\n        #\n\n\n        return waves\n    def get_spec(self, dp, is_training,flip=False):\n        spec_path = '/kaggle/input/hms-harmful-brain-activity-classification/test_spectrograms/%s.parquet' % (dp['spectrogram_id'])\n        spec = pd.read_parquet(spec_path)\n\n        spec=spec.values[:,1:]\n\n        images=[]\n        r = 0\n\n        for region in range(4):\n            img = spec[r:r + 300,\n                  region * 100:(region + 1) * 100].T\n            img = np.clip(img, np.exp(-4), np.exp(8))\n            img = np.log(img)\n\n            # STANDARDIZE PER IMAGE\n            # ep = 1e-6\n            # m = np.nanmean(img.flatten())\n            # s = np.nanstd(img.flatten())\n            # img = (img - m) / (s + ep)\n            img = np.nan_to_num(img, nan=0.0)\n\n            images.append(img)\n\n        images=np.stack(images,-1)\n        \n        if flip:\n            images = self.mirror_spec(images)\n        data = np.transpose(images, [2,0,1])\n\n\n        return data\n    \n    def get_mix(self, dp, is_training,flip=False):\n        \n        eeg_path = '/kaggle/input/hms-harmful-brain-activity-classification/test_eegs/%s.parquet' % (dp['eeg_id'])\n        eeg = pd.read_parquet(eeg_path)\n        \n        offset = 0\n\n        eeg = eeg.iloc[int(offset * 200):int(offset * 200) + 10000]\n\n        waves = eeg.values\n\n        waves = np.transpose(waves, axes=[1, 0])\n        for i in range(waves.shape[0]):\n            m = np.nanmean(waves[i])\n            if np.isnan(waves[i]).mean() < 1:\n                waves[i] = np.nan_to_num(waves[i], nan=m)\n            else:\n                waves[i] = 0\n\n        if flip:\n            waves = self.mirror_eeg(waves)\n        \n        waves = self.brain_lead(waves)\n        waves = np.array(waves, dtype=np.float64)\n\n        waves = np.clip(waves, -1024, 1024)\n\n        waves = mne.filter.filter_data(waves, 200, self.ll, self.rr, verbose=False)\n        #\n        # get spec\n        r = 0\n        spec_path = '/kaggle/input/hms-harmful-brain-activity-classification/test_spectrograms/%s.parquet' % (dp['spectrogram_id'])\n        spec = pd.read_parquet(spec_path)\n        spec=spec.values[:,1:]\n        \n\n        images = []\n        for region in range(4):\n            img = spec[r:r + 300, region * 100:(region + 1) * 100].T\n            img = np.clip(img, np.exp(-4), np.exp(8))\n            img = np.log(img)\n\n            # STANDARDIZE PER IMAGE\n            # ep = 1e-6\n            # m = np.nanmean(img.flatten())\n            # s = np.nanstd(img.flatten())\n            # img = (img - m) / (s + ep)\n            img = np.nan_to_num(img, nan=0.0)\n\n            images.append(img)\n\n        images = np.stack(images, -1)\n\n        if flip:\n            images = self.mirror_spec(images)\n\n        images = np.transpose(images, [2, 0, 1])\n        \n        \n        return waves, images\n    \n    def single_map_func(self, dp, is_training):\n        \"\"\"Data augmentation function.\"\"\"\n        ####customed here\n\n        \n        if self.use_eeg:\n            \n            data=self.get_eeg(dp,is_training,self.flip_eeg)\n        elif self.use_spec:\n            data=self.get_spec(dp,is_training,self.flip_eeg)\n        elif self.use_mix:\n            data,spec =self.get_mix(dp,is_training)\n            \n            return data.astype(np.float32),spec.astype(np.float32)\n        \n        return data.astype(np.float32)","metadata":{"execution":{"iopub.status.busy":"2024-04-07T11:49:34.542915Z","iopub.execute_input":"2024-04-07T11:49:34.543247Z","iopub.status.idle":"2024-04-07T11:49:34.582708Z","shell.execute_reply.started":"2024-04-07T11:49:34.543221Z","shell.execute_reply":"2024-04-07T11:49:34.581830Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n\n\nclass NetSpec(nn.Module):\n    def __init__(self, num_classes=1):\n        super().__init__()\n\n        \n\n        self.model = timm.create_model('efficientnet_b5',\n                                       pretrained=False,\n                                       in_chans=3)\n\n        self.fc = nn.Linear(2048, 6, bias=True)\n        # self.wave_encoder = ResNet1D(inp_ch=42, block=ResNetBlock1D, layers=[2, 2, 4, 2], num_classes=1)\n\n        # weight_init(self.fc)\n        self.dropout=nn.Dropout(0.5)\n\n        self.avg=nn.AdaptiveAvgPool2d(1)\n\n    def forward(self, x):\n        # do preprocess\n        bs = x.size(0)\n\n        # wave_fm = self.wave_encoder(x)  # 1x250\n        # wave_fm = wave_fm.unsqueeze(2)\n        #\n        x1 = [x[:,  i:i + 1,:, :] for i in range(4)]\n        x1 = torch.cat(x1,dim=2)\n        x= torch.cat([x1,x1,x1],dim=1)\n            \n\n        x = self.model.forward_features(x)\n        x = self.avg(x)\n\n        x = x.view(bs, -1)\n        x = self.dropout(x)\n\n        x = self.fc(x)\n        \n        x =torch.softmax(x,-1)\n        \n        ans=x\n        return ans\n\n\n","metadata":{"execution":{"iopub.status.busy":"2024-04-07T11:49:35.244335Z","iopub.execute_input":"2024-04-07T11:49:35.245327Z","iopub.status.idle":"2024-04-07T11:49:35.255193Z","shell.execute_reply.started":"2024-04-07T11:49:35.245282Z","shell.execute_reply":"2024-04-07T11:49:35.254356Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n\nclass Transform(nn.Module):\n    def __init__(self, ):\n        super().__init__()\n\n\n        self.wave_transform = torchaudio.transforms.Spectrogram(n_fft=512,\n                                                                win_length=128,\n                                                                hop_length=50,\n                                                                power=1)\n    def forward(self, x):\n        bs = x.size(0)\n        image=self.wave_transform(x)\n\n        n, c, h, w = image.size()\n        image = image[:, :, :int(20 / 100 * h + 2), :]\n\n        image = torch.clip(image, min=0, max=10000)/1000\n\n        image = torch.reshape(image, shape=[n, 2, -1, w])\n\n        x1 = image[:, 0:1, ...]\n        x2 = image[:, 1:2, ...]\n\n        image = torch.cat([x1, x2], dim=-1)\n\n        image = torch.cat([image, image, image], dim=1)\n\n        return image\n\n\n\n\n\nclass NetOneImage(nn.Module):\n    def __init__(self, num_classes=1):\n        super().__init__()\n\n        self.preprocess = Transform()\n\n        self.model=hgnetv2_b5()\n\n        self.fc = nn.Linear(2048, 6, bias=True)\n        # self.wave_encoder = ResNet1D(inp_ch=42, block=ResNetBlock1D, layers=[2, 2, 4, 2], num_classes=1)\n\n        # weight_init(self.fc)\n        self.dropout=nn.Dropout(0.5)\n\n        self.avg=nn.AdaptiveAvgPool2d(1)\n\n\n    def forward(self, x):\n        # do preprocess\n        bs = x.size(0)\n\n        x= self.preprocess(x)\n            \n        x = self.model.forward_features(x)\n        x = self.avg(x)\n        x = x.view(bs, -1)\n        x = self.dropout(x)\n\n        x = self.fc(x)\n        x =torch.softmax(x,-1)\n        \n        ans=x\n            \n        return ans\n","metadata":{"execution":{"iopub.status.busy":"2024-04-07T11:49:37.143264Z","iopub.execute_input":"2024-04-07T11:49:37.143651Z","iopub.status.idle":"2024-04-07T11:49:37.156219Z","shell.execute_reply.started":"2024-04-07T11:49:37.143622Z","shell.execute_reply":"2024-04-07T11:49:37.155318Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Transform50s(nn.Module):\n    def __init__(self, ):\n        super().__init__()\n\n\n        self.wave_transform = torchaudio.transforms.Spectrogram(n_fft=512,\n                                                                win_length=128,\n                                                                hop_length=50,\n                                                                power=1)\n\n    def forward(self, x):\n        bs = x.size(0)\n\n        image=self.wave_transform(x)\n        image = torch.clip(image, min=0,max=10000)/1000\n\n        n, c, h, w = image.size()\n        image = image[:, :, :int(20 / 100 * h + 10), :]\n        return image\n\n\nclass Transform10s(nn.Module):\n    def __init__(self, ):\n        super().__init__()\n\n\n        self.wave_transform = torchaudio.transforms.Spectrogram(n_fft=512,\n                                                                win_length=128,\n                                                                hop_length=10,\n                                                                power=1)\n\n    def forward(self, x):\n        bs = x.size(0)\n\n        image=self.wave_transform(x)\n        image = torch.clip(image, min=0,max=10000)/1000\n        n, c, h, w = image.size()\n        image = image[:, :, :int(20 / 100 * h + 10), :]\n        return image\n\nclass Modelx3d(nn.Module):\n    def __init__(self):\n        super().__init__()\n        model_name = \"x3d_l\"\n        self.net = create_x3d(input_clip_length=16,\n        input_crop_size=312,\n        depth_factor=5.0,)\n        \n        # self.net.blocks[5]=nn.Identity()\n        # self.net.avgpool = nn.Identity()\n        self.net.blocks[5].dropout=nn.Identity()\n        self.net.blocks[5].proj = nn.Identity()\n        self.net.blocks[5].activation = nn.Identity()\n        self.net.blocks[5].output_pool = nn.Identity()\n    def forward(self, x):\n\n        x = self.net(x)\n\n        return x\n\n\n\nclass Netx3d(nn.Module):\n    def __init__(self, num_classes=1):\n        super().__init__()\n\n        self.preprocess50s = Transform50s()\n        self.preprocess10s = Transform10s()\n\n        self.model = Modelx3d()\n\n        self.pool=nn.AdaptiveAvgPool3d(1)\n        self.fc = nn.Linear(2048, 6, bias=True)\n    \n    def forward(self, eeg):\n        \n        \n        bs = eeg.size(0)\n\n        eeg_50s=eeg\n        eeg_10s=eeg[:,:,4000:6000]\n        x_50 = self.preprocess50s(eeg_50s)\n        x_10 = self.preprocess10s(eeg_10s)\n        x= torch.cat([x_10,x_50],dim=1)\n\n\n        x = torch.unsqueeze(x,dim=1)\n\n        x = torch.cat([x,x,x],dim=1)\n            \n        x = self.model(x)\n\n#         x = self.pool(x)\n        x = x.view(bs, -1)\n            \n        \n        x= self.fc(x)\n        \n        x =torch.softmax(x,-1)\n        ans=x\n            \n        return ans","metadata":{"execution":{"iopub.status.busy":"2024-04-07T11:49:37.951845Z","iopub.execute_input":"2024-04-07T11:49:37.952193Z","iopub.status.idle":"2024-04-07T11:49:37.970239Z","shell.execute_reply.started":"2024-04-07T11:49:37.952165Z","shell.execute_reply":"2024-04-07T11:49:37.969401Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n\nclass Net1d(nn.Module):\n    def __init__(self,):\n        super(Net1d, self).__init__()\n        self.model=timm.create_model('efficientnet_b5', pretrained=False, in_chans=3)\n        self.pool=nn.AdaptiveAvgPool2d(1)\n        self.fc=nn.Linear(2048,out_features=6,bias=True)\n        self.dropout=nn.Dropout(p=0.5)\n\n    def extract_features(self, x):\n        feature1=self.model.forward_features(x)\n        return feature1\n    def forward(self, x):\n        \n        bs = x.size(0)\n        reshaped_tensor = x.view(bs,16,1000, 10)\n        reshaped_and_permuted_tensor = reshaped_tensor.permute(0,1,3,2)\n        reshaped_and_permuted_tensor= reshaped_and_permuted_tensor.reshape(bs,16*10,1000)\n        x=torch.unsqueeze(reshaped_and_permuted_tensor,dim=1)\n        x=torch.cat([x,x,x],dim=1)\n        bs=x.size(0)\n\n        x = self.extract_features(x)\n\n        # print(x.size())\n        x = self.pool(x)\n        x = x.view(bs, -1)\n        \n        x =self.dropout(x)\n        x = self.fc(x)\n        x =torch.softmax(x,-1)\n        \n        return x\n\n","metadata":{"execution":{"iopub.status.busy":"2024-04-07T11:49:39.747054Z","iopub.execute_input":"2024-04-07T11:49:39.747415Z","iopub.status.idle":"2024-04-07T11:49:39.757093Z","shell.execute_reply.started":"2024-04-07T11:49:39.747385Z","shell.execute_reply":"2024-04-07T11:49:39.756220Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n\nclass Net1dhgnet(nn.Module):\n    def __init__(self,):\n        super(Net1dhgnet, self).__init__()\n        self.model=hgnetv2_b5()\n        self.pool=nn.AdaptiveAvgPool2d(1)\n        self.fc=nn.Linear(2048,out_features=6,bias=True)\n        self.dropout=nn.Dropout(p=0.5)\n\n    def extract_features(self, x):\n        feature1=self.model.forward_features(x)\n        return feature1\n    def forward(self, x):\n        \n        bs = x.size(0)\n        reshaped_tensor = x.view(bs,16,1000, 10)\n        reshaped_and_permuted_tensor = reshaped_tensor.permute(0,1,3,2)\n        reshaped_and_permuted_tensor= reshaped_and_permuted_tensor.reshape(bs,16*10,1000)\n        x=torch.unsqueeze(reshaped_and_permuted_tensor,dim=1)\n        x=torch.cat([x,x,x],dim=1)\n        bs=x.size(0)\n\n        x = self.extract_features(x)\n\n        # print(x.size())\n        x = self.pool(x)\n        x = x.view(bs, -1)\n        \n        x =self.dropout(x)\n        x = self.fc(x)\n        x =torch.softmax(x,-1)\n        \n        return x\n\n","metadata":{"execution":{"iopub.status.busy":"2024-04-07T11:49:40.543836Z","iopub.execute_input":"2024-04-07T11:49:40.544648Z","iopub.status.idle":"2024-04-07T11:49:40.554417Z","shell.execute_reply.started":"2024-04-07T11:49:40.544617Z","shell.execute_reply":"2024-04-07T11:49:40.553459Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n\nclass Transformdoublehead(nn.Module):\n    def __init__(self, ):\n        super().__init__()\n\n        self.wave_transform = torchaudio.transforms.Spectrogram(n_fft=512,\n                                                                hop_length=50,\n                                                                power=1)\n        \n\n    def forward(self, x):\n        bs = x.size(0)\n\n        image = self.wave_transform(x)\n        # image = self.am2db(image)\n        image = torch.log10(image)\n        image = torch.clip(image, min=0)\n\n        n, c, h, w = image.size()\n\n        ## inference use 0-20hz filter,\n        image = image[:, :, :int(40 / 100 * h+10), :]\n\n        \n\n        return image\n\nclass Modeleeg(nn.Module):\n    def __init__(self,):\n        super(Modeleeg, self).__init__()\n\n        self.model=timm.create_model('efficientnet_b5',\n                                            pretrained=False,\n                                            in_chans=3,)\n\n\n\n        self.pool=nn.AdaptiveAvgPool2d(1)\n        \n\n    def extract_features(self, x):\n\n\n        x=self.model.forward_features(x)\n\n        return x\n\n    def forward(self, x):\n        bs = x.size(0)\n        reshaped_tensor = x.view(bs,16,1000, 10)\n\n        reshaped_and_permuted_tensor = reshaped_tensor.permute(0,1,3,2)\n\n        reshaped_and_permuted_tensor= reshaped_and_permuted_tensor.reshape(bs,16*10,1000)\n\n        x=torch.unsqueeze(reshaped_and_permuted_tensor,dim=1)\n\n        x=torch.cat([x,x,x],dim=1)\n        bs=x.size(0)\n\n\n        x = self.extract_features(x)\n\n        x = self.pool(x)\n        x = x.view(bs,-1)\n\n        return x\n\nclass Modelspec(nn.Module):\n    def __init__(self, num_classes=1):\n        super().__init__()\n\n        super().__init__()\n        model_name = \"x3d_l\"\n        self.net = create_x3d(input_clip_length=16,\n        input_crop_size=312,\n        depth_factor=5.0,)\n        \n        # self.net.blocks[5]=nn.Identity()\n        # self.net.avgpool = nn.Identity()\n        self.net.blocks[5].dropout=nn.Identity()\n        self.net.blocks[5].proj = nn.Identity()\n        self.net.blocks[5].activation = nn.Identity()\n        self.net.blocks[5].output_pool = nn.Identity()\n\n        self.avg = nn.AdaptiveAvgPool2d(1)\n\n\n\n    def forward(self, x):\n        x = torch.unsqueeze(x, dim=1)\n        x = torch.cat([x, x, x], dim=1)\n        x = self.net(x)\n        # do preprocess\n        bs = x.size(0)\n\n        x = x.view(bs, -1)\n\n        return x\n    \nclass Netdoublehead(nn.Module):\n    def __init__(self,):\n        super(Netdoublehead, self).__init__()\n        self.transform=Transformdoublehead()\n        self.model_wave = Modeleeg()\n\n        self.model_spec = Modelspec()\n\n        self.pool = nn.AdaptiveAvgPool3d(1)\n        self.fc = nn.Linear(2048*2, 6, bias=True)\n\n        self.droup = nn.Dropout(0.5)\n\n    \n    def forward(self, eeg):\n        \n        bs = eeg.size(0)\n\n        eeg_spec=self.transform(eeg)\n\n        x = self.model_wave(eeg)\n        y = self.model_spec(eeg_spec)\n#         x = (x + y) / 2\n        x = torch.cat([x,y],dim=1)\n\n        x = self.droup(x)\n        x = self.fc(x)\n        x =torch.softmax(x,-1)\n        \n        return x\n\n","metadata":{"execution":{"iopub.status.busy":"2024-04-07T11:49:41.043098Z","iopub.execute_input":"2024-04-07T11:49:41.043460Z","iopub.status.idle":"2024-04-07T11:49:41.063450Z","shell.execute_reply.started":"2024-04-07T11:49:41.043431Z","shell.execute_reply":"2024-04-07T11:49:41.062443Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n\nclass Transformdoubleheadbutterfilter(nn.Module):\n    def __init__(self, ):\n        super().__init__()\n\n        self.wave_transform = torchaudio.transforms.Spectrogram(n_fft=512,\n                                                                hop_length=50,\n                                                                power=1)\n        \n\n    def forward(self, x):\n        bs = x.size(0)\n\n        image = self.wave_transform(x)\n        # image = self.am2db(image)\n        image = torch.log10(image)\n        image = torch.clip(image, min=0)\n\n        n, c, h, w = image.size()\n\n        ## inference use 0-20hz filter,\n        image = image[:, :, :int(20 / 100 * h+5), :]\n\n        \n\n        return image\nclass Modeleegbutterfilter(nn.Module):\n    def __init__(self,):\n        super(Modeleegbutterfilter, self).__init__()\n\n        self.model=timm.create_model('efficientnet_b5',\n                                            pretrained=False,\n                                            in_chans=3,)\n\n\n\n        self.pool=nn.AdaptiveAvgPool2d(1)\n        self.fc=nn.Linear(2048,out_features=6,bias=True)\n\n    def extract_features(self, x):\n\n\n        x=self.model.forward_features(x)\n\n        return x\n\n    def forward(self, x):\n        bs = x.size(0)\n        reshaped_tensor = x.view(bs,16,1000, 10)\n\n        reshaped_and_permuted_tensor = reshaped_tensor.permute(0,1,3,2)\n\n        reshaped_and_permuted_tensor= reshaped_and_permuted_tensor.reshape(bs,16*10,1000)\n\n        x=torch.unsqueeze(reshaped_and_permuted_tensor,dim=1)\n\n        x=torch.cat([x,x,x],dim=1)\n        bs=x.size(0)\n\n\n        x = self.extract_features(x)\n\n        x = self.pool(x)\n        x = x.view(bs,-1)\n\n        return x\n\nclass Modelspecbutterfilter(nn.Module):\n    def __init__(self, num_classes=1):\n        super().__init__()\n\n        super().__init__()\n        model_name = \"x3d_l\"\n        self.net = create_x3d(input_clip_length=16,\n        input_crop_size=312,\n        depth_factor=5.0,)\n        \n        # self.net.blocks[5]=nn.Identity()\n        # self.net.avgpool = nn.Identity()\n        self.net.blocks[5].dropout=nn.Identity()\n        self.net.blocks[5].proj = nn.Identity()\n        self.net.blocks[5].activation = nn.Identity()\n        self.net.blocks[5].output_pool = nn.Identity()\n\n        self.avg = nn.AdaptiveAvgPool2d(1)\n\n\n\n    def forward(self, x):\n        x = torch.unsqueeze(x, dim=1)\n        x = torch.cat([x, x, x], dim=1)\n        x = self.net(x)\n        # do preprocess\n        bs = x.size(0)\n\n        x = x.view(bs, -1)\n\n        return x\n    \nclass Netdoubleheadbutterfilter(nn.Module):\n    def __init__(self,):\n        super(Netdoubleheadbutterfilter, self).__init__()\n        self.transform=Transformdoubleheadbutterfilter()\n        self.model_wave = Modeleegbutterfilter()\n\n        self.model_spec = Modelspecbutterfilter()\n\n        self.pool = nn.AdaptiveAvgPool3d(1)\n        self.fc = nn.Linear(2048*2, 6, bias=True)\n\n        self.droup = nn.Dropout(0.5)\n\n    \n    def forward(self, eeg):\n        \n        bs = eeg.size(0)\n\n        eeg_spec=self.transform(eeg)\n\n        x = self.model_wave(eeg)\n        y = self.model_spec(eeg_spec)\n#         x = (x + y) / 2\n        x = torch.cat([x,y],dim=1)\n\n        x = self.droup(x)\n        x = self.fc(x)\n        x =torch.softmax(x,-1)\n        \n        return x\n\n","metadata":{"execution":{"iopub.status.busy":"2024-04-07T11:49:41.543673Z","iopub.execute_input":"2024-04-07T11:49:41.544347Z","iopub.status.idle":"2024-04-07T11:49:41.564410Z","shell.execute_reply.started":"2024-04-07T11:49:41.544318Z","shell.execute_reply":"2024-04-07T11:49:41.563435Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def inference_function(test_loader, model, device,double_input=False):\n    model.eval()\n    \n    prediction_dict = {}\n    preds = []\n    with tqdm(test_loader, unit=\"test_batch\", desc='Inference') as tqdm_test_loader:\n        for step, X in enumerate(tqdm_test_loader):\n            \n            if double_input:\n                wave,spec=X\n                \n                wave = wave.to(device)\n                spec = spec.to(device)\n                \n                \n                batch_size = wave.size(0)\n                with torch.no_grad():\n                    y_preds = model(wave,spec)\n                    \n            else:\n                X = X.to(device)\n\n                batch_size = X.size(0)\n                with torch.no_grad():\n                    y_preds = model(X)\n\n            preds.append(y_preds.to('cpu').numpy()) \n                \n    prediction_dict[\"predictions\"] = np.concatenate(preds) \n    return prediction_dict","metadata":{"execution":{"iopub.status.busy":"2024-04-07T11:49:42.046578Z","iopub.execute_input":"2024-04-07T11:49:42.047153Z","iopub.status.idle":"2024-04-07T11:49:42.054934Z","shell.execute_reply.started":"2024-04-07T11:49:42.047126Z","shell.execute_reply":"2024-04-07T11:49:42.053990Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df=pd.read_csv(CFG['data'])\n\ntest_df.head(5)\n","metadata":{"execution":{"iopub.status.busy":"2024-04-07T11:49:42.642281Z","iopub.execute_input":"2024-04-07T11:49:42.643070Z","iopub.status.idle":"2024-04-07T11:49:42.657200Z","shell.execute_reply.started":"2024-04-07T11:49:42.643040Z","shell.execute_reply":"2024-04-07T11:49:42.656366Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def run_weight_spec():\n    print('infer with weights_spec')\n    device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n    predictions=[]\n    for model_weight in CFG['weights_spec']:\n\n        test_dataset = AlaskaDataIter(test_df, training_flag=False, shuffle=False,use_spec=True)\n        test_loader = DataLoader(test_dataset,\n                         CFG['batch_size'],\n                         num_workers=CFG['num_worker'],\n                         shuffle=False)\n\n        model = NetSpec()\n        state_dict = torch.load(model_weight, map_location=device)\n\n        model.load_state_dict(state_dict,strict=True)\n        model.to(device)\n        prediction_dict = inference_function(test_loader, model, device)\n        predictions.append(prediction_dict[\"predictions\"])\n\n        test_dataset = AlaskaDataIter(test_df, training_flag=False, shuffle=False,use_spec=True,flip=True)\n        test_loader = DataLoader(test_dataset,\n                         CFG['batch_size'],\n                         num_workers=CFG['num_worker'],\n                         shuffle=False)\n\n        model = NetSpec()\n        state_dict = torch.load(model_weight, map_location=device)\n\n        model.load_state_dict(state_dict,strict=True)\n        model.to(device)\n        prediction_dict = inference_function(test_loader, model, device)\n        predictions.append(prediction_dict[\"predictions\"])\n\n        torch.cuda.empty_cache()\n        gc.collect()\n    predictions = np.array(predictions)\n    predictions = np.mean(predictions, axis=0)\n    \n    return predictions","metadata":{"execution":{"iopub.status.busy":"2024-04-07T11:49:43.053337Z","iopub.execute_input":"2024-04-07T11:49:43.054031Z","iopub.status.idle":"2024-04-07T11:49:43.065195Z","shell.execute_reply.started":"2024-04-07T11:49:43.053992Z","shell.execute_reply":"2024-04-07T11:49:43.064168Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def run_weight_one_image():\n    print('infer with weights_one_image')\n    device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n    predictions=[]\n    for model_weight in CFG['weights_one_image']:\n        test_dataset = AlaskaDataIter(test_df, training_flag=False, shuffle=False,use_eeg=True,ll=0.5,rr=20)\n        test_loader = DataLoader(test_dataset,\n                         CFG['batch_size'],\n                         num_workers=CFG['num_worker'],\n                         shuffle=False)\n\n        model = NetOneImage()\n        state_dict = torch.load(model_weight, map_location=device)\n\n        model.load_state_dict(state_dict,strict=True)\n        model.to(device)\n        prediction_dict = inference_function(test_loader, model, device)\n        predictions.append(prediction_dict[\"predictions\"])\n        \n        test_dataset = AlaskaDataIter(test_df, training_flag=False, shuffle=False,use_eeg=True,flip=True,ll=0.5,rr=20)\n        test_loader = DataLoader(test_dataset,\n                         CFG['batch_size'],\n                         num_workers=CFG['num_worker'],\n                         shuffle=False)\n\n        model = NetOneImage()\n        state_dict = torch.load(model_weight, map_location=device)\n\n        model.load_state_dict(state_dict,strict=True)\n        model.to(device)\n        prediction_dict = inference_function(test_loader, model, device)\n        predictions.append(prediction_dict[\"predictions\"])\n\n        torch.cuda.empty_cache()\n        gc.collect()\n        \n        \n    predictions = np.array(predictions)\n    predictions = np.mean(predictions, axis=0)\n    \n    return predictions","metadata":{"execution":{"iopub.status.busy":"2024-04-07T11:49:44.043594Z","iopub.execute_input":"2024-04-07T11:49:44.044441Z","iopub.status.idle":"2024-04-07T11:49:44.057134Z","shell.execute_reply.started":"2024-04-07T11:49:44.044400Z","shell.execute_reply":"2024-04-07T11:49:44.055739Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def run_weight_eeg_raw():\n    device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n\n    print('infer with run_weight_eeg_raw')\n    predictions=[]\n    for model_weight in CFG['weights_eeg_raw']:\n        test_dataset = AlaskaDataIter(test_df, training_flag=False, shuffle=False,use_eeg=True,ll=0.5,rr=20)\n        test_loader = DataLoader(test_dataset,\n                         CFG['batch_size'],\n                         num_workers=CFG['num_worker'],\n                         shuffle=False)\n\n        model = Net1d()\n        state_dict = torch.load(model_weight, map_location=device)\n\n        model.load_state_dict(state_dict,strict=True)\n        model.to(device)\n        prediction_dict = inference_function(test_loader, model, device)\n        predictions.append(prediction_dict[\"predictions\"])\n\n\n        test_dataset = AlaskaDataIter(test_df, training_flag=False, shuffle=False,use_eeg=True,flip=True,ll=0.5,rr=20)\n        test_loader = DataLoader(test_dataset,\n                         CFG['batch_size'],\n                         num_workers=CFG['num_worker'],\n                         shuffle=False)\n\n        model = Net1d()\n        state_dict = torch.load(model_weight, map_location=device)\n\n        model.load_state_dict(state_dict,strict=True)\n        model.to(device)\n        prediction_dict = inference_function(test_loader, model, device)\n        predictions.append(prediction_dict[\"predictions\"])\n\n        torch.cuda.empty_cache()\n        gc.collect()\n    predictions = np.array(predictions)\n    predictions = np.mean(predictions, axis=0)\n    \n    return predictions","metadata":{"execution":{"iopub.status.busy":"2024-04-07T11:49:44.842628Z","iopub.execute_input":"2024-04-07T11:49:44.843357Z","iopub.status.idle":"2024-04-07T11:49:44.853383Z","shell.execute_reply.started":"2024-04-07T11:49:44.843325Z","shell.execute_reply":"2024-04-07T11:49:44.852473Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def run_weights_raw_eeg_butter_filter():\n    device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n\n    print('infer with weights_raw_eeg_butter_filter')\n    predictions=[]\n    for model_weight in CFG['weights_raw_eeg_butter_filter']:\n        test_dataset = AlaskaDataIter(test_df, training_flag=False, shuffle=False,use_eeg=True,ll=0.5,rr=20,use_mne_filter=False)\n        test_loader = DataLoader(test_dataset,\n                         CFG['batch_size'],\n                         num_workers=CFG['num_worker'],\n                         shuffle=False)\n\n        model = Net1d()\n        state_dict = torch.load(model_weight, map_location=device)\n\n        model.load_state_dict(state_dict,strict=True)\n        model.to(device)\n        prediction_dict = inference_function(test_loader, model, device)\n        predictions.append(prediction_dict[\"predictions\"])\n\n\n        test_dataset = AlaskaDataIter(test_df, training_flag=False, shuffle=False,use_eeg=True,flip=True,ll=0.5,rr=20,use_mne_filter=False)\n        test_loader = DataLoader(test_dataset,\n                         CFG['batch_size'],\n                         num_workers=CFG['num_worker'],\n                         shuffle=False)\n\n        model = Net1d()\n        state_dict = torch.load(model_weight, map_location=device)\n\n        model.load_state_dict(state_dict,strict=True)\n        model.to(device)\n        prediction_dict = inference_function(test_loader, model, device)\n        predictions.append(prediction_dict[\"predictions\"])\n\n        torch.cuda.empty_cache()\n        gc.collect()\n    predictions = np.array(predictions)\n    predictions = np.mean(predictions, axis=0)\n    \n    return predictions","metadata":{"execution":{"iopub.status.busy":"2024-04-07T11:49:45.252176Z","iopub.execute_input":"2024-04-07T11:49:45.252876Z","iopub.status.idle":"2024-04-07T11:49:45.263393Z","shell.execute_reply.started":"2024-04-07T11:49:45.252846Z","shell.execute_reply":"2024-04-07T11:49:45.262436Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def run_weights_hgnet():\n    device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n\n    print('infer with weights_hgnet')\n    predictions=[]\n    for model_weight in CFG['weights_hgnet']:\n        test_dataset = AlaskaDataIter(test_df, training_flag=False, shuffle=False,use_eeg=True,ll=0.5,rr=20,use_mne_filter=True)\n        test_loader = DataLoader(test_dataset,\n                         CFG['batch_size'],\n                         num_workers=CFG['num_worker'],\n                         shuffle=False)\n\n        model = Net1dhgnet()\n        state_dict = torch.load(model_weight, map_location=device)\n\n        model.load_state_dict(state_dict,strict=True)\n        model.to(device)\n        prediction_dict = inference_function(test_loader, model, device)\n        predictions.append(prediction_dict[\"predictions\"])\n\n\n        test_dataset = AlaskaDataIter(test_df, training_flag=False, shuffle=False,use_eeg=True,flip=True,ll=0.5,rr=20,use_mne_filter=True)\n        test_loader = DataLoader(test_dataset,\n                         CFG['batch_size'],\n                         num_workers=CFG['num_worker'],\n                         shuffle=False)\n\n        model = Net1dhgnet()\n        state_dict = torch.load(model_weight, map_location=device)\n\n        model.load_state_dict(state_dict,strict=True)\n        model.to(device)\n        prediction_dict = inference_function(test_loader, model, device)\n        predictions.append(prediction_dict[\"predictions\"])\n\n        torch.cuda.empty_cache()\n        gc.collect()\n    predictions = np.array(predictions)\n    predictions = np.mean(predictions, axis=0)\n    \n    return predictions","metadata":{"execution":{"iopub.status.busy":"2024-04-07T11:49:45.744259Z","iopub.execute_input":"2024-04-07T11:49:45.744970Z","iopub.status.idle":"2024-04-07T11:49:45.755777Z","shell.execute_reply.started":"2024-04-07T11:49:45.744936Z","shell.execute_reply":"2024-04-07T11:49:45.754886Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def run_weight_x3d():\n    device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n\n    print('infer with weights_x3d')\n    predictions=[]\n    for model_weight in CFG['weights_x3d']:\n        test_dataset = AlaskaDataIter(test_df, training_flag=False, shuffle=False,use_eeg=True,ll=0.5,rr=20,use_mne_filter=True)\n        test_loader = DataLoader(test_dataset,\n                         CFG['batch_size'],\n                         num_workers=CFG['num_worker'],\n                         shuffle=False)\n\n        model = Netx3d()\n        state_dict = torch.load(model_weight, map_location=device)\n\n        model.load_state_dict(state_dict,strict=True)\n        model.to(device)\n        prediction_dict = inference_function(test_loader, model, device)\n        predictions.append(prediction_dict[\"predictions\"])\n\n\n        test_dataset = AlaskaDataIter(test_df, training_flag=False, shuffle=False,use_eeg=True,flip=True,ll=0.5,rr=20,use_mne_filter=True)\n        test_loader = DataLoader(test_dataset,\n                         CFG['batch_size'],\n                         num_workers=CFG['num_worker'],\n                         shuffle=False)\n\n        model = Netx3d()\n        state_dict = torch.load(model_weight, map_location=device)\n\n        model.load_state_dict(state_dict,strict=True)\n        model.to(device)\n        prediction_dict = inference_function(test_loader, model, device)\n        predictions.append(prediction_dict[\"predictions\"])\n\n        torch.cuda.empty_cache()\n        gc.collect()\n    predictions = np.array(predictions)\n    predictions = np.mean(predictions, axis=0)\n    \n    return predictions","metadata":{"execution":{"iopub.status.busy":"2024-04-07T11:49:46.144752Z","iopub.execute_input":"2024-04-07T11:49:46.145535Z","iopub.status.idle":"2024-04-07T11:49:46.156191Z","shell.execute_reply.started":"2024-04-07T11:49:46.145491Z","shell.execute_reply":"2024-04-07T11:49:46.155171Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def run_weight_double_head():\n    device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n\n    print('infer with weights_doublehead')\n    predictions=[]\n    for model_weight in CFG['weights_doublehead']:\n        \n        test_dataset = AlaskaDataIter(test_df, training_flag=False, shuffle=False,use_eeg=True,ll=0.5,rr=40)\n        test_loader = DataLoader(test_dataset,\n                         CFG['batch_size']//2,\n                         num_workers=CFG['num_worker'],\n                         shuffle=False)\n\n        model = Netdoublehead()\n        state_dict = torch.load(model_weight, map_location=device)\n\n        model.load_state_dict(state_dict,strict=True)\n        model.to(device)\n        prediction_dict = inference_function(test_loader, model, device)\n        predictions.append(prediction_dict[\"predictions\"])\n\n\n        test_dataset = AlaskaDataIter(test_df, training_flag=False, shuffle=False,flip=True,use_eeg=True,ll=0.5,rr=40)\n        test_loader = DataLoader(test_dataset,\n                         CFG['batch_size']//2,\n                         num_workers=CFG['num_worker'],\n                         shuffle=False)\n\n        model = Netdoublehead()\n        state_dict = torch.load(model_weight, map_location=device)\n\n        model.load_state_dict(state_dict,strict=True)\n        model.to(device)\n        prediction_dict = inference_function(test_loader, model, device)\n        predictions.append(prediction_dict[\"predictions\"])\n\n        torch.cuda.empty_cache()\n        gc.collect()\n    predictions = np.array(predictions)\n    predictions = np.mean(predictions, axis=0)\n    \n    return predictions","metadata":{"execution":{"iopub.status.busy":"2024-04-07T11:49:47.041680Z","iopub.execute_input":"2024-04-07T11:49:47.042300Z","iopub.status.idle":"2024-04-07T11:49:47.052480Z","shell.execute_reply.started":"2024-04-07T11:49:47.042265Z","shell.execute_reply":"2024-04-07T11:49:47.051573Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def run_weight_double_headbutterfilter():\n    device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n\n    print('infer with weights_doubleheadbutterfilter')\n    predictions=[]\n    for model_weight in CFG['weights_doubleheadbutterfilter']:\n        \n        test_dataset = AlaskaDataIter(test_df, training_flag=False, shuffle=False,use_eeg=True,ll=0.5,rr=20,use_mne_filter=False)\n        test_loader = DataLoader(test_dataset,\n                         CFG['batch_size']//2,\n                         num_workers=CFG['num_worker'],\n                         shuffle=False)\n\n        model = Netdoubleheadbutterfilter()\n        state_dict = torch.load(model_weight, map_location=device)\n\n        model.load_state_dict(state_dict,strict=True)\n        model.to(device)\n        prediction_dict = inference_function(test_loader, model, device)\n        predictions.append(prediction_dict[\"predictions\"])\n\n\n        test_dataset = AlaskaDataIter(test_df, training_flag=False, shuffle=False,flip=True,use_eeg=True,ll=0.5,rr=20,use_mne_filter=False)\n        test_loader = DataLoader(test_dataset,\n                         CFG['batch_size']//2,\n                         num_workers=CFG['num_worker'],\n                         shuffle=False)\n\n        model = Netdoubleheadbutterfilter()\n        state_dict = torch.load(model_weight, map_location=device)\n\n        model.load_state_dict(state_dict,strict=True)\n        model.to(device)\n        prediction_dict = inference_function(test_loader, model, device)\n        predictions.append(prediction_dict[\"predictions\"])\n\n        torch.cuda.empty_cache()\n        gc.collect()\n    predictions = np.array(predictions)\n    predictions = np.mean(predictions, axis=0)\n    \n    return predictions","metadata":{"execution":{"iopub.status.busy":"2024-04-07T11:49:47.541544Z","iopub.execute_input":"2024-04-07T11:49:47.542113Z","iopub.status.idle":"2024-04-07T11:49:47.552431Z","shell.execute_reply.started":"2024-04-07T11:49:47.542084Z","shell.execute_reply":"2024-04-07T11:49:47.551588Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predictions = []\n\n# ***0.32 best spec  version1\n# ans=run_weight_spec()\n# predictions.append(ans)\n\n\n\n\n# #***0.25 best version 33\nans=run_weight_x3d()\npredictions.append(ans) \n\n# **0.26  version 8\nans=run_weight_one_image()\npredictions.append(ans) \n\n# #***0.24 bestversion 9 0.5-40hz\n# ans=run_weight_double_head()\n# predictions.append(ans) \n\n# #***0.23 best version29, add 1000 val data\nans=run_weight_eeg_raw()\npredictions.append(ans) \n\n#***0.23 best version5 butter filter\n# ans=run_weights_raw_eeg_butter_filter()\n# predictions.append(ans) \n\n#***0.23 best version1 butter filter double head\nans=run_weight_double_headbutterfilter()\npredictions.append(ans) \n\n#***0.23 best version3 mne filter hgnetb5\n# ans=run_weights_hgnet()\n# predictions.append(ans) \n\nweights_by_score=np.array([0.15,0.15,0.35,0.35])\n\npredictions=predictions[0]*weights_by_score[0]+\\\npredictions[1]*weights_by_score[1]+\\\npredictions[2]*weights_by_score[2]+\\\npredictions[3]*weights_by_score[3]\n# predictions=np.array(predictions)\n# predictions=np.mean(predictions,axis=0)\n\npredictions.shape\npredictions","metadata":{"execution":{"iopub.status.busy":"2024-04-07T12:06:40.046398Z","iopub.execute_input":"2024-04-07T12:06:40.046817Z","iopub.status.idle":"2024-04-07T12:08:48.376884Z","shell.execute_reply.started":"2024-04-07T12:06:40.046790Z","shell.execute_reply":"2024-04-07T12:08:48.375927Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TARGETS = ['seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']\nsub = pd.DataFrame({'eeg_id': test_df.eeg_id.values})\nsub[TARGETS] = predictions\nsub.to_csv('submission.csv',index=False)\nprint(f'Submissionn shape: {sub.shape}')\nsub.head()","metadata":{"execution":{"iopub.status.busy":"2024-04-07T12:08:50.225731Z","iopub.execute_input":"2024-04-07T12:08:50.226436Z","iopub.status.idle":"2024-04-07T12:08:50.246693Z","shell.execute_reply.started":"2024-04-07T12:08:50.226404Z","shell.execute_reply":"2024-04-07T12:08:50.245836Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}