{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","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"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"},{"sourceId":7392733,"sourceType":"datasetVersion","datasetId":4297749},{"sourceId":7465251,"sourceType":"datasetVersion","datasetId":4317718},{"sourceId":170600050,"sourceType":"kernelVersion"}],"dockerImageVersionId":30636,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Trying to implement a path signature model.","metadata":{}},{"cell_type":"markdown","source":"# Load Train Data","metadata":{}},{"cell_type":"code","source":"import pandas as pd, numpy as np, os\nimport matplotlib.pyplot as plt\n\ntrain = pd.read_csv('/kaggle/input/hms-harmful-brain-activity-classification/train.csv')\nprint( train.shape )\ndisplay( train.head() )\n\n# CHOICE TO CREATE OR LOAD EEGS FROM NOTEBOOK VERSION 1\nCREATE_EEGS = False\nTRAIN_MODEL = False","metadata":{"execution":{"iopub.status.busy":"2024-04-08T15:12:53.307647Z","iopub.execute_input":"2024-04-08T15:12:53.308072Z","iopub.status.idle":"2024-04-08T15:12:54.147853Z","shell.execute_reply.started":"2024-04-08T15:12:53.308037Z","shell.execute_reply":"2024-04-08T15:12:54.146650Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Raw EEG Features","metadata":{}},{"cell_type":"code","source":"df = pd.read_parquet('/kaggle/input/hms-harmful-brain-activity-classification/train_eegs/1000913311.parquet')\nFEATS = df.columns\nprint(f'There are {len(FEATS)} raw eeg features')\nprint( list(FEATS) )","metadata":{"execution":{"iopub.status.busy":"2024-04-08T15:12:54.150119Z","iopub.execute_input":"2024-04-08T15:12:54.151220Z","iopub.status.idle":"2024-04-08T15:12:54.328563Z","shell.execute_reply.started":"2024-04-08T15:12:54.151171Z","shell.execute_reply":"2024-04-08T15:12:54.327713Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('We will use the following subset of raw EEG features:')\nFEATS = ['Fp1','T3','C3','O1','Fp2','C4','T4','O2']\nFEAT2IDX = {x:y for x,y in zip(FEATS,range(len(FEATS)))}\nprint( list(FEATS) )","metadata":{"execution":{"iopub.status.busy":"2024-04-08T15:12:54.333263Z","iopub.execute_input":"2024-04-08T15:12:54.336023Z","iopub.status.idle":"2024-04-08T15:12:54.346165Z","shell.execute_reply.started":"2024-04-08T15:12:54.335977Z","shell.execute_reply":"2024-04-08T15:12:54.345320Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def eeg_from_parquet(parquet_path, display=False):\n    \n    # EXTRACT MIDDLE 50 SECONDS\n    eeg = pd.read_parquet(parquet_path, columns=FEATS)\n    rows = len(eeg)\n    offset = (rows-10_000)//2\n    eeg = eeg.iloc[offset:offset+10_000]\n    \n    if display: \n        plt.figure(figsize=(10,5))\n        offset = 0\n    \n    # CONVERT TO NUMPY\n    data = np.zeros((10_000,len(FEATS)))\n    for j,col in enumerate(FEATS):\n        \n        # FILL NAN\n        x = eeg[col].values.astype('float32')\n        m = np.nanmean(x)\n        if np.isnan(x).mean()<1: x = np.nan_to_num(x,nan=m)\n        else: x[:] = 0\n            \n        data[:,j] = x\n        \n        if display: \n            if j!=0: offset += x.max()\n            plt.plot(range(10_000),x-offset,label=col)\n            offset -= x.min()\n            \n    if display:\n        plt.legend()\n        name = parquet_path.split('/')[-1]\n        name = name.split('.')[0]\n        plt.title(f'EEG {name}',size=16)\n        plt.show()\n        \n    return data","metadata":{"execution":{"iopub.status.busy":"2024-04-08T15:12:54.350987Z","iopub.execute_input":"2024-04-08T15:12:54.351658Z","iopub.status.idle":"2024-04-08T15:12:54.366011Z","shell.execute_reply.started":"2024-04-08T15:12:54.351621Z","shell.execute_reply":"2024-04-08T15:12:54.364813Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\n\nall_eegs = {}\nDISPLAY = 4\nEEG_IDS = train.eeg_id.unique()\nPATH = '/kaggle/input/hms-harmful-brain-activity-classification/train_eegs/'\n\nfor i,eeg_id in enumerate(EEG_IDS):\n    if (i%100==0)&(i!=0): print(i,', ',end='') \n    \n    # SAVE EEG TO PYTHON DICTIONARY OF NUMPY ARRAYS\n    data = eeg_from_parquet(f'{PATH}{eeg_id}.parquet', display=i<DISPLAY)              \n    all_eegs[eeg_id] = data\n    \n    if i==DISPLAY:\n        if CREATE_EEGS:\n            print(f'Processing {train.eeg_id.nunique()} eeg parquets... ',end='')\n        else:\n            print(f'Reading {len(EEG_IDS)} eeg NumPys from disk.')\n            break\n            \nif CREATE_EEGS: \n    np.save('eegs',all_eegs)\nelse:\n    all_eegs = np.load('/kaggle/input/brain-eegs/eegs.npy',allow_pickle=True).item()","metadata":{"execution":{"iopub.status.busy":"2024-04-08T15:12:54.367560Z","iopub.execute_input":"2024-04-08T15:12:54.368298Z","iopub.status.idle":"2024-04-08T15:14:47.992983Z","shell.execute_reply.started":"2024-04-08T15:12:54.368253Z","shell.execute_reply":"2024-04-08T15:14:47.991053Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Deduplicate Train EEG Id","metadata":{}},{"cell_type":"code","source":"# LOAD TRAIN \ndf = pd.read_csv('/kaggle/input/hms-harmful-brain-activity-classification/train.csv')\nTARGETS = df.columns[-6:]\nTARS = {'Seizure':0, 'LPD':1, 'GPD':2, 'LRDA':3, 'GRDA':4, 'Other':5}\nTARS2 = {x:y for y,x in TARS.items()}\n\ntrain = df.groupby('eeg_id')[['patient_id']].agg('first')\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']].agg('first')\ntrain['target'] = tmp\n\ntrain = train.reset_index()\ntrain = train.loc[train.eeg_id.isin(EEG_IDS)]\nprint('Train Data with unique eeg_id shape:', train.shape )\ntrain.head()","metadata":{"execution":{"iopub.status.busy":"2024-04-08T15:14:47.995859Z","iopub.execute_input":"2024-04-08T15:14:47.996304Z","iopub.status.idle":"2024-04-08T15:14:48.475034Z","shell.execute_reply.started":"2024-04-08T15:14:47.996263Z","shell.execute_reply":"2024-04-08T15:14:48.473814Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Filters","metadata":{}},{"cell_type":"code","source":"from scipy.signal import firwin, lfilter, iirnotch\n\ndef fir_lowpass_filter(data, cutoff_low=0.5, cutoff_high=50, sampling_rate=200, numtaps=101):\n    # Calculate the Nyquist frequency\n    nyquist_rate = sampling_rate / 2.0\n    \n    # Create the filter coefficients (taps) for an FIR filter\n    taps = firwin(numtaps, [cutoff_low, cutoff_high], nyq=nyquist_rate, pass_zero=False)\n    \n    # Apply the filter to the data using lfilter\n    a, b = iirnotch(60,30, fs=200)\n    filtered_data_tmp = lfilter(b, a, data, axis=0)\n    filtered_data = lfilter(taps, 1.0, filtered_data_tmp, axis=0)  \n    \n    return filtered_data","metadata":{"execution":{"iopub.status.busy":"2024-04-08T15:14:48.476516Z","iopub.execute_input":"2024-04-08T15:14:48.476877Z","iopub.status.idle":"2024-04-08T15:14:49.624178Z","shell.execute_reply.started":"2024-04-08T15:14:48.476845Z","shell.execute_reply":"2024-04-08T15:14:49.622811Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Loader with FIR filter with a frequency range of 0.5-50 Hz","metadata":{}},{"cell_type":"code","source":"import torch\nfrom torch.utils.data import Dataset, DataLoader\n\nclass CustomDataset(Dataset):\n    def __init__(self, data, eegs, feat2idx, targets, downsample=1, mode='train'):\n        self.data = data\n        self.eegs = eegs\n        self.feat2idx = feat2idx\n        self.targets = targets\n        self.downsample = downsample\n        self.mode = mode\n\n    def __len__(self):\n        return len(self.data)\n\n    def __getitem__(self, idx):\n        row = self.data.iloc[idx]\n        X = np.zeros((10_000, 8), dtype='float32')\n        y = np.zeros(6, dtype='float32')\n        data = self.eegs[row.eeg_id]\n\n        # === Feature engineering ===\n        X[:,0] = data[:,self.feat2idx['Fp1']] - data[:,self.feat2idx['T3']]\n        X[:,1] = data[:,self.feat2idx['T3']] - data[:,self.feat2idx['O1']]\n\n        X[:,2] = data[:,self.feat2idx['Fp1']] - data[:,self.feat2idx['C3']]\n        X[:,3] = data[:,self.feat2idx['C3']] - data[:,self.feat2idx['O1']]\n\n        X[:,4] = data[:,self.feat2idx['Fp2']] - data[:,self.feat2idx['C4']]\n        X[:,5] = data[:,self.feat2idx['C4']] - data[:,self.feat2idx['O2']]\n\n        X[:,6] = data[:,self.feat2idx['Fp2']] - data[:,self.feat2idx['T4']]\n        X[:,7] = data[:,self.feat2idx['T4']] - data[:,self.feat2idx['O2']]\n\n        # === Standarize ===\n        X = np.clip(X,-1024, 1024)\n        X = np.nan_to_num(X, nan=0) / 32.0\n\n        # FIR Low-pass Filter\n        for i in range(X.shape[1]):\n            X[:, i] = fir_lowpass_filter(X[:, i])\n\n        # Downsampling\n        X = X[::self.downsample, :]\n        \n        \n        X = torch.from_numpy(X).float()\n\n        \n        if self.mode != 'test':\n            y = row[self.targets].values.astype(np.float32)\n            y = y / np.sum(y)  # Ensure it's a valid distribution\n            y = torch.from_numpy(y).float()\n            return X, y\n        \n        else:\n            return X\n\n    \ndataset = CustomDataset(data=train, eegs=all_eegs, feat2idx=FEAT2IDX, targets=TARGETS)\ndataloader = DataLoader(dataset, batch_size=32, shuffle=True, num_workers=4)","metadata":{"execution":{"iopub.status.busy":"2024-04-08T15:14:49.626764Z","iopub.execute_input":"2024-04-08T15:14:49.627145Z","iopub.status.idle":"2024-04-08T15:14:52.357177Z","shell.execute_reply.started":"2024-04-08T15:14:49.627102Z","shell.execute_reply":"2024-04-08T15:14:52.355917Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Display Data Loader","metadata":{}},{"cell_type":"code","source":"first_batch = next(iter(dataloader))\nX_batch, y_batch = first_batch\n\n# Visualizing the first four samples from the batch\nfor k in range(4):  # For each of the first four samples\n    plt.figure(figsize=(20, 4))\n    offset = 0\n    X_np = X_batch[k].numpy()  # Convert to numpy array for plotting\n    y_np = y_batch[k].numpy()  # Convert to numpy array\n    \n    for j in range(X_np.shape[-1]):  # For each feature/channel in the sample\n        if j != 0: \n            offset -= X_np[:, j].min()  # Calculate offset for better visualization\n        plt.plot(range(X_np.shape[0]), X_np[:, j] + offset, label=f'Feature {j+1}')\n        offset += X_np[:, j].max()\n    \n    # Constructing the target string for the title\n    tt = f'{y_np[0]:0.1f}'\n    for t in y_np[1:]:\n        tt += f', {t:0.1f}'\n    \n    plt.title(f'Sample {k+1} - Target = {tt}', size=14)\n    plt.legend()\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-04-08T15:14:52.358793Z","iopub.execute_input":"2024-04-08T15:14:52.360144Z","iopub.status.idle":"2024-04-08T15:14:58.149857Z","shell.execute_reply.started":"2024-04-08T15:14:52.360106Z","shell.execute_reply":"2024-04-08T15:14:58.148169Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Building model","metadata":{}},{"cell_type":"markdown","source":"## Path segmenting and signature transform","metadata":{}},{"cell_type":"code","source":"import torch.nn as nn","metadata":{"execution":{"iopub.status.busy":"2024-04-08T15:14:58.155284Z","iopub.execute_input":"2024-04-08T15:14:58.155723Z","iopub.status.idle":"2024-04-08T15:14:58.162483Z","shell.execute_reply.started":"2024-04-08T15:14:58.155676Z","shell.execute_reply":"2024-04-08T15:14:58.160560Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"! python -m pip install --no-index --find-links=../input/signatory -r ../input/signatory/requirements.txt","metadata":{"execution":{"iopub.status.busy":"2024-04-08T15:14:58.164046Z","iopub.execute_input":"2024-04-08T15:14:58.164441Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import signatory","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SegmentEEG(nn.Module):\n    def __init__(self, sampling_rate=50):\n        super(SegmentEEG, self).__init__()\n        self.sampling_rate = sampling_rate  # Number of samples in one second\n\n    def forward(self, x):\n        # X is of shape (batch_size, sequence_length, num_channels)\n        num_segments = x.shape[1] // self.sampling_rate\n        segmented_x = x[:, :num_segments*self.sampling_rate, :].reshape(x.shape[0], num_segments, self.sampling_rate, -1)\n        \n        # Create a time channel: evenly spaced values from 0 to 1 second, given the sampling rate\n        time_channel = torch.linspace(0, 1, steps=self.sampling_rate, device=x.device).unsqueeze(0).unsqueeze(0).unsqueeze(-1)\n        time_channel = time_channel.expand(segmented_x.size(0), segmented_x.size(1), self.sampling_rate, 1)\n        \n        # Concatenate the time channel to the segmented data\n        # New shape: (batch_size, num_segments, samples_per_segment, num_channels + 1)\n        segmented_x_with_time = torch.cat((segmented_x, time_channel), dim=-1)\n        \n        return segmented_x_with_time\n    \nclass SigFeats(nn.Module):\n    def __init__(self, input_channels, depth):\n        super(SigFeats, self).__init__()\n        self.input_channels = input_channels\n        self.depth = depth\n        self.signature_dim = signatory.signature_channels(input_channels, depth)\n    \n    def forward(self, x):\n        batch_size, num_segments, segment_length, num_channels = x.size()\n        x = x.view(-1, segment_length, num_channels)  # Reshape for signatory input: combining batch and segments\n        signatures = signatory.signature(x, depth=self.depth)\n        signatures = signatures.view(batch_size, num_segments, -1)  # Reshape back to separate segments out: (N, num_seg, sig_dim)\n        \n        return signatures","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Putting everything together","metadata":{}},{"cell_type":"code","source":"import torch.nn.functional as F","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class AttentionMechanism(nn.Module):\n    def __init__(self, hidden_dim):\n        super(AttentionMechanism, self).__init__()\n        self.W = nn.Linear(hidden_dim, hidden_dim)\n        self.uw = nn.Parameter(torch.rand(hidden_dim, 1))\n        self.tanh = nn.Tanh()\n    \n    def forward(self, lstm_out):\n        \"\"\"\n        lstm_out: [batch_size, seq_length, hidden_dim * 2]\n        \"\"\"\n        u = self.tanh(self.W(lstm_out))  # [batch_size, seq_length, hidden_dim * 2]\n        scores = torch.matmul(u, self.uw)  # [batch_size, seq_length, 1]\n        attention_weights = F.softmax(scores.squeeze(-1), dim=1)  # [batch_size, seq_length]\n        \n        # Expand weights for weighted sum\n        expanded_weights = attention_weights.unsqueeze(-1).expand_as(lstm_out)  # [batch_size, seq_length, hidden_dim * 2]\n        weighted_sum = torch.sum(expanded_weights * lstm_out, dim=1)  # [batch_size, hidden_dim * 2]\n        \n        return weighted_sum, attention_weights\n\nclass BiLSTMWithAttention(nn.Module):\n    def __init__(self, input_dim, hidden_dim, num_layers):\n        super(BiLSTMWithAttention, self).__init__()\n        self.bilstm = nn.LSTM(input_size=input_dim, hidden_size=hidden_dim,\n                              num_layers=num_layers, batch_first=True,\n                              bidirectional=True)\n        self.attention = AttentionMechanism(hidden_dim*2)\n    \n    def forward(self, x):\n        lstm_out, _ = self.bilstm(x)  # lstm_out: [batch_size, seq_length, hidden_dim * 2]\n        context_vector, attention_weights = self.attention(lstm_out)\n        \n        return context_vector, attention_weights","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ResGRUBlock(nn.Module):\n    def __init__(self, input_size, hidden_size):\n        super(ResGRUBlock, self).__init__()\n        self.gru = nn.GRU(input_size, hidden_size, batch_first=True)\n        self.bn = nn.BatchNorm1d(hidden_size)\n        self.relu = nn.ReLU()\n        # Adjustment function g(x) - a linear layer to match dimensions\n        self.adjustment = nn.Linear(input_size, hidden_size)\n\n    def forward(self, x):\n        # x is of shape [batch_size, seq_length, input_size]\n        batch_size, seq_length, _ = x.shape\n        original_x = x\n        # GRU layer\n        y, _ = self.gru(x)\n        # Reshape for batch normalization\n        y_reshaped = y.contiguous().view(batch_size * seq_length, -1)\n        y_bn = self.bn(self.relu(y_reshaped))\n        y_bn = y_bn.view(batch_size, seq_length, -1)\n        # Adjust original input dimensions if necessary\n        adjusted_x = self.adjustment(original_x)\n        # Adding the residual\n        yR = y_bn + adjusted_x\n        return yR","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CustomModel(nn.Module):\n    def __init__(self, sampling_rate, input_channels, depth, hidden_size, num_layers, num_classes=6):\n        super(CustomModel, self).__init__()\n        self.segment_eeg = SegmentEEG(sampling_rate=sampling_rate)\n        self.sig_feats = SigFeats(input_channels=input_channels + 1, depth=depth)\n        self.bilstm_attention = BiLSTMWithAttention(input_dim=self.sig_feats.signature_dim, \n                                                    hidden_dim=hidden_size, num_layers=num_layers)\n        self.res_gru_blocks = nn.ModuleList([ResGRUBlock(hidden_size if i > 0 else hidden_size * 2 + self.sig_feats.signature_dim, \n                                                         hidden_size) for i in range(num_layers)])\n        self.gru_attention = AttentionMechanism(hidden_size)\n        self.classifier = nn.Linear(hidden_size, num_classes)\n\n    def forward(self, x):\n        sig_feats = self.segment_eeg(x)\n        sig_feats = self.sig_feats(sig_feats)  #  [batch_size, seq_length, signature_dim]\n        \n        # Process through BiLSTM with Attention\n        context_vector, _ = self.bilstm_attention(sig_feats)  # [batch_size, hidden_dim * 2]\n        \n        # Expand context_vector to match sig_feats seq_length and concatenate\n        expanded_context = context_vector.unsqueeze(1).expand(-1, sig_feats.size(1), -1)\n        gru_input = torch.cat((sig_feats, expanded_context), dim=2)  # [batch_size, seq_length, signature_dim + hidden_dim * 2]\n        \n        # Process through ResGRUBlock(s)\n        for res_gru_block in self.res_gru_blocks:\n            gru_input = res_gru_block(gru_input)\n        \n        gru_context_vector, _ = self.gru_attention(gru_input)\n        logits = self.classifier(gru_context_vector)\n        return logits\n        ","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Using Pytorch Lightning","metadata":{}},{"cell_type":"code","source":"import pytorch_lightning as pl\nfrom pytorch_lightning import Trainer, seed_everything\nfrom pytorch_lightning.callbacks import ModelCheckpoint\nimport torch\nfrom torch.utils.data import DataLoader, Dataset\nfrom sklearn.model_selection import KFold\n\nseed_everything(42)\n\nclass LightningModel(pl.LightningModule):\n    def __init__(self, sampling_rate, input_channels, depth,hidden_size, num_layers, num_classes=6):\n        super().__init__()\n        self.model = CustomModel(sampling_rate, \n                                 input_channels, \n                                 depth, \n                                 hidden_size, \n                                 num_layers, \n                                 num_classes=6)\n        self.save_hyperparameters()\n        self.validation_outputs = []\n    \n    def forward(self, x):\n        return self.model(x)\n    \n    def training_step(self, batch, batch_idx):\n        X, y = batch\n        y_hat = self(X)\n        loss = F.kl_div(F.log_softmax(y_hat, dim=1), y, reduction='batchmean')\n        self.log('train_loss', loss, on_step=True, on_epoch=True, prog_bar=True, logger=True)\n        return loss\n    \n    def validation_step(self, batch, batch_idx):\n        X, y = batch\n        y_hat = self(X)\n        loss = F.kl_div(F.log_softmax(y_hat, dim=1), y, reduction='batchmean')\n        self.log('val_loss', loss, on_step=True, on_epoch=True, prog_bar=True, logger=True)\n        self.log('val_loss', loss, on_step=True, on_epoch=True, prog_bar=True, logger=True)\n        self.validation_outputs.append({'val_loss': loss, 'predictions': y_hat.detach(), 'targets': y.detach()})\n    \n    def on_validation_epoch_end(self):\n        avg_loss = torch.stack([x['val_loss'] for x in self.validation_outputs]).mean()\n        self.log('avg_val_loss', avg_loss, on_epoch=True, prog_bar=True, logger=True)\n        \n        # Save predictions and targets for access after validation\n        self.val_predictions = torch.cat([x['predictions'] for x in self.validation_outputs], dim=0)\n        self.val_targets = torch.cat([x['targets'] for x in self.validation_outputs], dim=0)\n        \n        # Clear the list for the next validation epoch\n        self.validation_outputs.clear()\n\n        self.validation_outputs.clear()\n    \n    def configure_optimizers(self):\n        optimizer = torch.optim.Adam(self.parameters(), lr=8e-4)\n        return optimizer","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Cross Validation","metadata":{}},{"cell_type":"code","source":"class EEGDataModule(pl.LightningDataModule):\n    def __init__(self, train_folds, valid_folds):\n        super().__init__()\n        self.train_folds = train_folds\n        self.valid_folds = valid_folds\n\n    def setup(self, stage=None):\n        train_ids = set(train_folds['eeg_id'])\n        val_ids = set(valid_folds['eeg_id'])\n        \n        # Filter the all_eegs dictionary to obtain separate dictionaries for training and validation\n        train_eegs = {eeg_id: eeg_data for eeg_id, eeg_data in all_eegs.items() if eeg_id in train_ids}\n        val_eegs = {eeg_id: eeg_data for eeg_id, eeg_data in all_eegs.items() if eeg_id in val_ids}\n        \n        self.train_dataset = CustomDataset(data=self.train_folds, eegs=train_eegs, feat2idx=FEAT2IDX, targets=TARGETS)\n        self.valid_dataset = CustomDataset(data=self.valid_folds, eegs=val_eegs, feat2idx=FEAT2IDX, targets=TARGETS)\n        \n\n    def train_dataloader(self):\n        return DataLoader(self.train_dataset, batch_size=32, shuffle=True, num_workers=3, pin_memory=True, drop_last=True)\n\n    def val_dataloader(self):\n        return DataLoader(self.valid_dataset, batch_size=32, shuffle=False, num_workers=3, pin_memory=True, drop_last=False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"n_splits = 5\nkf = KFold(n_splits=n_splits, shuffle=True, random_state=42)\n\nall_predictions = []\nall_targets = []\nmodels_dir = '/kaggle/working/models'\n\nfor fold, (train_idx, val_idx) in enumerate(kf.split(train)):\n    print(f\"FOLD {fold+1}\")\n    train_folds = train.iloc[train_idx]\n    valid_folds = train.iloc[val_idx]\n    \n    # Create a LightningDataModule for this fold\n    eeg_data_module = EEGDataModule(train_folds=train_folds, valid_folds=valid_folds)\n    \n    # Setup model checkpointing\n    checkpoint_callback = ModelCheckpoint(\n        monitor='avg_val_loss',\n        dirpath=models_dir,  # Use the defined models directory\n        filename=f'ps_blst_att_fold_{fold}_best',  # Specific naming convention\n        save_top_k=1,\n        mode='min',\n    )\n    \n    # Initialize the Lightning Trainer\n    trainer = Trainer(max_epochs=6, callbacks=[checkpoint_callback],\n                     strategy='auto', accelerator='auto', devices=\"auto\")\n    \n    # Instantiate and fit the model\n    model = LightningModel(sampling_rate=25, \n                           input_channels=8, \n                           depth=2, \n                           hidden_size=64, \n                           num_layers=2, \n                           num_classes=6)\n    trainer.fit(model, datamodule=eeg_data_module)\n    \n    # Collect and log predictions for each fold\n    validation_output = trainer.validate(model, datamodule=eeg_data_module)[0]\n    \n    # Access the saved predictions and targets\n    all_predictions.append(model.val_predictions)\n    all_targets.append(model.val_targets)\n\n# Concatenate all fold predictions and targets for final scoring\nall_predictions = torch.cat(all_predictions, dim=0)\nall_targets = torch.cat(all_targets, dim=0)\n\n# Compute final metric\npredictions_log_prob = F.log_softmax(all_predictions, dim=1)\ntargets_prob = all_targets\nkl_loss = F.kl_div(predictions_log_prob, targets_prob, reduction='batchmean')\nprint(f\"Final KLDiv Loss across all folds: {kl_loss.item()}\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"! zip -r models.zip /kaggle/working/models","metadata":{"execution":{"iopub.status.busy":"2024-04-08T13:28:48.052867Z","iopub.execute_input":"2024-04-08T13:28:48.053552Z","iopub.status.idle":"2024-04-08T13:28:54.738459Z","shell.execute_reply.started":"2024-04-08T13:28:48.053503Z","shell.execute_reply":"2024-04-08T13:28:54.737435Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference and Submission","metadata":{}},{"cell_type":"markdown","source":"## Test dataset","metadata":{}},{"cell_type":"code","source":"from glob import glob\nfrom tqdm import tqdm","metadata":{"execution":{"iopub.status.busy":"2024-04-08T13:28:54.739981Z","iopub.execute_input":"2024-04-08T13:28:54.740311Z","iopub.status.idle":"2024-04-08T13:28:54.745325Z","shell.execute_reply.started":"2024-04-08T13:28:54.740280Z","shell.execute_reply":"2024-04-08T13:28:54.744361Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class paths:\n    OUTPUT_DIR = \"/kaggle/working/\"\n    TEST_CSV = \"/kaggle/input/hms-harmful-brain-activity-classification/test.csv\"\n    TEST_EEGS = \"/kaggle/input/hms-harmful-brain-activity-classification/test_eegs/\"","metadata":{"execution":{"iopub.status.busy":"2024-04-08T13:28:54.746511Z","iopub.execute_input":"2024-04-08T13:28:54.746817Z","iopub.status.idle":"2024-04-08T13:28:54.757176Z","shell.execute_reply.started":"2024-04-08T13:28:54.746788Z","shell.execute_reply":"2024-04-08T13:28:54.756323Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test = pd.read_csv(paths.TEST_CSV)\nprint(f\"Test dataframe shape is: {test.shape}\")","metadata":{"execution":{"iopub.status.busy":"2024-04-08T13:28:54.758304Z","iopub.execute_input":"2024-04-08T13:28:54.758603Z","iopub.status.idle":"2024-04-08T13:28:54.784015Z","shell.execute_reply.started":"2024-04-08T13:28:54.758576Z","shell.execute_reply":"2024-04-08T13:28:54.783103Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_eegs = {}\neeg_paths = glob(paths.TEST_EEGS + \"*.parquet\")\neeg_ids = test.eeg_id.unique()\n\nfor i, eeg_id in tqdm(enumerate(eeg_ids)):  \n    # Save EEG to Python dictionary of numpy arrays\n    eeg_path = paths.TEST_EEGS + str(eeg_id) + \".parquet\"\n    data = eeg_from_parquet(eeg_path)              \n    test_eegs[eeg_id] = data\n\ntest_dataset = CustomDataset(data=test, eegs=test_eegs, feat2idx=FEAT2IDX, targets=TARGETS, mode='test')\n\ntest_X = test_dataset[0]\nprint(f\"X shape: {test_X.shape}\")","metadata":{"execution":{"iopub.status.busy":"2024-04-08T13:28:54.785235Z","iopub.execute_input":"2024-04-08T13:28:54.785585Z","iopub.status.idle":"2024-04-08T13:28:54.869686Z","shell.execute_reply.started":"2024-04-08T13:28:54.785552Z","shell.execute_reply":"2024-04-08T13:28:54.868739Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Inference","metadata":{}},{"cell_type":"code","source":"model_checkpoints_dir = '/kaggle/working/models'\nmodel_ckpts = glob(f\"{model_checkpoints_dir}/*.ckpt\")\nmodel_ckpts","metadata":{"execution":{"iopub.status.busy":"2024-04-08T13:28:54.870884Z","iopub.execute_input":"2024-04-08T13:28:54.871160Z","iopub.status.idle":"2024-04-08T13:28:54.879633Z","shell.execute_reply.started":"2024-04-08T13:28:54.871135Z","shell.execute_reply":"2024-04-08T13:28:54.878668Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_loader = DataLoader(test_dataset, batch_size=32, shuffle=False, num_workers=4)\n\npredictions = []\n\nfor model_ckpt in model_ckpts:\n    # Load the model\n    model = LightningModel.load_from_checkpoint(checkpoint_path=model_ckpt)\n    \n    # Initialize the trainer\n    trainer = Trainer()\n    \n    # Use the trainer to perform prediction\n    raw_predictions = trainer.predict(model, dataloaders=test_loader)\n    \n    # Process raw_predictions with softmax\n    softmax_predictions = [torch.softmax(batch, dim=1).numpy() for batch in raw_predictions]\n    predictions.append(np.concatenate(softmax_predictions, axis=0))\n    \n# Aggregate predictions from all models\npredictions = np.mean(np.array(predictions), axis=0)","metadata":{"execution":{"iopub.status.busy":"2024-04-08T13:28:54.880733Z","iopub.execute_input":"2024-04-08T13:28:54.881189Z","iopub.status.idle":"2024-04-08T13:29:16.519122Z","shell.execute_reply.started":"2024-04-08T13:28:54.881153Z","shell.execute_reply":"2024-04-08T13:29:16.518116Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Submission","metadata":{}},{"cell_type":"code","source":"TARGETS = ['seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']\nsub = pd.DataFrame({'eeg_id': test.eeg_id.values})\nsub[TARGETS] = predictions\nsub.to_csv('submission.csv',index=False)\nprint(f'Submission shape: {sub.shape}')\nsub.head()","metadata":{"execution":{"iopub.status.busy":"2024-04-08T13:29:16.520607Z","iopub.execute_input":"2024-04-08T13:29:16.520964Z","iopub.status.idle":"2024-04-08T13:29:16.557565Z","shell.execute_reply.started":"2024-04-08T13:29:16.520929Z","shell.execute_reply":"2024-04-08T13:29:16.556706Z"},"trusted":true},"execution_count":null,"outputs":[]}]}