{"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":"none","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"},{"sourceId":7823432,"sourceType":"datasetVersion","datasetId":4584064},{"sourceId":7831494,"sourceType":"datasetVersion","datasetId":4589786}],"dockerImageVersionId":30664,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np \nimport pandas as pd \nimport torch\nimport torch.nn as nn\nimport sys","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-03-14T18:14:50.600442Z","iopub.execute_input":"2024-03-14T18:14:50.600738Z","iopub.status.idle":"2024-03-14T18:14:56.214227Z","shell.execute_reply.started":"2024-03-14T18:14:50.600714Z","shell.execute_reply":"2024-03-14T18:14:56.213397Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv('/kaggle/input/training-data-meta-model/training_data_meta.csv')\ndisplay(df)","metadata":{"execution":{"iopub.status.busy":"2024-03-14T18:14:56.215781Z","iopub.execute_input":"2024-03-14T18:14:56.216179Z","iopub.status.idle":"2024-03-14T18:14:56.738833Z","shell.execute_reply.started":"2024-03-14T18:14:56.216154Z","shell.execute_reply":"2024-03-14T18:14:56.737922Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Get additional information to feed to the meta model\nlabels = ['seizure', 'lpd', 'gpd', 'lrda', 'grda', 'other']\nmax_df = df.groupby(['training_instance'])[labels].agg('max').reset_index()\nmin_df = df.groupby(['training_instance'])[labels].agg('min').reset_index()\nmean_df = df.groupby(['training_instance'])[labels].agg('mean').reset_index()\nmedian_df = df.groupby(['training_instance'])[labels].agg('median').reset_index()\nstd_df = df.groupby(['training_instance'])[labels].agg('std').reset_index()\ndisplay(mean_df)","metadata":{"execution":{"iopub.status.busy":"2024-03-14T18:24:57.890290Z","iopub.execute_input":"2024-03-14T18:24:57.891020Z","iopub.status.idle":"2024-03-14T18:24:57.993429Z","shell.execute_reply.started":"2024-03-14T18:24:57.890989Z","shell.execute_reply":"2024-03-14T18:24:57.992516Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def df2tensor(df):\n    max_list = df[labels].values.tolist()\n    max_tensor = torch.tensor(max_list)\n    max_tensor.shape\n    return max_tensor\n\nmax_tensor = df2tensor(max_df)\nmin_tensor = df2tensor(min_df)\nmean_tensor = df2tensor(mean_df)\nmedian_tensor = df2tensor(median_df)\nstd_tensor = df2tensor(std_df)\nmean_tensor.shape","metadata":{"execution":{"iopub.status.busy":"2024-03-14T18:24:58.746640Z","iopub.execute_input":"2024-03-14T18:24:58.747221Z","iopub.status.idle":"2024-03-14T18:24:58.921921Z","shell.execute_reply.started":"2024-03-14T18:24:58.747190Z","shell.execute_reply":"2024-03-14T18:24:58.920902Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def df2tensor_alldata(meta_df):\n    labels = ['seizure', 'lpd', 'gpd', 'lrda', 'grda', 'other']\n    data = []\n\n    for i in range(int(meta_df['training_instance'].max()) + 1):\n        instance_data = []\n        instance_df = meta_df[meta_df['training_instance'] == i]\n\n        for j in range(len(instance_df)):\n            model_data = [instance_df.iloc[j][label] for label in labels]\n            instance_data.append(model_data)\n\n        data.append(instance_data)\n    data_tensor = torch.tensor(data)\n    return data_tensor\n\ninput_data = df2tensor_alldata(df)\ninput_data.shape","metadata":{"execution":{"iopub.status.busy":"2024-03-14T18:26:06.047078Z","iopub.execute_input":"2024-03-14T18:26:06.047746Z","iopub.status.idle":"2024-03-14T18:26:51.081054Z","shell.execute_reply.started":"2024-03-14T18:26:06.047714Z","shell.execute_reply":"2024-03-14T18:26:51.080091Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"input_meta = torch.cat([max_tensor.unsqueeze(1), min_tensor.unsqueeze(1), \n                        mean_tensor.unsqueeze(1), std_tensor.unsqueeze(1), median_tensor.unsqueeze(1)], 1)\ninput_meta.shape","metadata":{"execution":{"iopub.status.busy":"2024-03-14T18:26:51.082974Z","iopub.execute_input":"2024-03-14T18:26:51.083743Z","iopub.status.idle":"2024-03-14T18:26:51.122187Z","shell.execute_reply.started":"2024-03-14T18:26:51.083707Z","shell.execute_reply":"2024-03-14T18:26:51.121319Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = pd.read_csv(\"/kaggle/input/hms-harmful-brain-activity-classification/train.csv\")\ntrain_features = pd.DataFrame()\n\nfor label in labels:\n    train_grouped_by_spectrogram_id = train_df[f'{label}_vote'].groupby(train_df['spectrogram_id']).sum()\n\n    label_vote_sum = pd.DataFrame()\n    label_vote_sum[\"spectrogram_id\"] = train_grouped_by_spectrogram_id.index\n    label_vote_sum[f\"{label}_vote_sum\"] = train_grouped_by_spectrogram_id.values\n\n    if label == labels[0]:\n        train_features = label_vote_sum\n    else:\n        train_features = train_features.merge(label_vote_sum, on='spectrogram_id', how='left')\n\n# Add a column to sum all votes\ntrain_features['total_vote'] = 0\nfor label in labels:\n    train_features['total_vote'] += train_features[f'{label}_vote_sum']\n\n# Calculate and store the normalized vote for each label\nfor label in labels:\n    train_features[f'{label}_vote'] = train_features[f'{label}_vote_sum'] / train_features['total_vote']\n\n# Select relevant columns for the training features\nchoose_cols = ['spectrogram_id']\nfor label in labels:\n    choose_cols += [f'{label}_vote']\ntrain_features = train_features[choose_cols]\n\n# Add a column with the path to the spectrogram files\ntrain_features['path'] = train_features['spectrogram_id'].apply(lambda x: \"/kaggle/input/hms-harmful-brain-activity-classification/train_spectrograms/\" + str(x) + \".parquet\")\n\nlabels_meta = train_features[['seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']].values.tolist()\nlabels_meta_tensor = torch.tensor(labels_meta)\nlabels_meta_tensor.shape","metadata":{"execution":{"iopub.status.busy":"2024-03-14T18:26:51.123552Z","iopub.execute_input":"2024-03-14T18:26:51.124057Z","iopub.status.idle":"2024-03-14T18:26:51.464120Z","shell.execute_reply.started":"2024-03-14T18:26:51.124025Z","shell.execute_reply":"2024-03-14T18:26:51.463106Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MLP(nn.Module):\n    def __init__(self, input_size):\n        super(MLP, self).__init__()\n        self.layers = nn.Sequential(\n            nn.Linear(input_size*6, 600),\n            nn.ReLU(),\n            nn.Linear(600, 1200),\n            nn.ReLU(),\n            nn.Linear(1200, 1000),\n            nn.ReLU(),\n            nn.Linear(1000, 600),\n            nn.ReLU(),\n            nn.Linear(600, 6),\n            nn.Softmax(1)\n        )\n        \n    def forward(self, x):\n        # convert tensor (64, 20, 6) --> (64, 20*6)\n        x = x.view(x.size(0), -1)\n        x = self.layers(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2024-03-14T18:26:51.466136Z","iopub.execute_input":"2024-03-14T18:26:51.466426Z","iopub.status.idle":"2024-03-14T18:26:51.473383Z","shell.execute_reply.started":"2024-03-14T18:26:51.466400Z","shell.execute_reply":"2024-03-14T18:26:51.472533Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class KLDivergenceLoss(nn.Module):\n    def __init__(self, epsilon=1e-15):\n        super(KLDivergenceLoss, self).__init__()\n        self.epsilon = epsilon\n\n    def forward(self, p, q):\n        # Clip probabilities to avoid log(0)\n        p = torch.clamp(p, self.epsilon, 1 - self.epsilon)\n\n        # Compute logarithms\n        log_p = torch.log(p)\n        log_q = nn.functional.log_softmax(q, dim=1)\n\n        # Calculate element-wise KL divergence\n        kl_divergence_per_point = p * (log_p - log_q)\n\n        # Sum over classes to get KL divergence per sample\n        kl_divergence_per_sample = torch.sum(kl_divergence_per_point, dim=1)\n\n        # Compute mean over samples\n        kl_loss = torch.mean(kl_divergence_per_sample)\n\n        return kl_loss\n","metadata":{"execution":{"iopub.status.busy":"2024-03-14T18:26:51.474596Z","iopub.execute_input":"2024-03-14T18:26:51.474889Z","iopub.status.idle":"2024-03-14T18:26:51.489347Z","shell.execute_reply.started":"2024-03-14T18:26:51.474855Z","shell.execute_reply":"2024-03-14T18:26:51.488342Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport sys\n\ndef train_mlp(input_tensor, target_tensor, model, criterion_ce, criterion_kl, optimizer, batch_size=32, num_epochs=10, validation_size=0.1,lambda_=1):\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    model.to(device)\n    \n    # Split data into training and validation sets\n    total_size = len(input_tensor)\n    val_size = int(total_size * validation_size)\n    train_size = total_size - val_size\n    \n    train_indices = torch.randperm(train_size)\n    val_indices = torch.arange(train_size, total_size)\n    \n    train_input = input_tensor[train_indices].to(device)\n    train_target = target_tensor[train_indices].to(device)\n    \n    val_input = input_tensor[val_indices].to(device)\n    val_target = target_tensor[val_indices].to(device)\n    \n    best_total_loss = sys.maxsize\n    not_improved_counter = 0\n    \n    for epoch in range(num_epochs):\n        model.train()\n        running_ce_loss = 0.0\n        running_kl_loss = 0.0\n        tel = 0\n        \n        # Shuffle indices for training\n        indices = torch.randperm(train_size)\n        \n        # Iterate over batches for training\n        for i in range(0, len(indices), batch_size):\n            batch_indices = indices[i:i+batch_size]\n            \n            inputs = train_input[batch_indices]\n            targets = train_target[batch_indices]\n\n            # Forward pass\n            outputs = model(inputs)\n            ce_loss = criterion_ce(outputs, targets)\n            kl_loss = criterion_kl(targets, outputs)\n\n            # Combined loss\n            total_loss = ce_loss + lambda_ * kl_loss\n\n            # Backward pass and optimization\n            optimizer.zero_grad()\n            total_loss.backward()\n            optimizer.step()\n\n            running_ce_loss += ce_loss.item()\n            running_kl_loss += kl_loss.item()\n            tel += 1\n\n        average_ce_loss = running_ce_loss / tel\n        average_kl_loss = running_kl_loss / tel\n        print(f\"Epoch [{epoch + 1}/{num_epochs}], CE Loss: {average_ce_loss:.4f}, KL Loss: {average_kl_loss:.4f}\")\n        \n        # Validation\n        model.eval()\n        with torch.no_grad():\n            val_outputs = model(val_input)\n            val_ce_loss = criterion_ce(val_outputs, val_target)\n            val_accuracy = accuracy(val_outputs, val_target)\n            val_kl_loss = criterion_kl(val_outputs, val_target)\n            val_total_loss = val_ce_loss + lambda_ * val_kl_loss\n            \n        if val_total_loss < best_total_loss:\n            torch.save(model.state_dict(), 'meta_model_best.pth')\n            best_total_loss = val_total_loss\n            print(f\"Model saved on epoch: {epoch} with total loss: {best_total_loss} and accuracy: {val_accuracy}\")\n            not_improved_counter = 0\n        else:\n            not_improved_counter += 1\n        if not_improved_counter >= 20:\n            break\n\ndef accuracy(outputs, targets):\n    _, predicted = torch.max(outputs, 1)\n    _, label = torch.max(targets, 1)\n    correct = (predicted == label).sum().item()\n    total = targets.size(0)\n    return correct / total\n","metadata":{"execution":{"iopub.status.busy":"2024-03-14T18:26:51.491419Z","iopub.execute_input":"2024-03-14T18:26:51.491736Z","iopub.status.idle":"2024-03-14T18:26:51.507326Z","shell.execute_reply.started":"2024-03-14T18:26:51.491712Z","shell.execute_reply":"2024-03-14T18:26:51.506467Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\ninput_tensor = input_meta.float().to(device)\nprint(input_tensor.shape)\nlabels_meta_tensor = labels_meta_tensor.float().to(device)\n\nmeta_model = MLP(input_tensor.shape[1]).to(device)\n\noptimizer = torch.optim.Adam(meta_model.parameters(), lr=0.001)\nloss_ce = nn.CrossEntropyLoss()\nloss_bce = nn.BCELoss()\nloss_kl = KLDivergenceLoss()","metadata":{"execution":{"iopub.status.busy":"2024-03-14T18:26:51.508603Z","iopub.execute_input":"2024-03-14T18:26:51.508871Z","iopub.status.idle":"2024-03-14T18:26:54.957599Z","shell.execute_reply.started":"2024-03-14T18:26:51.508849Z","shell.execute_reply":"2024-03-14T18:26:54.956783Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"meta_model","metadata":{"execution":{"iopub.status.busy":"2024-03-14T13:07:23.718220Z","iopub.execute_input":"2024-03-14T13:07:23.718590Z","iopub.status.idle":"2024-03-14T13:07:23.726093Z","shell.execute_reply.started":"2024-03-14T13:07:23.718565Z","shell.execute_reply":"2024-03-14T13:07:23.725279Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mean_tensor.unsqueeze(1).shape","metadata":{"execution":{"iopub.status.busy":"2024-03-14T18:25:13.276063Z","iopub.execute_input":"2024-03-14T18:25:13.276665Z","iopub.status.idle":"2024-03-14T18:25:13.304118Z","shell.execute_reply.started":"2024-03-14T18:25:13.276632Z","shell.execute_reply":"2024-03-14T18:25:13.303264Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labels_meta_tensor.shape","metadata":{"execution":{"iopub.status.busy":"2024-03-14T18:26:54.958659Z","iopub.execute_input":"2024-03-14T18:26:54.959047Z","iopub.status.idle":"2024-03-14T18:26:54.964468Z","shell.execute_reply.started":"2024-03-14T18:26:54.959022Z","shell.execute_reply":"2024-03-14T18:26:54.963577Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(loss_kl(mean_tensor.to(device),labels_meta_tensor))","metadata":{"execution":{"iopub.status.busy":"2024-03-14T18:32:18.652340Z","iopub.execute_input":"2024-03-14T18:32:18.652715Z","iopub.status.idle":"2024-03-14T18:32:18.659700Z","shell.execute_reply.started":"2024-03-14T18:32:18.652687Z","shell.execute_reply":"2024-03-14T18:32:18.658947Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_mlp(input_tensor, labels_meta_tensor, meta_model, loss_ce, loss_kl, optimizer, batch_size=64, num_epochs=80, validation_size=0.1, lambda_=2)\ntorch.save(meta_model.state_dict(), 'meta_model_final.pth')","metadata":{"execution":{"iopub.status.busy":"2024-03-14T18:28:50.886230Z","iopub.execute_input":"2024-03-14T18:28:50.886646Z","iopub.status.idle":"2024-03-14T18:29:28.126409Z","shell.execute_reply.started":"2024-03-14T18:28:50.886617Z","shell.execute_reply":"2024-03-14T18:29:28.125395Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with torch.no_grad():\n    output = meta_model(input_tensor)","metadata":{"execution":{"iopub.status.busy":"2024-03-14T18:30:01.582891Z","iopub.execute_input":"2024-03-14T18:30:01.583240Z","iopub.status.idle":"2024-03-14T18:30:01.590224Z","shell.execute_reply.started":"2024-03-14T18:30:01.583215Z","shell.execute_reply":"2024-03-14T18:30:01.589313Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(loss_kl(output,labels_meta_tensor))","metadata":{"execution":{"iopub.status.busy":"2024-03-14T18:30:02.616736Z","iopub.execute_input":"2024-03-14T18:30:02.617513Z","iopub.status.idle":"2024-03-14T18:30:02.624073Z","shell.execute_reply.started":"2024-03-14T18:30:02.617478Z","shell.execute_reply":"2024-03-14T18:30:02.623120Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"1.1502 input, mean, std\n1.1211 max, min, mean, std, median\n1.1325 input, max, min, mean, std, median\nCurrent version: max, min, mean, std, median","metadata":{}}]}