{"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":7899529,"sourceType":"datasetVersion","datasetId":4639130},{"sourceId":7908175,"sourceType":"datasetVersion","datasetId":4645532},{"sourceId":7930406,"sourceType":"datasetVersion","datasetId":4646663}],"dockerImageVersionId":30665,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"%pip install /kaggle/input/tmp63271321/einops-0.7.0-py3-none-any.whl","metadata":{"execution":{"iopub.status.busy":"2024-03-24T14:10:31.331473Z","iopub.execute_input":"2024-03-24T14:10:31.331734Z","iopub.status.idle":"2024-03-24T14:11:04.613183Z","shell.execute_reply.started":"2024-03-24T14:10:31.331711Z","shell.execute_reply":"2024-03-24T14:11:04.612055Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport einops\nimport cv2\nimport pandas as pd\nimport numpy as np\nfrom glob import glob\nimport matplotlib.pyplot as plt \nimport torch as t\nfrom torch.utils.data import Dataset, DataLoader\nimport torch as t\nimport torch.nn as nn\nimport torch.nn.functional as F\n\ndevice = 'cuda' if t.cuda.is_available() else 'cpu'","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-03-24T14:11:04.614968Z","iopub.execute_input":"2024-03-24T14:11:04.615309Z","iopub.status.idle":"2024-03-24T14:11:08.505000Z","shell.execute_reply.started":"2024-03-24T14:11:04.615281Z","shell.execute_reply":"2024-03-24T14:11:08.504229Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_path = '/kaggle/input/hms-harmful-brain-activity-classification/test_eegs/'\nBASE_PATH = \"/kaggle/input/hms-harmful-brain-activity-classification\"\nclass_names = ['Seizure', 'LPD', 'GPD', 'LRDA','GRDA', 'Other']\nFEATS_FOR_REAL = ['Fp1', 'F3', 'C3', 'P3', 'F7', 'T3', 'T5', 'O1', 'Fz', 'Cz', 'Pz', 'Fp2', 'F4', 'C4', 'P4', 'F8', 'T4', 'T6', 'O2', 'EKG']\nGROUPS_IDS = [\n    [0, 1, 2, 3, 7],\n    [0, 4, 5, 6, 7],\n    [11, 12, 13, 14, 18],\n    [11, 15, 16, 17, 18],\n]","metadata":{"execution":{"iopub.status.busy":"2024-03-24T14:11:08.506185Z","iopub.execute_input":"2024-03-24T14:11:08.506580Z","iopub.status.idle":"2024-03-24T14:11:08.512775Z","shell.execute_reply.started":"2024-03-24T14:11:08.506556Z","shell.execute_reply":"2024-03-24T14:11:08.511846Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df = pd.read_csv(f'{BASE_PATH}/test.csv')\n\n# test_df['eeg_path'] = f'{BASE_PATH}/test_eegs/'+test_df['eeg_id'].astype(str)+'.parquet'\n# test_df['spec_path'] = f'{BASE_PATH}/test_spectrograms/'+test_df['spectrogram_id'].astype(str)+'.parquet'","metadata":{"execution":{"iopub.status.busy":"2024-03-24T14:11:08.514666Z","iopub.execute_input":"2024-03-24T14:11:08.514939Z","iopub.status.idle":"2024-03-24T14:11:08.535399Z","shell.execute_reply.started":"2024-03-24T14:11:08.514916Z","shell.execute_reply":"2024-03-24T14:11:08.534685Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class EEGTestDataset(Dataset):\n    def __init__(self):\n        super().__init__()\n        self.dataframe = test_df\n\n    def __len__(self):\n        return len(self.dataframe)\n\n    def __getitem__(self, idx):\n        row = self.dataframe.iloc[idx]\n        eeg_id = row['eeg_id']\n        parq_path = f'{test_path}{eeg_id}.parquet'\n        eeg = pd.read_parquet(parq_path)\n        eeg = eeg.ffill(axis=0)\n        eeg = eeg.fillna(0)\n        samples = t.tensor(eeg[FEATS_FOR_REAL].values)\n        \n        return samples","metadata":{"execution":{"iopub.status.busy":"2024-03-24T14:11:08.536362Z","iopub.execute_input":"2024-03-24T14:11:08.536620Z","iopub.status.idle":"2024-03-24T14:11:08.542918Z","shell.execute_reply.started":"2024-03-24T14:11:08.536598Z","shell.execute_reply":"2024-03-24T14:11:08.542053Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch_size = 150\ntest_dataset = EEGTestDataset()\ntest_dataloader = DataLoader(test_dataset, batch_size=batch_size, num_workers=4, prefetch_factor=2)","metadata":{"execution":{"iopub.status.busy":"2024-03-24T14:11:08.544005Z","iopub.execute_input":"2024-03-24T14:11:08.544285Z","iopub.status.idle":"2024-03-24T14:11:08.552904Z","shell.execute_reply.started":"2024-03-24T14:11:08.544263Z","shell.execute_reply":"2024-03-24T14:11:08.552067Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ConvBlock(nn.Module):\n    def __init__(self, d_in, d_out, kernel_size, drop):\n        super().__init__()\n        self.model = nn.Sequential(\n            nn.Conv1d(d_in, d_out, kernel_size=kernel_size, padding='same', stride=1),\n            nn.ReLU(),\n            nn.Dropout(drop),\n            nn.Conv1d(d_out, d_out, kernel_size=kernel_size, padding='same', stride=1),\n            nn.ReLU(),\n            nn.Dropout(drop),\n            nn.Conv1d(d_out, d_out, kernel_size=kernel_size, padding='same', stride=1),\n            nn.ReLU(),\n            nn.Dropout(drop),\n            nn.MaxPool1d(kernel_size=2, stride=2, padding=0), # reduce sequence size by 2\n        )\n    def forward(self, x):\n        # TODO: add skip for training speed\n        return self.model(x)\n        \nclass Model(nn.Module):\n    def __init__(self, in_channels=5, gru_hidden_size=128, drop=0.1):\n        super().__init__()\n        self.d_split = len(GROUPS_IDS)\n        self.pre_out = in_channels * 4\n        self.gru_hidden_size = gru_hidden_size\n        \n        self.pre_process = nn.Sequential(\n            nn.BatchNorm1d(in_channels, momentum=None),\n            # use conv1d as a denoiser\n            # block 1\n            ConvBlock(in_channels, in_channels * 2, kernel_size=3, drop=drop),\n            # nn.BatchNorm1d(in_channels * 2, momentum=None),\n            \n            # block 2\n            ConvBlock(in_channels * 2, in_channels * 4, kernel_size=5, drop=drop),\n            # nn.BatchNorm1d(self.pre_out, momentum=None),\n\n            # block 3\n            ConvBlock(in_channels * 4, in_channels * 4, kernel_size=7, drop=drop),\n            # nn.BatchNorm1d(self.pre_out, momentum=None),\n        )\n        \n        # TODO: add a learnable first state for GRU or check what is the default\n        self.gru = nn.GRU(self.pre_out, self.gru_hidden_size, num_layers=1, batch_first=True, bidirectional=True)\n\n        self.post_gru = nn.Sequential(\n            nn.Linear(self.gru_hidden_size * 2, self.gru_hidden_size),\n            nn.ReLU(),\n            nn.Linear(self.gru_hidden_size, self.gru_hidden_size),\n        )\n\n        self.head = nn.Sequential(\n            nn.Linear(self.gru_hidden_size * self.d_split, self.gru_hidden_size * 2),\n            nn.ReLU(),\n            nn.Dropout(drop),\n            nn.Linear(self.gru_hidden_size * 2, 6)\n        )\n\n    def forward(self, x: ('batch', 'seq', 'channel')):\n        # separate the input into 4 splits (LP, LL, RP, RR)\n        splits = [x[:, :, group] for group in GROUPS_IDS]\n        # fold it into batch so we can run in parallel\n        x = einops.rearrange(t.stack(splits, dim=0), 'group batch seq channel -> (group batch) seq channel')\n\n        # pre_process: (batch, channel, seq) → (batch / 4, channel * 4, seq)\n        x = x.permute((0, 2, 1))\n        x = self.pre_process(x)\n        x = x.permute((0, 2, 1))\n\n        # GRU: (batch, seq, input_size), [(2 * num_layers, batch, hidden_size)] → (batch, seq, 2 * hidden_size)\n        x, _ = self.gru(x)\n        x = x[:, -1, :]\n\n        # MLP post GRU\n        x = self.post_gru(x)\n\n        # unfold the splits\n        x = einops.rearrange(x, '(group batch) hidden -> batch (hidden group)', group=self.d_split)\n\n        # head: (batch, 2 * hidden_size) → (batch, 6)\n        x = self.head(x)\n\n        # out: → (batch, 6)\n        return x","metadata":{"execution":{"iopub.status.busy":"2024-03-24T14:11:08.554340Z","iopub.execute_input":"2024-03-24T14:11:08.554644Z","iopub.status.idle":"2024-03-24T14:11:08.572135Z","shell.execute_reply.started":"2024-03-24T14:11:08.554622Z","shell.execute_reply":"2024-03-24T14:11:08.571327Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = Model().to(device)\nmodel.load_state_dict(t.load('/kaggle/input/tmp63271321/gru-4-splits_2024-03-24_15h03.pt', map_location=device))","metadata":{"execution":{"iopub.status.busy":"2024-03-24T14:11:08.573690Z","iopub.execute_input":"2024-03-24T14:11:08.574285Z","iopub.status.idle":"2024-03-24T14:11:08.946305Z","shell.execute_reply.started":"2024-03-24T14:11:08.574255Z","shell.execute_reply":"2024-03-24T14:11:08.945414Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Validation Stuff kthx\n\nres = []\nwith t.no_grad():\n    model.train()\n\n    for batch in test_dataloader:\n        batch = batch.to(device)\n        prob = model(batch).softmax(-1)\n        res.append(prob.detach().cpu())\n\n    res = t.cat(res, dim=0)","metadata":{"execution":{"iopub.status.busy":"2024-03-24T14:11:08.947610Z","iopub.execute_input":"2024-03-24T14:11:08.947978Z","iopub.status.idle":"2024-03-24T14:11:12.483839Z","shell.execute_reply.started":"2024-03-24T14:11:08.947944Z","shell.execute_reply":"2024-03-24T14:11:12.482914Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pred_df = test_df[[\"eeg_id\"]].copy()\ntarget_cols = [x.lower()+'_vote' for x in class_names]\n\npred_df[target_cols] = res.tolist()\n\nsub_df = pd.read_csv(f'{BASE_PATH}/sample_submission.csv')\nsub_df = sub_df[[\"eeg_id\"]].copy()\nsub_df = sub_df.merge(pred_df, on=\"eeg_id\", how=\"left\")\nsub_df.to_csv(\"submission.csv\", index=False)\nsub_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-03-24T14:11:12.486260Z","iopub.execute_input":"2024-03-24T14:11:12.486681Z","iopub.status.idle":"2024-03-24T14:11:12.519975Z","shell.execute_reply.started":"2024-03-24T14:11:12.486656Z","shell.execute_reply":"2024-03-24T14:11:12.519096Z"},"trusted":true},"execution_count":null,"outputs":[]}]}