{"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":[{"sourceType":"competition","sourceId":59093,"databundleVersionId":7469972},{"sourceType":"datasetVersion","sourceId":7908175,"datasetId":4645532,"databundleVersionId":8015693},{"sourceType":"datasetVersion","sourceId":7899529,"datasetId":4639130,"databundleVersionId":8006480},{"sourceType":"datasetVersion","sourceId":7960540,"datasetId":4646663,"databundleVersionId":8070283}],"dockerImageVersionId":30674,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Karina's Data Loading Code","metadata":{}},{"cell_type":"code","source":" !pip3 install torch torchvision torchaudio --extra-index-url https://download.pytorch.org/whl/cu115\n\n print(\"Done.\")","metadata":{"execution":{"iopub.status.busy":"2024-03-27T23:59:23.627514Z","iopub.execute_input":"2024-03-27T23:59:23.628338Z","iopub.status.idle":"2024-03-27T23:59:36.814924Z","shell.execute_reply.started":"2024-03-27T23:59:23.628304Z","shell.execute_reply":"2024-03-27T23:59:36.813512Z"},"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\nfrom sklearn.model_selection import train_test_split\n\nfrom datetime import datetime\n\nimport torch as t\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\n\nprint(\"done\")","metadata":{"execution":{"iopub.status.busy":"2024-03-27T23:59:36.817607Z","iopub.execute_input":"2024-03-27T23:59:36.818027Z","iopub.status.idle":"2024-03-27T23:59:36.826726Z","shell.execute_reply.started":"2024-03-27T23:59:36.817990Z","shell.execute_reply":"2024-03-27T23:59:36.825675Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Utility Functions\nimport gc \ndef GC():\n    gc.collect()\n    t.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2024-03-27T23:59:38.863046Z","iopub.execute_input":"2024-03-27T23:59:38.863416Z","iopub.status.idle":"2024-03-27T23:59:38.868715Z","shell.execute_reply.started":"2024-03-27T23:59:38.863389Z","shell.execute_reply":"2024-03-27T23:59:38.867683Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Define Paths\nBASE_PATH = \"/kaggle/input/hms-harmful-brain-activity-classification\"\nTRAIN_EEG_PATH = f'{BASE_PATH}/train_eegs'\n\nWORKING_BASE_PATH = \"/kaggle/working\"\nPROCESSED_EEG_PATH = f'{WORKING_BASE_PATH}/train_eegs_processed'\nos.makedirs(PROCESSED_EEG_PATH, exist_ok=True)","metadata":{"execution":{"iopub.status.busy":"2024-03-27T23:59:40.963000Z","iopub.execute_input":"2024-03-27T23:59:40.963656Z","iopub.status.idle":"2024-03-27T23:59:40.969063Z","shell.execute_reply.started":"2024-03-27T23:59:40.963626Z","shell.execute_reply":"2024-03-27T23:59:40.968069Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Get CSV data\nLABELED_DATA_df = pd.read_csv(f'{BASE_PATH}/train.csv')\nUNLABELED_DATA_df = pd.read_csv(f'{BASE_PATH}/test.csv')\n\n# Define target columns(votes for each activity pattern)\nLABEL_COLS = LABELED_DATA_df.columns[-6:]\n\n#---- CELL OUTPUT ------------:\nprint(f\"\\nLABEL COLUMNS:\")\nfor col in LABEL_COLS:\n    print(\"  -\",col)\n    \nprint('\\n\\nLABELED DATA:', LABELED_DATA_df.shape, '\\n' )\ndisplay(LABELED_DATA_df.head() )\n\nprint('\\n\\nUNLABELED DATA:', UNLABELED_DATA_df.shape, '\\n' )\ndisplay(UNLABELED_DATA_df.head() )","metadata":{"execution":{"iopub.status.busy":"2024-03-27T23:59:46.498007Z","iopub.execute_input":"2024-03-27T23:59:46.498473Z","iopub.status.idle":"2024-03-27T23:59:46.733741Z","shell.execute_reply.started":"2024-03-27T23:59:46.498444Z","shell.execute_reply":"2024-03-27T23:59:46.732686Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#SELECT SUBSET OF TRAINING DATA\n\nLABELED_DATA_df = LABELED_DATA_df.iloc[:1000]\nprint(len(LABELED_DATA_df))","metadata":{"execution":{"iopub.status.busy":"2024-03-27T23:59:50.036991Z","iopub.execute_input":"2024-03-27T23:59:50.037381Z","iopub.status.idle":"2024-03-27T23:59:50.043595Z","shell.execute_reply.started":"2024-03-27T23:59:50.037350Z","shell.execute_reply":"2024-03-27T23:59:50.042500Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Find all EEG snippets and save as individual files for faster loading by dataloader\n\nrun_preprocessing = True\n\ndef preprocess_and_save(df, eeg_path, output_path):\n    old_parq_path = ''\n    for idx, row in df.iterrows():\n        eeg_id = row['eeg_id']\n        eeg_sub_id = row['eeg_sub_id']\n        processed_path = os.path.join(output_path, f'{eeg_id}_{eeg_sub_id}_processed.npy')\n        if ~os.path.exists(processed_path): #create process eeg file if does not already exist\n            new_parq_path = f'{eeg_path}/{eeg_id}.parquet'\n            if new_parq_path != old_parq_path: #load new parquet file if current eeg is not draw from last opened parquet file\n                eeg_df = pd.read_parquet(new_parq_path)\n                old_parq_path = new_parq_path\n\n            # Crop and clean the data\n            start_time_second = row['eeg_label_offset_seconds']\n            dur_s = 50\n            samp_rate = 200\n            offset_dp = int(start_time_second * samp_rate)\n            dur_dp = dur_s * samp_rate\n            snippet_df = eeg_df.iloc[offset_dp:offset_dp+dur_dp]\n            assert(snippet_df.shape[0] == dur_dp)\n            snippet_df = snippet_df.ffill(axis=0).fillna(0)\n            #Save\n            np.save(processed_path, snippet_df.values)\n        if idx%1000 <1:\n            print(f\"{idx} of {len(df)} processed\")\n\nif run_preprocessing:\n    preprocess_and_save(LABELED_DATA_df, TRAIN_EEG_PATH, PROCESSED_EEG_PATH)\n    output = \"EEG data preprocessed to extract relevant intervals for each dp\"\nelse:\n    output = \"preprocessing turned off (presumably because preprocessed files already saved)\"\n\n#---- CELL OUTPUT ------------:\nprint(output)","metadata":{"execution":{"iopub.status.busy":"2024-03-27T23:59:51.957218Z","iopub.execute_input":"2024-03-27T23:59:51.957598Z","iopub.status.idle":"2024-03-27T23:59:55.115457Z","shell.execute_reply.started":"2024-03-27T23:59:51.957566Z","shell.execute_reply":"2024-03-27T23:59:55.114329Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create Dataset Class\n\nclass TrainOrValDataset(Dataset):\n    def __init__(self, df):\n        self.df = df\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx, label_cols = LABEL_COLS, eeg_path = TRAIN_EEG_PATH):    \n        \n        row_series = self.df.iloc[idx]\n        \n        #GET FEATURE DATA (EEG channels)\n        #Read in EEG data\n        eeg_id = row_series['eeg_id']\n        eeg_sub_id = row_series['eeg_sub_id']\n        \n        #Load EEG clip\n        processed_eeg_path = os.path.join(PROCESSED_EEG_PATH, f'{eeg_id}_{eeg_sub_id}_processed.npy')\n        eeg_nparray = np.load(processed_eeg_path)\n        \n        #GET LABEL DATA expert votes)\n        label = row_series[label_cols].values.astype(np.float64)\n        label = label/np.sum(label) #convert to pct of total votes\n      \n        return t.tensor(eeg_nparray, dtype=t.float32), t.tensor(label,dtype=t.float64)","metadata":{"execution":{"iopub.status.busy":"2024-03-27T23:59:57.941086Z","iopub.execute_input":"2024-03-27T23:59:57.941482Z","iopub.status.idle":"2024-03-27T23:59:57.950309Z","shell.execute_reply.started":"2024-03-27T23:59:57.941454Z","shell.execute_reply":"2024-03-27T23:59:57.949340Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create datasets (train, val, test)\n\nTRAIN_df, VAL_df = train_test_split(LABELED_DATA_df, test_size = .1, random_state = 1)\n\nTRAIN_ds = TrainOrValDataset(TRAIN_df)\nVAL_ds = TrainOrValDataset(VAL_df)\nTEST_ds = TrainOrValDataset(UNLABELED_DATA_df)\n\n#---- CELL OUTPUT ------------:\nprint(f\"Train size: {len(TRAIN_ds)}\")\nprint(f\"Val size: {len(VAL_ds)}\")\nprint(f\"Test size: {len(TEST_ds)}*\")\nprint(\"\"\"\n* NOTE: test size is 1 because this is a code submission competition.\nThey provide one example that can be used to setup submission code to take actual\ntest data, which we don't get to see in advance\"\"\")","metadata":{"execution":{"iopub.status.busy":"2024-03-28T00:00:00.491155Z","iopub.execute_input":"2024-03-28T00:00:00.492085Z","iopub.status.idle":"2024-03-28T00:00:00.502339Z","shell.execute_reply.started":"2024-03-28T00:00:00.492046Z","shell.execute_reply":"2024-03-28T00:00:00.501326Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create Dataloaders\nfrom torch.utils.data import DataLoader\n\nBATCH_SZ = 64\nN_WORKERS = 1\nSHUFFLE_STATE = False #Should be set to True, but ran out of space to load all data so we set to false to not try to pull later data\n\nTRAIN_dl = DataLoader(TRAIN_ds, batch_size=BATCH_SZ, shuffle=SHUFFLE_STATE, num_workers=N_WORKERS)\nVAL_dl = DataLoader(VAL_ds, batch_size=BATCH_SZ, shuffle=SHUFFLE_STATE, num_workers=N_WORKERS)\nTEST_dl = DataLoader(TEST_ds, batch_size=BATCH_SZ, shuffle=SHUFFLE_STATE, num_workers=N_WORKERS)\n\n#---- CELL OUTPUT ------------:\nprint(\"data loaders created\")","metadata":{"execution":{"iopub.status.busy":"2024-03-28T00:00:02.657219Z","iopub.execute_input":"2024-03-28T00:00:02.657939Z","iopub.status.idle":"2024-03-28T00:00:02.666118Z","shell.execute_reply.started":"2024-03-28T00:00:02.657902Z","shell.execute_reply":"2024-03-28T00:00:02.664980Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Regis Code\n","metadata":{}},{"cell_type":"code","source":"%pip install /kaggle/input/tmp63271321/einops-0.7.0-py3-none-any.whl\n","metadata":{"execution":{"iopub.status.busy":"2024-03-28T00:00:08.713094Z","iopub.execute_input":"2024-03-28T00:00:08.713782Z","iopub.status.idle":"2024-03-28T00:00:21.926598Z","shell.execute_reply.started":"2024-03-28T00:00:08.713749Z","shell.execute_reply":"2024-03-28T00:00:21.925222Z"},"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\nfrom scipy.signal import butter, sosfilt\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-28T00:00:29.882836Z","iopub.execute_input":"2024-03-28T00:00:29.883357Z","iopub.status.idle":"2024-03-28T00:00:29.890350Z","shell.execute_reply.started":"2024-03-28T00:00:29.883322Z","shell.execute_reply":"2024-03-28T00:00:29.889207Z"},"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']\nTARGETS = ['seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']\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-28T00:00:32.382055Z","iopub.execute_input":"2024-03-28T00:00:32.382792Z","iopub.status.idle":"2024-03-28T00:00:32.390003Z","shell.execute_reply.started":"2024-03-28T00:00:32.382758Z","shell.execute_reply":"2024-03-28T00:00:32.389002Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df = pd.read_csv(f'{BASE_PATH}/test.csv')\n# test_df = pd.read_csv(f'{BASE_PATH}/train.csv').iloc[:5]\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-28T00:00:35.025925Z","iopub.execute_input":"2024-03-28T00:00:35.026345Z","iopub.status.idle":"2024-03-28T00:00:35.035725Z","shell.execute_reply.started":"2024-03-28T00:00:35.026313Z","shell.execute_reply":"2024-03-28T00:00:35.034814Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def logify(data):\n#     # numpy absolute value\n#     log_data = np.log1p(np.abs(data))\n#     log_data[data < 0] *= -1\n#     return log_data\n\n# def filters(data):\n#     return logify(data)\n\ndef clip(data, bound=300):\n    return np.clip(data, -bound, bound)\n\ndef robust_norm(data):\n    median = np.median(data, axis=0)\n    q75, q25 = np.percentile(data, [75 ,25], axis=0)\n    iqr = q75 - q25\n    iqr[iqr < 1e-6] = 1e-6 # numerical stability\n    return (data - median) / iqr\n\ndef band_filter(data, low=1, high=70, fs=200, order=4):\n    sos = butter(N=order, Wn=[low, high], btype='bandpass', fs=fs, output='sos')\n    return sosfilt(sos, data, axis=0)\n\ndef filters(data):\n    data = band_filter(data)\n    data = robust_norm(data)\n    return data\n\nclass 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#         parq_path = f'/kaggle/input/hms-harmful-brain-activity-classification/train_eegs/{eeg_id}.parquet'\n        eeg = pd.read_parquet(parq_path)\n        eeg = eeg.ffill(axis=0)\n        eeg = eeg.fillna(0)\n        # XXX\n        filtered_eeg = filters(eeg[FEATS_FOR_REAL].values)\n        samples = t.tensor(filtered_eeg).float()\n#         samples = t.tensor(eeg[FEATS_FOR_REAL].values)\n        \n        return samples","metadata":{"execution":{"iopub.status.busy":"2024-03-28T00:00:38.894174Z","iopub.execute_input":"2024-03-28T00:00:38.894574Z","iopub.status.idle":"2024-03-28T00:00:38.907481Z","shell.execute_reply.started":"2024-03-28T00:00:38.894541Z","shell.execute_reply":"2024-03-28T00:00:38.906434Z"},"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-28T00:00:42.662273Z","iopub.execute_input":"2024-03-28T00:00:42.662689Z","iopub.status.idle":"2024-03-28T00:00:42.668741Z","shell.execute_reply.started":"2024-03-28T00:00:42.662657Z","shell.execute_reply":"2024-03-28T00:00:42.667551Z"},"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        \n# class Model(nn.Module):\n#     def __init__(self, in_channels=5, gru_hidden_size=128, drop=0.2):\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(in_channels * 4, momentum=None),\n\n#             # block 3\n#             ConvBlock(in_channels * 4, in_channels * 4, kernel_size=7, drop=drop),\n#             nn.BatchNorm1d(in_channels * 4, momentum=None), # re-enable one more to force training ?\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-27T23:52:26.039070Z","iopub.status.idle":"2024-03-27T23:52:26.039467Z","shell.execute_reply.started":"2024-03-27T23:52:26.039284Z","shell.execute_reply":"2024-03-27T23:52:26.039300Z"},"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.InstanceNorm1d(d_out),\n            nn.ReLU(),\n            nn.Dropout(drop),\n            nn.Conv1d(d_out, d_out, kernel_size=kernel_size, padding='same', stride=1),\n            nn.InstanceNorm1d(d_out),\n            nn.ReLU(),\n            nn.Dropout(drop),\n            nn.Conv1d(d_out, d_out, kernel_size=kernel_size, padding='same', stride=1),\n            nn.InstanceNorm1d(d_out),\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=4, gru_hidden_size=128, drop=0.2):\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.InstanceNorm1d(in_channels),\n            # nn.LayerNorm(normalized_shape=[in_channels, 10000]),\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            # nn.LayerNorm(normalized_shape=[in_channels * 2, 5000]),\n            # -- block 2 --\n            ConvBlock(in_channels * 2, in_channels * 4, kernel_size=5, drop=drop),\n            # nn.BatchNorm1d(in_channels * 4, momentum=None),\n            # nn.LayerNorm(normalized_shape=[in_channels * 4, 2500]),\n            # -- block 3 --\n            ConvBlock(in_channels * 4, in_channels * 4, kernel_size=7, drop=drop),\n            # nn.BatchNorm1d(in_channels * 4, momentum=None), # re-enable one more to force training ?\n            # nn.LayerNorm(normalized_shape=[in_channels * 4, 1250]),\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.BatchNorm1d(self.gru_hidden_size * 2, momentum=None),\n            nn.LayerNorm(normalized_shape=[self.gru_hidden_size * 2]),\n            nn.Dropout(drop),\n            nn.Linear(self.gru_hidden_size * 2, self.gru_hidden_size),\n            nn.ReLU(),\n            nn.Dropout(drop),\n            nn.Linear(self.gru_hidden_size, self.gru_hidden_size),\n        )\n\n        self.head = nn.Sequential(\n            # nn.BatchNorm1d(self.gru_hidden_size * self.d_split, momentum=None),\n            nn.LayerNorm(normalized_shape=[self.gru_hidden_size * self.d_split]),\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 montage(self, x):\n        splits = [x[:, :, group] for group in GROUPS_IDS]\n        splits = [s[:, :, :-1] - s[:, :, 1:] for s in splits]\n        return einops.rearrange(t.stack(splits, dim=0), 'group batch seq channel -> (group batch) seq channel')\n\n    def forward(self, x: ('batch', 'seq', 'channel')):\n        # separate the input into 4 montages (LP, LL, RP, RR)\n        x = self.montage(x)\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        # 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        # MLP post GRU\n        x = self.post_gru(x)\n        # unfold the splits\n        x = einops.rearrange(x, '(group batch) hidden -> batch (hidden group)', group=self.d_split)\n        # head: (batch, 2 * hidden_size) → (batch, 6)\n        x = self.head(x)\n        # out: → (batch, 6)\n        return x","metadata":{"execution":{"iopub.status.busy":"2024-03-28T00:00:48.002080Z","iopub.execute_input":"2024-03-28T00:00:48.002468Z","iopub.status.idle":"2024-03-28T00:00:48.026744Z","shell.execute_reply.started":"2024-03-28T00:00:48.002437Z","shell.execute_reply":"2024-03-28T00:00:48.025748Z"},"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-27_14h50.pt', map_location=device))","metadata":{"execution":{"iopub.status.busy":"2024-03-28T00:00:54.281288Z","iopub.execute_input":"2024-03-28T00:00:54.281670Z","iopub.status.idle":"2024-03-28T00:00:54.311523Z","shell.execute_reply.started":"2024-03-28T00:00:54.281639Z","shell.execute_reply":"2024-03-28T00:00:54.310512Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Predict\n\nlabels=[]\nprobs = []\nlosses=[]\nwith t.no_grad():\n    model.eval()\n\n    for input_batch, label_batch in TRAIN_dl:\n        input_batch = input_batch.to(device)\n        for ind, sample in enumerate(input_batch):\n            sample = sample.unsqueeze(0) #one sample, but still need to have batch dimension\n            label = label_batch[ind]\n            prob = model(sample).softmax(-1)\n            kl_loss = nn.KLDivLoss(reduction=\"batchmean\")\n            loss = kl_loss(prob, label.to(device))\n            probs.append(prob.detach().cpu())\n            losses.append(loss)\n            labels.append(label)\n\n    #probs = t.cat(probs, dim=0)\n    \nprint(\"Probabilities Calculated\")","metadata":{"execution":{"iopub.status.busy":"2024-03-28T01:00:02.524492Z","iopub.execute_input":"2024-03-28T01:00:02.524898Z","iopub.status.idle":"2024-03-28T01:00:10.161080Z","shell.execute_reply.started":"2024-03-28T01:00:02.524864Z","shell.execute_reply":"2024-03-28T01:00:10.159869Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# loop through examples\n\n\ncorrect_counts= np.zeros((11, 6))\nall_counts = np.zeros((11, 6))\ncum_loss = np.zeros((11, 6))\nfor ind, prob in enumerate(probs):\n    label = labels[ind]\n    loss = losses[ind]\n    \n    label_activity = np.argmax(label)\n    pred_activity = np.argmax(prob)\n    \n    label_prob = np.round(max(label),1)\n    pred_prob = np.round(max(prob),1) #round prob to one decimal place\n    \n    row = int(label_prob*10)\n    col = label_activity\n    all_counts[row][col] += 1\n    cum_loss[row][col] += loss\n    if label_activity == pred_activity:\n        correct_counts[row][col] += 1\n\n#epsilon = .0000000001\navg_loss = cum_loss/(all_counts)\npct_correct = 100* correct_counts/(all_counts)\n\n\nprint(all_counts)\nprint(\"-----------------------\")\nprint(np.round(avg_loss,1))\n#print(np.round(pct_correct,0))\n\n    \n","metadata":{"execution":{"iopub.status.busy":"2024-03-28T01:15:21.654887Z","iopub.execute_input":"2024-03-28T01:15:21.655839Z","iopub.status.idle":"2024-03-28T01:15:22.031360Z","shell.execute_reply.started":"2024-03-28T01:15:21.655781Z","shell.execute_reply":"2024-03-28T01:15:22.030364Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Make submission file\n\nif False:\n    sub = test_df[[\"eeg_id\"]].copy()\n    sub[TARGETS] = res\n    sub.to_csv('submission.csv',index=False)\n\n    print('Submission shape',sub.shape)\n    sub.head()","metadata":{"execution":{"iopub.status.busy":"2024-03-27T23:58:55.918445Z","iopub.execute_input":"2024-03-27T23:58:55.919370Z","iopub.status.idle":"2024-03-27T23:58:55.924994Z","shell.execute_reply.started":"2024-03-27T23:58:55.919332Z","shell.execute_reply":"2024-03-27T23:58:55.923915Z"},"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":[]}]}