{"metadata":{"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"},{"sourceId":2132855,"sourceType":"datasetVersion","datasetId":900016},{"sourceId":7392775,"sourceType":"datasetVersion","datasetId":4297782},{"sourceId":161531087,"sourceType":"kernelVersion"}],"dockerImageVersionId":30636,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true},"kernelspec":{"display_name":"Python 3","language":"python","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"},"papermill":{"default_parameters":{},"duration":102.485133,"end_time":"2024-01-16T12:14:21.10455","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2024-01-16T12:12:38.619417","version":"2.4.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import sys\nsys.path.append('/kaggle/input/efficientnet-pytorch/EfficientNet-PyTorch/EfficientNet-PyTorch-master')\nfrom efficientnet_pytorch import EfficientNet","metadata":{"execution":{"iopub.status.busy":"2024-04-28T18:06:31.525593Z","iopub.execute_input":"2024-04-28T18:06:31.525946Z","iopub.status.idle":"2024-04-28T18:06:31.530815Z","shell.execute_reply.started":"2024-04-28T18:06:31.525918Z","shell.execute_reply":"2024-04-28T18:06:31.529805Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd, numpy as np, os\nimport matplotlib.pyplot as plt, gc\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom scipy.signal import butter, lfilter\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom scipy.signal import freqz\n\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.nn.functional as F\nfrom tqdm import tqdm\nfrom sklearn.model_selection import train_test_split\nfrom sklearn import preprocessing\nfrom sklearn.model_selection import KFold, GroupKFold\n\nfrom torchvision import datasets, transforms","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","papermill":{"duration":0.693827,"end_time":"2024-01-16T12:12:42.606147","exception":false,"start_time":"2024-01-16T12:12:41.91232","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-04-28T18:06:31.532443Z","iopub.execute_input":"2024-04-28T18:06:31.532739Z","iopub.status.idle":"2024-04-28T18:06:31.541133Z","shell.execute_reply.started":"2024-04-28T18:06:31.532712Z","shell.execute_reply":"2024-04-28T18:06:31.540157Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv('/kaggle/input/hms-harmful-brain-activity-classification/train.csv')\nTARGETS = df.columns[-6:]\nprint('Train shape:', df.shape )\nprint('Targets', list(TARGETS))\ndf.head()","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","papermill":{"duration":0.693827,"end_time":"2024-01-16T12:12:42.606147","exception":false,"start_time":"2024-01-16T12:12:41.91232","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-04-28T18:06:31.543201Z","iopub.execute_input":"2024-04-28T18:06:31.543581Z","iopub.status.idle":"2024-04-28T18:06:31.726826Z","shell.execute_reply.started":"2024-04-28T18:06:31.543556Z","shell.execute_reply":"2024-04-28T18:06:31.725908Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = df.groupby('eeg_id')[['spectrogram_id','spectrogram_label_offset_seconds']].agg(\n    {'spectrogram_id':'first','spectrogram_label_offset_seconds':'min'})\ntrain.columns = ['spec_id','min']\n\ntmp = df.groupby('eeg_id')[['spectrogram_id','spectrogram_label_offset_seconds']].agg(\n    {'spectrogram_label_offset_seconds':'max'})\ntrain['max'] = tmp\n\ntmp = df.groupby('eeg_id')[['patient_id']].agg('first')\ntrain['patient_id'] = tmp\n\ntmp = df.groupby('eeg_id')[TARGETS].agg('sum')\nfor t in TARGETS:\n    train[t] = tmp[t].values\n    \ny_data = train[TARGETS].values\ny_data = y_data / y_data.sum(axis=1,keepdims=True)\ntrain[TARGETS] = y_data\n\ntmp = df.groupby('eeg_id')[['expert_consensus']].apply(lambda x: x.mode().iloc[0]).reset_index()\ntmp2 = df.groupby(['eeg_id','expert_consensus'])[['eeg_sub_id']].agg(min).reset_index()\ntmp = pd.merge(tmp,tmp2,on=['eeg_id','expert_consensus'],how='left')\ntrain['target'] = tmp['expert_consensus'].values\ntrain['eeg_sub_id'] = tmp['eeg_sub_id'].values\n\n\ntrain = train.reset_index()\nprint('Train non-overlapp eeg_id shape:', train.shape )\ntrain.head()","metadata":{"execution":{"iopub.status.busy":"2024-04-28T18:06:31.728535Z","iopub.execute_input":"2024-04-28T18:06:31.728869Z","iopub.status.idle":"2024-04-28T18:06:39.854213Z","shell.execute_reply.started":"2024-04-28T18:06:31.728832Z","shell.execute_reply":"2024-04-28T18:06:39.853234Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"IS_TRAINING=True\n\nif IS_TRAINING:\n    spectrograms = np.load('/kaggle/input/brain-spectrograms/specs.npy',allow_pickle=True).item()\n    eegs_data = np.load('/kaggle/input/hms-eeg-raw-dataset-16waves/16_waves_eeg_specs_partial_train.npy',allow_pickle=True).item()","metadata":{"execution":{"iopub.status.busy":"2024-04-28T18:06:39.855315Z","iopub.execute_input":"2024-04-28T18:06:39.855568Z","iopub.status.idle":"2024-04-28T18:08:03.514444Z","shell.execute_reply.started":"2024-04-28T18:06:39.855546Z","shell.execute_reply":"2024-04-28T18:08:03.513419Z"},"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']]\n\ndef butter_bandpass(lowcut, highcut, fs, order=5):\n    return butter(order, [lowcut, highcut], fs=fs, btype='band')\n\ndef butter_bandpass_filter(data, lowcut, highcut, fs, order=5):\n    b, a = butter_bandpass(lowcut, highcut, fs, order=order)\n    y = lfilter(b, a, data)\n    return y\n\n\ndef denoise_filter(x):\n    # Sample rate and desired cutoff frequencies (in Hz).\n    fs = 200.0\n    lowcut = 1.0\n    highcut = 25.0\n    \n    # Filter a noisy signal.\n    T = 50\n    nsamples = T * fs\n    t = np.arange(0, nsamples) / fs\n    y = butter_bandpass_filter(x, lowcut, highcut, fs, order=6)\n    y = (y + np.roll(y,-1)+ np.roll(y,-2)+ np.roll(y,-3))/4\n    y = y[0:-1:4]\n    \n    return y","metadata":{"execution":{"iopub.status.busy":"2024-04-28T18:08:03.516821Z","iopub.execute_input":"2024-04-28T18:08:03.517157Z","iopub.status.idle":"2024-04-28T18:08:03.526258Z","shell.execute_reply.started":"2024-04-28T18:08:03.517124Z","shell.execute_reply":"2024-04-28T18:08:03.525303Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataloader","metadata":{"papermill":{"duration":0.028035,"end_time":"2024-01-16T12:13:50.519598","exception":false,"start_time":"2024-01-16T12:13:50.491563","status":"completed"},"tags":[]}},{"cell_type":"code","source":"test_eeg_path = '/kaggle/input/hms-harmful-brain-activity-classification/test_eegs/'\nclass CustomDataset(Dataset):\n    def __init__(self, dataframe,specs,eegs_data,mode='Train', transform=None):\n        self.dataframe = dataframe\n        self.specs = specs\n        self.mode = mode\n        self.eegs_data = eegs_data\n        \n    def __len__(self):\n        return len(self.dataframe)\n    \n    def __getitem__(self, idx):\n        \n        if self.mode=='Test':\n            row = self.dataframe.iloc[idx]\n            eeg_id = row['eeg_id']\n            parq_path = f'{test_eeg_path}{eeg_id}.parquet'\n            eeg = pd.read_parquet(parq_path)\n            rows = len(eeg)\n            offset = (rows-10_000)//2\n            eeg = eeg.iloc[offset:offset+10_000]\n            signals = []\n            for k in range(4):\n                    COLS = FEATS[k]\n\n                    # COMPUTE PAIR DIFFERENCES AND AVERAGE\n        #             x = eeg[COLS[0]].values - eeg[COLS[1]].values\n                    for j in range(4):\n                        x = eeg[COLS[j]].values - eeg[COLS[j+1]].values\n                        x = denoise_filter(x)\n                        signals.append(x)\n            signal_eeg = np.array(signals,dtype='float32')\n        else:\n            row = self.dataframe.iloc[idx]\n            eeg_id = row['eeg_id']\n            eeg_sub_id = row['eeg_sub_id']\n            eeg_key = f'{eeg_id}_{eeg_sub_id}'\n            signal_eeg = self.eegs_data[eeg_key]\n        ##############################\n        \n        row = self.dataframe.iloc[idx]\n        if self.mode=='Test': \n                r = 0\n        else: \n            r = int( (row['min'] + row['max'])//4 )\n        X = np.zeros((128,256,4))#,dtype='float64')\n        for k in range(4):\n                # EXTRACT 300 ROWS OF SPECTROGRAM\n                img = self.specs[row.spec_id][r:r+300,k*100:(k+1)*100].T\n                \n                # LOG TRANSFORM SPECTROGRAM\n                img = np.clip(img,np.exp(-4),np.exp(8))\n                img = np.log(img)\n                \n                # 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                # CROP TO 256 TIME STEPS\n                X[14:-14,:,k] = img[:,22:-22] / 2.0\n                \n        X = torch.tensor(X,dtype=torch.float32)\n#         X = X.permute(2, 0, 1)\n#         print('test',X.shape)\n        if self.mode=='Test': \n            return X,torch.tensor(signal_eeg,dtype=torch.float32)\n        else:\n            labels = row[TARGETS].values.astype(np.float32) #np.array(row[-6:]).reshape(6,1)\n            labels = labels/np.sum(labels)\n            return X,torch.tensor(signal_eeg,dtype=torch.float32), torch.tensor(labels,dtype=torch.float32)","metadata":{"execution":{"iopub.status.busy":"2024-04-28T18:08:03.527675Z","iopub.execute_input":"2024-04-28T18:08:03.528242Z","iopub.status.idle":"2024-04-28T18:08:03.545524Z","shell.execute_reply.started":"2024-04-28T18:08:03.528209Z","shell.execute_reply":"2024-04-28T18:08:03.544646Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_name = 'efficientnet-b0'\nfile_name  = '/kaggle/input/efficientnet-pytorch/efficientnet-b0-08094119.pth'","metadata":{"execution":{"iopub.status.busy":"2024-04-28T18:08:03.546759Z","iopub.execute_input":"2024-04-28T18:08:03.547045Z","iopub.status.idle":"2024-04-28T18:08:03.556526Z","shell.execute_reply.started":"2024-04-28T18:08:03.547022Z","shell.execute_reply":"2024-04-28T18:08:03.555417Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class wave_residual_block(nn.Module):\n    def __init__(self,in_channels,out_channels,layer_num):\n        super(wave_residual_block, self).__init__()\n        # kernel size will always be 2,strid=1 (comparing two samples of time)\n        dilatn = 2**(layer_num-1)\n        self.dilatn = dilatn\n        self.filter_conv = nn.Conv1d(in_channels, out_channels, 2, stride=1,padding=dilatn, dilation=dilatn, bias=False)\n        self.gate_conv = nn.Conv1d(in_channels, out_channels, 2, stride=1,padding=dilatn, dilation=dilatn, bias=False)\n        self.conv = nn.Conv1d(out_channels, out_channels, 1,1)\n\n    def forward(self, x):\n        \n        y = F.tanh(self.filter_conv(x))*F.sigmoid(self.gate_conv(x))\n        y = y[:, :, :-self.dilatn]\n        y = self.conv(y)\n        x = x + y\n        return x,y\n    \n# Define the 1D CNN model\nclass WaveClassifier(nn.Module):\n    def __init__(self,in_channels=16,num_layers = 6):  # 6 layers will cover 64 samples which correspond to 64/50 seconds\n        super(WaveClassifier, self).__init__()\n        self.waveblock_0 = wave_residual_block(16,16,1)\n        self.waveblocks = nn.ModuleList([wave_residual_block(16,16,i) for i in range(2,num_layers+1)])\n#         for i in range(2,num_layers+1):\n        self.conv = nn.Conv1d(in_channels, 16, 1,1)\n        self.num_layers = num_layers\n        \n        self.conv1 = nn.Conv1d(16, 48, 20, 12)\n        self.conv2 = nn.Conv1d(48, 32, 10, 6)\n        self.flatten = nn.Flatten()\n#         self.fc1 = nn.Linear(1536, 6)  # Adjust the input size based on your input dimensions\n#         self.softmax = nn.Softmax(dim=1)\n        self.dropout = nn.Dropout(p=0.3)\n    def forward(self, x):\n        \n        x = self.conv(x)\n        skip_connections = []\n        x,y = self.waveblock_0(x)\n        skip_connections.append(y)\n        for i in range(self.num_layers-1):\n            x,y = self.waveblocks[i](x)\n            skip_connections.append(y)\n        \n        y_list = torch.stack(skip_connections) \n        x = torch.sum(y_list,dim=0,keepdim=True)\n        x = torch.squeeze(x,dim=0)\n        x = F.relu(self.conv1(x))\n        x = self.dropout(x)\n        x = F.relu(self.conv2(x))\n        x = self.flatten(x)\n#         x = self.fc1(x)\n#         x = self.softmax(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2024-04-28T18:08:03.557586Z","iopub.execute_input":"2024-04-28T18:08:03.557875Z","iopub.status.idle":"2024-04-28T18:08:03.573731Z","shell.execute_reply.started":"2024-04-28T18:08:03.557852Z","shell.execute_reply":"2024-04-28T18:08:03.572814Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nclass EffNET_HMS(nn.Module):\n    def __init__(self,eff_model):  # 6 layers will cover 64 samples which correspond to 64/50 seconds\n        super(EffNET_HMS, self).__init__()\n        self.pretrained_model = eff_model\n#         self.flatten = nn.Flatten()\n#         self.fc1 = nn.Linear(1000, 6)  # Adjust the input size based on your input dimensions\n#         self.softmax = nn.Softmax(dim=1)\n\n    def forward(self, x):\n#         print(x.shape)\n        x1 = [x[:,:,:,i:i+1] for i in range(4)]\n        x = torch.cat(x1,dim=1)\n        x = torch.cat([x,x,x],dim=3)\n#         print(x.shape)\n        x = x.permute(0,3, 1, 2)\n        x = self.pretrained_model(x)\n#         x = self.fc1(x)\n#         x = self.softmax(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2024-04-28T18:08:03.574968Z","iopub.execute_input":"2024-04-28T18:08:03.575388Z","iopub.status.idle":"2024-04-28T18:08:03.583723Z","shell.execute_reply.started":"2024-04-28T18:08:03.575358Z","shell.execute_reply":"2024-04-28T18:08:03.582753Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nclass Combined_HMS(nn.Module):\n    def __init__(self,eff_model):  # 6 layers will cover 64 samples which correspond to 64/50 seconds\n        super(Combined_HMS, self).__init__()\n        self.waveclassifier = WaveClassifier()\n        self.efficient_net = EffNET_HMS(eff_model)\n#         self.flatten = nn.Flatten()\n        self.fc1 = nn.Linear((1056+500)//3, 6)  # Adjust the input size based on your input dimensions\n        self.softmax = nn.Softmax(dim=1)\n        self.dropout1 = nn.Dropout(p=0.6)\n        self.dropout2 = nn.Dropout(p=0.4)\n        self.maxpool1 = nn.AvgPool1d(2,2)#AvgPool1d\n        self.maxpool2 = nn.AvgPool1d(3,3)#AvgPool1d\n    def forward(self, x,y):\n        \n        x = self.efficient_net(x)\n        x = self.maxpool1(x)\n        x = self.dropout1(x)\n#         print(x.shape)\n#         print(y.shape)\n        y = self.waveclassifier(y)\n#         print(y.shape)\n        x = torch.cat([x,y],dim=1)\n        x = self.maxpool2(x)\n        x = self.dropout2(x)\n#         print(x.shape)\n        x = self.fc1(x)\n        x = self.softmax(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2024-04-28T18:08:03.584746Z","iopub.execute_input":"2024-04-28T18:08:03.585034Z","iopub.status.idle":"2024-04-28T18:08:03.594380Z","shell.execute_reply.started":"2024-04-28T18:08:03.585002Z","shell.execute_reply":"2024-04-28T18:08:03.593410Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n","metadata":{"execution":{"iopub.status.busy":"2024-04-28T18:08:03.597384Z","iopub.execute_input":"2024-04-28T18:08:03.597738Z","iopub.status.idle":"2024-04-28T18:08:03.604074Z","shell.execute_reply.started":"2024-04-28T18:08:03.597716Z","shell.execute_reply":"2024-04-28T18:08:03.603243Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not os.path.exists('combined_model'):\n        os.makedirs('combined_model')","metadata":{"execution":{"iopub.status.busy":"2024-04-28T18:08:03.605213Z","iopub.execute_input":"2024-04-28T18:08:03.605965Z","iopub.status.idle":"2024-04-28T18:08:03.612379Z","shell.execute_reply.started":"2024-04-28T18:08:03.605934Z","shell.execute_reply":"2024-04-28T18:08:03.611475Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def validatn_loss(data_loader, model):\n        \n    model.to(device)\n    model.eval()    \n    predictions = []\n    true = []\n#     loss_lst = []\n    for batch in data_loader:\n        with torch.no_grad():\n            inp1,inp2, label = batch\n            outputs = our_model(inp1.to(device),inp2.to(device))\n            # inputs = {key:val.reshape(val.shape[0], -1).to(config.device) for key,val in batch.items()}\n        predictions.extend(outputs.detach())\n        true.extend(label)\n#         loss_lst.extend(criterion(torch.log(outputs.detach()), y.to(device)))\n        \n\n    predictions = torch.vstack(predictions)\n    true = torch.vstack(true)\n#     loss_lst = torch.vstack(loss_lst)\n#     print(loss_lst.mean())\n    loss = criterion(torch.log(predictions), true.to(device))\n#     print(loss)\n    return loss.item()","metadata":{"execution":{"iopub.status.busy":"2024-04-28T18:08:03.613482Z","iopub.execute_input":"2024-04-28T18:08:03.615394Z","iopub.status.idle":"2024-04-28T18:08:03.622534Z","shell.execute_reply.started":"2024-04-28T18:08:03.615369Z","shell.execute_reply":"2024-04-28T18:08:03.621744Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_splits = 4\nif IS_TRAINING:\n    all_oof = []\n    all_true = []\n    gkf = GroupKFold(n_splits=num_splits)\n    for i, (train_index, valid_index) in enumerate(gkf.split(train, train.target, train.patient_id)):\n        dataset_train = CustomDataset(dataframe=train.iloc[train_index],specs = spectrograms,eegs_data=eegs_data)\n        dataset_test = CustomDataset(dataframe=train.iloc[valid_index],specs = spectrograms,eegs_data=eegs_data)\n        train_dataloader = DataLoader(dataset_train, batch_size=16, shuffle=True,drop_last=True)\n        val_loader = DataLoader(dataset_test, batch_size=16,shuffle=False)\n\n        min_val_loss = 99\n        epochs = 4\n        eff_model = EfficientNet.from_name(model_name)\n        eff_model.load_state_dict(torch.load(file_name))\n        our_model = Combined_HMS(eff_model)\n        our_model.to(device)\n        \n        optimizer = optim.AdamW(our_model.parameters(), lr=0.001,weight_decay=0.04)\n        criterion = nn.KLDivLoss(reduction=\"batchmean\").cuda()\n        \n        for epoch in range(epochs):\n            pbar = tqdm(train_dataloader)\n            running_loss=0\n            cnt=0\n            our_model.train()\n            for batch in pbar:\n                cnt+=1\n                inp1,inp2, label = batch\n                pred = our_model(inp1.to(device),inp2.to(device))\n                loss = criterion(torch.log(pred), label.to(device))\n                # Backward pass and optimization\n                optimizer.zero_grad()\n                loss.backward()\n                optimizer.step()\n                running_loss += loss*inp1.size(0)/len(train_dataloader.dataset)\n\n                pbar.set_description(f\"Batch loss : | Runn: {loss.item():.2f} | {running_loss.item():.2f}\")\n                \n                if cnt % len(train_dataloader)//2 == 0 :\n                    val_loss = validatn_loss(val_loader,our_model)\n                    if min_val_loss > val_loss :\n                        min_val_loss = val_loss\n                        torch.save(our_model.state_dict(), f'combined_model/model_best_fold_{i}.pt')\n                        print('i,epoch_half,val loss : ',i,epoch,val_loss)\n                    our_model.train()\n                    \n            val_loss = validatn_loss(val_loader,our_model)\n            if min_val_loss > val_loss :\n                min_val_loss = val_loss\n                torch.save(our_model.state_dict(), f'combined_model/model_best_fold_{i}.pt')\n                print('i,epoch,val loss : ',i,epoch,val_loss)","metadata":{"execution":{"iopub.status.busy":"2024-04-28T18:08:03.623681Z","iopub.execute_input":"2024-04-28T18:08:03.623958Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Submit to Kaggle LB\n","metadata":{"papermill":{"duration":0.039743,"end_time":"2024-01-16T12:14:15.75862","exception":false,"start_time":"2024-01-16T12:14:15.718877","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# del train; gc.collect()\ntest = pd.read_csv('/kaggle/input/hms-harmful-brain-activity-classification/test.csv')\nprint('Test shape:',test.shape)\ntest.head()","metadata":{"papermill":{"duration":0.059958,"end_time":"2024-01-16T12:14:15.85819","exception":false,"start_time":"2024-01-16T12:14:15.798232","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# READ ALL SPECTROGRAMS\nPATH2 = '/kaggle/input/hms-harmful-brain-activity-classification/test_spectrograms/'\nfiles2 = os.listdir(PATH2)\nprint(f'There are {len(files2)} test spectrogram parquets')\n    \nspectrograms2 = {}\nfor i,f in enumerate(files2):\n    if i%100==0: print(i,', ',end='')\n    tmp = pd.read_parquet(f'{PATH2}{f}')\n    name = int(f.split('.')[0])\n    spectrograms2[name] = tmp.iloc[:,1:].values\n    \n# RENAME FOR DATALOADER\ntest = test.rename({'spectrogram_id':'spec_id'},axis=1)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset_test = CustomDataset(dataframe=test,specs = spectrograms2,mode='Test',eegs_data=None)\n# CustomDataset(dataframe=train.iloc[train_index],specs = spectrograms,eegs_data=eegs_data)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_loader = DataLoader(dataset_test, batch_size=8,shuffle=False)\n# model = EffNET_HMS().double()\n\neff_model = EfficientNet.from_name(model_name)\n# eff_model.load_state_dict(torch.load(file_name))\n#         our_model.pretrained_model.load_state_dict(torch.load(file_name))\nour_model = Combined_HMS(eff_model)\n\npreds_all_fold = []\nfor i in range(num_splits):\n#     our_model.load_state_dict(torch.load('/kaggle/input/16waves-combined-models-fold5/model_best_fold_'+str(i)+'.pt'))\n    our_model.load_state_dict(torch.load(f'/kaggle/input/hms-combined-4split/model_best_fold_{i}.pt'))\n    our_model.to(device)\n    our_model.eval()\n    \n    preds = []\n    for batch in test_loader:\n        inp1,inp2 = batch\n        pred = our_model(inp1.to(device),inp2.to(device),)\n        preds.append(pred.detach().cpu().numpy())\n    preds = np.vstack(preds)\n    preds_all_fold.append(preds)\n\nprediction_all_fold = np.mean(preds_all_fold,axis=0)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# CREATE SUBMISSION.CSV\nfrom IPython.display import display\n\nsub = pd.DataFrame({'eeg_id':test.eeg_id.values})\nsub[TARGETS] = prediction_all_fold\nsub.to_csv('submission.csv',index=False)\nprint('Submission shape',sub.shape)\ndisplay( sub.head() )\n\n# SANITY CHECK TO CONFIRM PREDICTIONS SUM TO ONE\nprint('Sub row 0 sums to:',sub.iloc[0,-6:].sum())","metadata":{"papermill":{"duration":0.07372,"end_time":"2024-01-16T12:14:18.05987","exception":false,"start_time":"2024-01-16T12:14:17.98615","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}