{"metadata":{"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"},{"sourceId":7464325,"sourceType":"datasetVersion","datasetId":4338463},{"sourceId":7544414,"sourceType":"datasetVersion","datasetId":4393462},{"sourceId":7546083,"sourceType":"datasetVersion","datasetId":4391884},{"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":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.10.12"},"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 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","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-02-04T13:33:15.219222Z","iopub.execute_input":"2024-02-04T13:33:15.219685Z","iopub.status.idle":"2024-02-04T13:33:19.336138Z","shell.execute_reply.started":"2024-02-04T13:33:15.219644Z","shell.execute_reply":"2024-02-04T13:33:19.335276Z"},"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-02-04T13:33:19.337763Z","iopub.execute_input":"2024-02-04T13:33:19.338184Z","iopub.status.idle":"2024-02-04T13:33:19.649267Z","shell.execute_reply.started":"2024-02-04T13:33:19.338154Z","shell.execute_reply":"2024-02-04T13:33:19.648300Z"},"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-02-04T13:33:19.650443Z","iopub.execute_input":"2024-02-04T13:33:19.650756Z","iopub.status.idle":"2024-02-04T13:33:27.686075Z","shell.execute_reply.started":"2024-02-04T13:33:19.650727Z","shell.execute_reply":"2024-02-04T13:33:27.685156Z"},"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-02-04T13:33:27.689121Z","iopub.execute_input":"2024-02-04T13:33:27.690072Z","iopub.status.idle":"2024-02-04T13:33:27.699873Z","shell.execute_reply.started":"2024-02-04T13:33:27.690031Z","shell.execute_reply":"2024-02-04T13:33:27.699062Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"IS_TRAINING=False\n\nif IS_TRAINING:\n#     eegs_data = np.load('/kaggle/input/hms-eeg-raw-dataset/16_waves_eeg_specs_partial_train_quantize.npy',allow_pickle=True).item()\n    eegs_data = np.load('/kaggle/input/hms-eeg-raw-dataset/16_waves_eeg_specs_partial_train.npy',allow_pickle=True).item()","metadata":{"execution":{"iopub.status.busy":"2024-02-04T13:33:27.700922Z","iopub.execute_input":"2024-02-04T13:33:27.701195Z","iopub.status.idle":"2024-02-04T13:34:07.299441Z","shell.execute_reply.started":"2024-02-04T13:33:27.701169Z","shell.execute_reply":"2024-02-04T13:34:07.298367Z"},"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,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#             signal_eeg = quantize_data(signal_eeg,1)\n        ##############################\n        \n        row = self.dataframe.iloc[idx]\n#         X = X.permute(2, 0, 1)\n#         print('test',X.shape)\n        if self.mode=='Test': \n            return 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 torch.tensor(signal_eeg,dtype=torch.float32), torch.tensor(labels,dtype=torch.float32)","metadata":{"execution":{"iopub.status.busy":"2024-02-04T13:34:07.300846Z","iopub.execute_input":"2024-02-04T13:34:07.301593Z","iopub.status.idle":"2024-02-04T13:34:07.313183Z","shell.execute_reply.started":"2024-02-04T13:34:07.301555Z","shell.execute_reply":"2024-02-04T13:34:07.312185Z"},"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_skip = nn.Conv1d(out_channels, out_channels, 1,1)\n        self.conv_res = nn.Conv1d(out_channels, out_channels, 1,1)\n\n    def forward(self, x):\n        y = F.tanh(self.filter_conv(x))*F.sigmoid(self.gate_conv(x))\n        y = y[:, :, :-self.dilatn]\n        y_skip = self.conv_skip(y)\n        y_res = self.conv_res(y)\n        x = x + y_res\n        return x,y_skip\n    \nclass WaveBlock(nn.Module):\n    def __init__(self,in_channels=4,num_layers = 6):  # 6 layers will cover 64 samples which correspond to 64/50 seconds\n        super(WaveBlock, self).__init__()\n        \n        self.waveblocks = nn.ModuleList([wave_residual_block(16,16,i) for i in range(1,num_layers+1)])\n        self.conv0 = nn.Conv1d(in_channels,16, 1,1)\n\n        self.num_layers = num_layers\n        \n        self.conv1 = nn.Conv1d(16, 4, 1)\n\n    def forward(self, x):\n        x = self.conv0(x)\n        \n        skip_connections = []\n        for i in range(self.num_layers):\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 = self.conv1(x)\n        x = F.relu(x)\n        \n        \n        return x\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        \n        self.waveblock = WaveBlock()\n        self.conv1 = nn.Conv1d(16, 48, 20, 10)\n        self.conv2 = nn.Conv1d(48, 32, 10, 5)\n#         self.conv3 = nn.Conv1d(32, 32, 98, 1)\n        self.flatten = nn.Flatten()\n#         self.pool = nn.MaxPool1d(kernel_size=2, stride=2)\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.2)\n#         self.bn = nn.BatchNorm1d(16)\n    def forward(self,inp):\n#         inp = self.bn(inp)\n        x = []\n        for i in range(4):\n            x.append(self.waveblock(inp[:,i:i+4]))\n        x = torch.concat(x,dim=1)\n        \n        x = F.relu(self.conv1(x))\n        x = self.dropout(x)\n#         print(x.shape)\n        x = F.relu(self.conv2(x))\n        x = self.dropout(x)\n#         print(x.shape)\n        x = self.flatten(x)\n        x = self.fc1(x)\n#         print(x)\n        x = self.softmax(x)\n#         print(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2024-02-04T13:47:32.113261Z","iopub.execute_input":"2024-02-04T13:47:32.114115Z","iopub.status.idle":"2024-02-04T13:47:32.131998Z","shell.execute_reply.started":"2024-02-04T13:47:32.114081Z","shell.execute_reply":"2024-02-04T13:47:32.131014Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# model = WaveClassifier()\n# model(torch.rand(32,16,2500))","metadata":{"execution":{"iopub.status.busy":"2024-02-04T13:34:07.333665Z","iopub.execute_input":"2024-02-04T13:34:07.333945Z","iopub.status.idle":"2024-02-04T13:34:07.343989Z","shell.execute_reply.started":"2024-02-04T13:34:07.333910Z","shell.execute_reply":"2024-02-04T13:34:07.343124Z"},"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            x, y = batch\n            x = x.to(device)\n            # inputs = {key:val.reshape(val.shape[0], -1).to(config.device) for key,val in batch.items()}\n            outputs = model(x)\n        predictions.extend(outputs.detach())\n        true.extend(y)\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-02-04T13:35:10.908599Z","iopub.execute_input":"2024-02-04T13:35:10.909441Z","iopub.status.idle":"2024-02-04T13:35:10.916182Z","shell.execute_reply.started":"2024-02-04T13:35:10.909396Z","shell.execute_reply":"2024-02-04T13:35:10.915263Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not os.path.exists('wavenet_model'):\n        os.makedirs('wavenet_model')","metadata":{"execution":{"iopub.status.busy":"2024-02-04T13:34:07.356863Z","iopub.execute_input":"2024-02-04T13:34:07.357145Z","iopub.status.idle":"2024-02-04T13:34:07.363361Z","shell.execute_reply.started":"2024-02-04T13:34:07.357121Z","shell.execute_reply":"2024-02-04T13:34:07.362541Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{"execution":{"iopub.status.busy":"2024-02-04T13:34:07.364536Z","iopub.execute_input":"2024-02-04T13:34:07.364850Z","iopub.status.idle":"2024-02-04T13:34:07.423257Z","shell.execute_reply.started":"2024-02-04T13:34:07.364824Z","shell.execute_reply":"2024-02-04T13:34:07.422285Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if IS_TRAINING:\n    all_oof = []\n    all_true = []\n    gkf = GroupKFold(n_splits=5)\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],eegs_data=eegs_data)\n        dataset_test = CustomDataset(dataframe=train.iloc[valid_index],eegs_data=eegs_data)\n        train_dataloader = DataLoader(dataset_train, batch_size=32, shuffle=True,drop_last=True)\n        val_loader = DataLoader(dataset_test, batch_size=16,shuffle=False)\n\n        min_val_loss = 99\n        epochs = 6\n        \n        our_model = WaveClassifier()\n        our_model.to(device)\n        optimizer = optim.AdamW(our_model.parameters(), lr=0.001,weight_decay=0.01)\n        criterion = nn.KLDivLoss(reduction=\"batchmean\").cuda()\n#         our_model.to(torch.float32)\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, label = batch\n                pred = our_model(inp1.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:\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'wavenet_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'wavenet_model/model_best_fold_{i}.pt')\n                print('i,epoch,val loss : ',i,epoch,val_loss)","metadata":{"execution":{"iopub.status.busy":"2024-02-04T13:47:40.230992Z","iopub.execute_input":"2024-02-04T13:47:40.231968Z","iopub.status.idle":"2024-02-04T14:09:25.108700Z","shell.execute_reply.started":"2024-02-04T13:47:40.231926Z","shell.execute_reply":"2024-02-04T14:09:25.107462Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"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":"dataset_test = CustomDataset(dataframe=test,mode='Test',eegs_data=None)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_loader = DataLoader(dataset_test, batch_size=16,shuffle=False)\n\nour_model = WaveClassifier().float()\n\npreds_all_fold = []\nfor i in range(5):\n    our_model.load_state_dict(torch.load('/kaggle/input/16-wavenet-model/model_best_fold_'+str(i)+'.pt'))\n    our_model.to(device)\n    our_model.eval()\n    \n    preds = []\n    for batch in test_loader:\n        inp1 = batch\n        pred = our_model(inp1.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)\n","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":[]}]}