{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"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":7857591,"sourceType":"datasetVersion","datasetId":4606875}],"dockerImageVersionId":30673,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nos.environ[\"CUDA_VISIBLE_DEVICES\"]=\"6\"\n\nimport IPython.display as ipd\nimport pandas as pd\nimport numpy as np\nimport seaborn as sns\nimport librosa\nimport soundfile as sf\nimport re\nimport scipy\nimport torch\nimport matplotlib.pyplot as plt\nimport nnAudio\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchaudio\nimport h5py\nimport gc\nimport timeit\nimport albumentations as A\nimport lightning\nimport timm\n\nfrom typing import Tuple, Dict, Any, Optional, Callable, Union\nfrom glob import glob\nfrom tqdm import tqdm\nfrom joblib import Parallel, delayed\nfrom tqdm import tqdm\nfrom collections import Counter\nfrom nnAudio.features.mel import MelSpectrogram\nfrom torchaudio.transforms import AmplitudeToDB\nfrom albumentations.pytorch import ToTensorV2\nfrom lightning.pytorch.callbacks import Callback\nfrom torchaudio.transforms import FrequencyMasking, TimeMasking\nfrom sklearn.model_selection import StratifiedGroupKFold\nfrom lightning.pytorch import loggers as pl_loggers\nfrom lightning.pytorch.callbacks import LearningRateMonitor, ModelCheckpoint\nfrom time import time\nfrom pprint import pprint\nfrom collections import OrderedDict\n\n%matplotlib inline","metadata":{"execution":{"iopub.status.busy":"2024-03-21T21:13:34.241716Z","iopub.execute_input":"2024-03-21T21:13:34.242016Z","iopub.status.idle":"2024-03-21T21:13:43.070381Z","shell.execute_reply.started":"2024-03-21T21:13:34.241987Z","shell.execute_reply":"2024-03-21T21:13:43.069141Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install nnAudio lightning","metadata":{"execution":{"iopub.status.busy":"2024-03-21T21:13:19.746396Z","iopub.execute_input":"2024-03-21T21:13:19.746764Z","iopub.status.idle":"2024-03-21T21:13:34.239432Z","shell.execute_reply.started":"2024-03-21T21:13:19.746738Z","shell.execute_reply":"2024-03-21T21:13:34.238616Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Task\n\nFor this lecture we will consider [HMS - Harmful Brain Activity Classification](https://www.kaggle.com/competitions/hms-harmful-brain-activity-classification/overview) Kaggle competition. \n\nIt is not audio but rather signal classification task:\n\n```\nThe goal of this competition is to detect and classify seizures and other types of harmful brain activity. You will develop a model trained on electroencephalography (EEG) signals recorded from critically ill hospital patients.\n```\n\nBut you can apply pretty much the same methods as for audio classification.","metadata":{}},{"cell_type":"code","source":"# # Uncomment it to download\n# !kaggle competitions download -c hms-harmful-brain-activity-classification\n# !unzip hms-harmful-brain-activity-classification.zip -d ../../data/hms","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DATA_ROOT = \"/kaggle/input/hms-harmful-brain-activity-classification/\"","metadata":{"execution":{"iopub.status.busy":"2024-03-21T21:14:37.269189Z","iopub.execute_input":"2024-03-21T21:14:37.269646Z","iopub.status.idle":"2024-03-21T21:14:37.275789Z","shell.execute_reply.started":"2024-03-21T21:14:37.269611Z","shell.execute_reply":"2024-03-21T21:14:37.274276Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Agenda\n1. [EDA](#EDA)\n2. [Metric](#Metric)\n3. [Data Preparation](#Data_Preparation)\n4. [Validation](#Validation)\n5. [Code Preparation](#Code_Preparation)\n6. [Model](#Model)\n7. [Training](#Training)\n8. [Evaluation](#Evaluation)\n9. [Homework](#Homework)","metadata":{}},{"cell_type":"markdown","source":"<a id='EDA'></a>\n# Exploratory Data Analysis (EDA)\n\n## Data Description\n\nThe goal of this competition is to detect and classify seizures and other types of harmful brain activity in electroencephalography (EEG) data. Even experts find this to be a challenging task and often disagree about the correct labels.\n\nThis is a code competition. Only a few examples from the test set are available for download. When your submission is scored the test folders will be replaced with versions containing the complete test set.\n\n### Files\n\n**train.csv** Metadata for the train set. The expert annotators reviewed 50 second long EEG samples plus matched spectrograms covering 10 a minute window centered at the same time and labeled the central 10 seconds. Many of these samples overlapped and have been consolidated. `train.csv` provides the metadata that allows you to extract the original subsets that the raters annotated.\n\n- `eeg_id` - A unique identifier for the entire EEG recording.\n- `eeg_sub_id` - An ID for the specific 50 second long subsample this row's labels apply to.\n- `eeg_label_offset_seconds` - The time between the beginning of the consolidated EEG and this subsample.\n- `spectrogram_id` - A unique identifier for the entire EEG recording.\n- `spectrogram_sub_id` - An ID for the specific 10 minute subsample this row's labels apply to.\n- `spectogram_label_offset_seconds` - The time between the beginning of the consolidated spectrogram and this subsample.\n- `label_id` - An ID for this set of labels.\n- `patient_id` - An ID for the patient who donated the data.\n- `expert_consensus` - The consensus annotator label. Provided for convenience only.\n- `[seizure/lpd/gpd/lrda/grda/other]_vote` - The count of annotator votes for a given brain activity class. The full names of the activity classes are as follows: lpd: lateralized periodic discharges, gpd: generalized periodic discharges, lrd: lateralized rhythmic delta activity, and grda: generalized rhythmic delta activity . A detailed explanations of these patterns is [available here](https://www.acns.org/UserFiles/file/ACNSStandardizedCriticalCareEEGTerminology_rev2021.pdf).\n\n**test.csv** Metadata for the test set. As there are no overlapping samples in the test set, many columns in the train metadata don't apply.\n\n- `eeg_id`\n- `spectrogram_id`\n- `patient_id`\n\n**sample_submission.csv**\n\n- eeg_id\n- `[seizure/lpd/gpd/lrda/grda/other]_vote` - The target columns. Your predictions must be probabilities. Note that the test samples had between 3 and 20 annotators.\n\n**train_eegs/** EEG data from one or more overlapping samples. Use the metadata in **train.csv** to select specific annotated subsets. The column names are [the names of the individual electrode locations for EEG leads](https://en.wikipedia.org/wiki/10%E2%80%9320_system_%28EEG%29), with one exception. The EKG column is for an electrocardiogram lead that records data from the heart. All of the EEG data (for both train and test) was collected at a frequency of 200 samples per second.\n\n**test_eegs/** Exactly 50 seconds of EEG data.\n\n**train_spectrograms/** Spectrograms assembled EEG data. Use the metadata in train.csv to select specific annotated subsets. The column names indicate the frequency in hertz and the recording regions of the EEG electrodes. The latter are abbreviated as LL = left lateral; RL = right lateral; LP = left parasagittal; RP = right parasagittal.\n\n**test_spectrograms/** Spectrograms assembled using exactly 10 minutes of EEG data.\n\n**example_figures/** Larger copies of the example case images used on the overview tab.","metadata":{}},{"cell_type":"markdown","source":"## Meta Visualization","metadata":{}},{"cell_type":"code","source":"train = pd.read_csv(os.path.join(DATA_ROOT,\"train.csv\"))\nsample_submission = pd.read_csv(os.path.join(DATA_ROOT,\"sample_submission.csv\")) \ntest = pd.read_csv(os.path.join(DATA_ROOT,\"test.csv\"))","metadata":{"execution":{"iopub.status.busy":"2024-03-21T21:14:41.464662Z","iopub.execute_input":"2024-03-21T21:14:41.465481Z","iopub.status.idle":"2024-03-21T21:14:41.662691Z","shell.execute_reply.started":"2024-03-21T21:14:41.465402Z","shell.execute_reply":"2024-03-21T21:14:41.661898Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train.shape","metadata":{"execution":{"iopub.status.busy":"2024-03-21T19:23:46.275746Z","iopub.execute_input":"2024-03-21T19:23:46.276107Z","iopub.status.idle":"2024-03-21T19:23:46.284383Z","shell.execute_reply.started":"2024-03-21T19:23:46.276081Z","shell.execute_reply":"2024-03-21T19:23:46.283163Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train.head()","metadata":{"execution":{"iopub.status.busy":"2024-03-21T19:23:50.382015Z","iopub.execute_input":"2024-03-21T19:23:50.382406Z","iopub.status.idle":"2024-03-21T19:23:50.407522Z","shell.execute_reply.started":"2024-03-21T19:23:50.382375Z","shell.execute_reply":"2024-03-21T19:23:50.406467Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_submission","metadata":{"execution":{"iopub.status.busy":"2024-03-21T19:23:52.67826Z","iopub.execute_input":"2024-03-21T19:23:52.678646Z","iopub.status.idle":"2024-03-21T19:23:52.693129Z","shell.execute_reply.started":"2024-03-21T19:23:52.678615Z","shell.execute_reply":"2024-03-21T19:23:52.691856Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test","metadata":{"execution":{"iopub.status.busy":"2024-03-21T19:23:54.993418Z","iopub.execute_input":"2024-03-21T19:23:54.993823Z","iopub.status.idle":"2024-03-21T19:23:55.003847Z","shell.execute_reply.started":"2024-03-21T19:23:54.993791Z","shell.execute_reply":"2024-03-21T19:23:55.002734Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"Unique EEG IDs:\", len(set(train[\"eeg_id\"])))\nplt.title(\"EEG IDs Value Counts Histogram\")\ntrain[\"eeg_id\"].value_counts().hist(bins=100)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-03-21T19:23:56.918216Z","iopub.execute_input":"2024-03-21T19:23:56.918606Z","iopub.status.idle":"2024-03-21T19:23:57.384026Z","shell.execute_reply.started":"2024-03-21T19:23:56.918575Z","shell.execute_reply":"2024-03-21T19:23:57.382844Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"Unique Spectogram IDs:\", len(set(train[\"spectrogram_id\"])))\nplt.title(\"Spectogram IDs Value Counts Histogram\")\ntrain[\"spectrogram_id\"].value_counts().hist(bins=100)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-03-21T19:23:59.95345Z","iopub.execute_input":"2024-03-21T19:23:59.953871Z","iopub.status.idle":"2024-03-21T19:24:00.334873Z","shell.execute_reply.started":"2024-03-21T19:23:59.953838Z","shell.execute_reply":"2024-03-21T19:24:00.333641Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def sub_id_sanity_check(sub_ids):\n    sub_ids = list(sub_ids)\n    return sub_ids == list(range(max(sub_ids) + 1))\n\ntrain.groupby(\"eeg_id\")[\"eeg_sub_id\"].apply(sub_id_sanity_check).all(), train.groupby(\"spectrogram_id\")[\"spectrogram_sub_id\"].apply(sub_id_sanity_check).all()","metadata":{"execution":{"iopub.status.busy":"2024-03-21T19:24:04.888616Z","iopub.execute_input":"2024-03-21T19:24:04.889017Z","iopub.status.idle":"2024-03-21T19:24:05.592819Z","shell.execute_reply.started":"2024-03-21T19:24:04.888989Z","shell.execute_reply":"2024-03-21T19:24:05.591988Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(set(train[\"label_id\"])) == len(train)","metadata":{"execution":{"iopub.status.busy":"2024-03-21T19:24:09.967727Z","iopub.execute_input":"2024-03-21T19:24:09.968925Z","iopub.status.idle":"2024-03-21T19:24:10.002359Z","shell.execute_reply.started":"2024-03-21T19:24:09.968886Z","shell.execute_reply":"2024-03-21T19:24:10.001156Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"Unique Patient IDs:\", len(set(train[\"patient_id\"])))\nplt.title(\"Global Patient IDs Value Counts Histogram\")\ntrain[\"patient_id\"].value_counts().hist(bins=100)\nplt.show()\nplt.title(\"Global Patient IDs Value Counts by EEG ID Histogram\")\ntrain.drop_duplicates(\"eeg_id\")[\"patient_id\"].value_counts().hist(bins=100)\nplt.show()\nplt.title(\"Global Patient IDs Value Counts by Spectogram ID Histogram\")\ntrain.drop_duplicates(\"eeg_id\")[\"spectrogram_id\"].value_counts().hist(bins=100)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-03-21T19:24:11.451914Z","iopub.execute_input":"2024-03-21T19:24:11.452286Z","iopub.status.idle":"2024-03-21T19:24:12.566594Z","shell.execute_reply.started":"2024-03-21T19:24:11.452258Z","shell.execute_reply":"2024-03-21T19:24:12.565483Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train[\"expert_consensus\"].value_counts()","metadata":{"execution":{"iopub.status.busy":"2024-03-21T19:24:16.086291Z","iopub.execute_input":"2024-03-21T19:24:16.086949Z","iopub.status.idle":"2024-03-21T19:24:16.100564Z","shell.execute_reply.started":"2024-03-21T19:24:16.086917Z","shell.execute_reply":"2024-03-21T19:24:16.099818Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train[\"expert_consensus\"].value_counts(normalize=True)","metadata":{"execution":{"iopub.status.busy":"2024-03-21T19:24:18.049215Z","iopub.execute_input":"2024-03-21T19:24:18.049881Z","iopub.status.idle":"2024-03-21T19:24:18.063799Z","shell.execute_reply.started":"2024-03-21T19:24:18.049848Z","shell.execute_reply":"2024-03-21T19:24:18.063003Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"Value Counts By the number of classes for one EEG waveform\")\ntrain.groupby(\"eeg_id\")[\"expert_consensus\"].apply(set).apply(len).value_counts()","metadata":{"execution":{"iopub.status.busy":"2024-03-21T19:24:20.501388Z","iopub.execute_input":"2024-03-21T19:24:20.502173Z","iopub.status.idle":"2024-03-21T19:24:20.928486Z","shell.execute_reply.started":"2024-03-21T19:24:20.502136Z","shell.execute_reply":"2024-03-21T19:24:20.927446Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"Value Counts By the number of classes for one EEG waveform\")\ntrain.groupby(\"spectrogram_id\")[\"expert_consensus\"].apply(set).apply(len).value_counts()","metadata":{"execution":{"iopub.status.busy":"2024-03-21T19:24:23.37237Z","iopub.execute_input":"2024-03-21T19:24:23.373189Z","iopub.status.idle":"2024-03-21T19:24:23.645233Z","shell.execute_reply.started":"2024-03-21T19:24:23.373143Z","shell.execute_reply":"2024-03-21T19:24:23.644386Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train.drop_duplicates(\"eeg_id\")[\"expert_consensus\"].value_counts(normalize=True)","metadata":{"execution":{"iopub.status.busy":"2024-03-21T19:24:25.390687Z","iopub.execute_input":"2024-03-21T19:24:25.391583Z","iopub.status.idle":"2024-03-21T19:24:25.405487Z","shell.execute_reply.started":"2024-03-21T19:24:25.391547Z","shell.execute_reply":"2024-03-21T19:24:25.40432Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train.drop_duplicates(\"spectrogram_id\")[\"expert_consensus\"].value_counts(normalize=True)","metadata":{"execution":{"iopub.status.busy":"2024-03-21T19:24:27.789995Z","iopub.execute_input":"2024-03-21T19:24:27.790437Z","iopub.status.idle":"2024-03-21T19:24:27.804022Z","shell.execute_reply.started":"2024-03-21T19:24:27.790406Z","shell.execute_reply":"2024-03-21T19:24:27.802786Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def pick_seq_class_v1(all_classes):\n    counter = Counter(all_classes)\n    most_common_element, _ = counter.most_common(1)[0]  # Get the most common element\n    return most_common_element\n\ndef pick_seq_class_v2(all_classes):\n    all_classes_no_other = [el for el in all_classes if el != \"Other\"]\n    if len(all_classes_no_other) > 0:\n        counter = Counter(all_classes_no_other)\n    else:\n        counter = Counter(all_classes)\n    most_common_element, _ = counter.most_common(1)[0]  # Get the most common element\n    return most_common_element","metadata":{"execution":{"iopub.status.busy":"2024-03-21T19:24:31.445735Z","iopub.execute_input":"2024-03-21T19:24:31.446108Z","iopub.status.idle":"2024-03-21T19:24:31.45346Z","shell.execute_reply.started":"2024-03-21T19:24:31.446081Z","shell.execute_reply":"2024-03-21T19:24:31.452161Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train.groupby(\"eeg_id\")[\"expert_consensus\"].apply(pick_seq_class_v1).value_counts(normalize=True)","metadata":{"execution":{"iopub.status.busy":"2024-03-21T19:24:34.723098Z","iopub.execute_input":"2024-03-21T19:24:34.723507Z","iopub.status.idle":"2024-03-21T19:24:35.225302Z","shell.execute_reply.started":"2024-03-21T19:24:34.723476Z","shell.execute_reply":"2024-03-21T19:24:35.224274Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train.groupby(\"eeg_id\")[\"expert_consensus\"].apply(pick_seq_class_v2).value_counts(normalize=True)","metadata":{"execution":{"iopub.status.busy":"2024-03-21T19:24:42.683409Z","iopub.execute_input":"2024-03-21T19:24:42.684084Z","iopub.status.idle":"2024-03-21T19:24:43.228966Z","shell.execute_reply.started":"2024-03-21T19:24:42.684049Z","shell.execute_reply":"2024-03-21T19:24:43.228205Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"> **TODO**: Continue...","metadata":{}},{"cell_type":"code","source":"train.keys()","metadata":{"execution":{"iopub.status.busy":"2024-03-21T21:16:21.243594Z","iopub.execute_input":"2024-03-21T21:16:21.244754Z","iopub.status.idle":"2024-03-21T21:16:21.251821Z","shell.execute_reply.started":"2024-03-21T21:16:21.244716Z","shell.execute_reply":"2024-03-21T21:16:21.250635Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Class Distribution\nclass_counts = train[['seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']].sum()\n\n# Plot class distribution\nplt.figure(figsize=(10, 6))\nclass_counts.plot(kind='bar', color='skyblue')\nplt.title('Class Distribution')\nplt.xlabel('Class')\nplt.ylabel('Count')\nplt.xticks(rotation=45)\nplt.grid(axis='y', linestyle='--', alpha=0.7)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-03-21T21:31:47.994817Z","iopub.execute_input":"2024-03-21T21:31:47.995154Z","iopub.status.idle":"2024-03-21T21:31:48.216057Z","shell.execute_reply.started":"2024-03-21T21:31:47.995127Z","shell.execute_reply":"2024-03-21T21:31:48.214653Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The class distribution analysis reveals a notable class imbalance, with the \"other_vote\" class exhibiting nearly twice the number of samples compared to the other classes. Addressing this imbalance is crucial for ensuring the robustness and fairness of machine learning models trained on this dataset.","metadata":{}},{"cell_type":"code","source":"# Compute correlation matrix\nclass_correlation = train[['seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']].corr()\n\n# Plot correlation heatmap\nplt.figure(figsize=(8, 6))\nsns.heatmap(class_correlation, annot=True, cmap='coolwarm', fmt=\".2f\", linewidths=0.5)\nplt.title('Correlation Between Brain Activity Classes')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-03-21T21:31:38.260422Z","iopub.execute_input":"2024-03-21T21:31:38.260774Z","iopub.status.idle":"2024-03-21T21:31:38.555308Z","shell.execute_reply.started":"2024-03-21T21:31:38.260748Z","shell.execute_reply":"2024-03-21T21:31:38.553634Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The correlation matrix reveals that seizure activity (`seizure_vote`) exhibits a positive correlation with generalized periodic discharges (`gpd_vote`) but negative correlations with other classes. Conversely, lateralized periodic discharges (`lpd_vote`) demonstrate a negative correlation with seizure activity and generalized periodic discharges. Other classes show varied correlations with each other, indicating distinct patterns of brain activity.\n","metadata":{}},{"cell_type":"code","source":"# Temporal Analysis\ntrain_sorted = train.sort_values(by='eeg_id')\nplt.figure(figsize=(12, 6))\nfor class_name in class_counts.index:\n    class_vote_col = f'{class_name}'\n    plt.plot(train_sorted['eeg_id'], train_sorted[class_vote_col].cumsum(), label=class_name)\n\nplt.title('Cumulative Occurrence of Classes Over Time')\nplt.xlabel('EEG ID')\nplt.ylabel('Cumulative Count')\nplt.legend()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-03-21T21:36:22.631773Z","iopub.execute_input":"2024-03-21T21:36:22.632386Z","iopub.status.idle":"2024-03-21T21:36:23.535852Z","shell.execute_reply.started":"2024-03-21T21:36:22.632348Z","shell.execute_reply":"2024-03-21T21:36:23.535148Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data Samples Visualization","metadata":{}},{"cell_type":"code","source":"NEURAL_SENSORS = (\n    \"Fp1\",\n    \"F3\",\n    \"C3\",\n    \"P3\",\n    \"F7\",\n    \"T3\",\n    \"T5\",\n    \"O1\",\n    \"Fz\",\n    \"Cz\",\n    \"Pz\",\n    \"Fp2\",\n    \"F4\",\n    \"C4\",\n    \"P4\",\n    \"F8\",\n    \"T4\",\n    \"T6\",\n    \"O2\",\n    \"EKG\",\n)\nSPEC_TYPES = (\"LL\", \"RL\", \"LP\", \"RP\")\nREVERSED_SPEC_TYPES = ('LL','LP','RP','RR')\nREVERSED_FEATS = [['Fp1','F7','T3','T5','O1'],\n         ['Fp1','F3','C3','P3','O1'],\n         ['Fp2','F8','T4','T6','O2'],\n         ['Fp2','F4','C4','P4','O2']]\nTARGETS = (\"Seizure\", \"LPD\", \"GPD\", \"LRDA\", \"GRDA\", \"Other\")\nTARGET2ID = {\"Seizure\": 0, \"LPD\": 1, \"GPD\": 2, \"LRDA\": 3, \"GRDA\": 4, \"Other\": 5}\nID2TARGET = {v: k for k, v in TARGET2ID.items()}\nDEFAULT_SAMPLE_RATE = 200","metadata":{"execution":{"iopub.status.busy":"2024-03-21T22:00:48.370526Z","iopub.execute_input":"2024-03-21T22:00:48.371303Z","iopub.status.idle":"2024-03-21T22:00:48.37948Z","shell.execute_reply.started":"2024-03-21T22:00:48.371267Z","shell.execute_reply":"2024-03-21T22:00:48.377966Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Converting to [h5py](https://docs.h5py.org/en/stable/)\n\nWhy do we need to this?\n\nBecause we want to access to chunks of files, without reading the whole file.\n\n```\nHDF5 lets you store huge amounts of numerical data, and easily manipulate that data from NumPy. For example, you can slice into multi-terabyte datasets stored on disk, as if they were real NumPy arrays. Thousands of datasets can be stored in a single file, categorized and tagged however you want.\n```","metadata":{}},{"cell_type":"code","source":"class ProgressParallel(Parallel):\n    def __init__(self, use_tqdm=True, total=None, *args, **kwargs):\n        self._use_tqdm = use_tqdm\n        self._total = total\n        super().__init__(*args, **kwargs)\n\n    def __call__(self, *args, **kwargs):\n        with tqdm(disable=not self._use_tqdm, total=self._total) as self._pbar:\n            return Parallel.__call__(self, *args, **kwargs)\n\n    def print_progress(self):\n        if self._total is None:\n            self._pbar.total = self.n_dispatched_tasks\n        self._pbar.n = self.n_completed_tasks\n        self._pbar.refresh()\n\ndef compose_spec_from_df(input_df, spec_type, to_db=True, return_freq_and_time=False):\n    spec_cols = [col for col in input_df.columns if spec_type in col]\n    spec_cols = sorted(spec_cols, key=lambda x: float(x.split(\"_\")[1]))\n    spec = np.stack([input_df[el].values for el in spec_cols])\n    if to_db:\n        spec = librosa.amplitude_to_db(spec)\n    if return_freq_and_time:\n        freqs = [float(col.split(\"_\")[1]) for col in spec_cols]\n        times = list(input_df[\"time\"])\n        return spec, freqs, times\n\ndef process_and_save_spec(spec_src_path, spec_tgt_path):\n    sample_df = pd.read_parquet(spec_src_path, engine=\"pyarrow\")\n    prev_times, prev_freqs = None, None\n    with h5py.File(spec_tgt_path, \"w\") as data_file:\n        for spec_type in SPEC_TYPES:\n            spec, freqs, times = compose_spec_from_df(sample_df, spec_type, to_db=False, return_freq_and_time=True)\n            if prev_times is not None:\n                assert np.all(prev_times == times)\n                assert np.all(prev_freqs == freqs)\n            data_file.create_dataset(spec_type, data=spec)\n            prev_times, prev_freqs = times, freqs\n        data_file.create_dataset(\"freqs\", data=np.array(freqs))\n        data_file.create_dataset(\"times\", data=np.array(times))\n\ndef process_and_save_eeg(eeg_src_path, eeg_tgt_path):\n    sample_df = pd.read_parquet(eeg_src_path, engine=\"pyarrow\")\n    assert set(sample_df.columns) == set(NEURAL_SENSORS)\n    with h5py.File(eeg_tgt_path, \"w\") as data_file:\n        for sensor in NEURAL_SENSORS:\n            data_file.create_dataset(sensor, data=sample_df[sensor].values)\n\ndef read_h5py_file(file_path: str):\n    with h5py.File(file_path, \"r\") as data_file:\n        data = {key: data_file[key][:] for key in data_file.keys()}\n    return data","metadata":{"execution":{"iopub.status.busy":"2024-03-21T22:00:55.616983Z","iopub.execute_input":"2024-03-21T22:00:55.61734Z","iopub.status.idle":"2024-03-21T22:00:55.631777Z","shell.execute_reply.started":"2024-03-21T22:00:55.617313Z","shell.execute_reply":"2024-03-21T22:00:55.630718Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<font style=\"color:red\">Uncomment next 2 blocks:</font>","metadata":{}},{"cell_type":"code","source":"# eeg_spec_file_pathes = glob(\n#     os.path.join(DATA_ROOT,\"train_spectrograms/*.parquet\")\n# )\n# print(f\"Found {len(eeg_spec_file_pathes)} EEG Spectogram files\")\n# os.makedirs(\n#     os.path.join(DATA_ROOT,\"train_spectrograms_npy\"), exist_ok=True\n# )\n# ProgressParallel(n_jobs=4, total=len(eeg_spec_file_pathes))(\n#     delayed(process_and_save_spec)(\n#         spec_src_path=spec_src_path, \n#         spec_tgt_path=os.path.join(DATA_ROOT, \"train_spectrograms_npy\", os.path.basename(spec_src_path).replace(\".parquet\", \".h5\"))\n#     ) for spec_src_path in eeg_spec_file_pathes\n# );","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# eeg_file_pathes = glob(\n#     os.path.join(DATA_ROOT,\"train_eegs/*.parquet\")\n# )\n# print(f\"Found {len(eeg_file_pathes)} EEG files\")\n\n# os.makedirs(\n#     os.path.join(DATA_ROOT,\"train_eegs_npy\"), exist_ok=True\n# )\n# ProgressParallel(n_jobs=4, total=len(eeg_file_pathes))(\n#     delayed(process_and_save_eeg)(\n#         eeg_src_path=eeg_src_path,\n#         eeg_tgt_path=os.path.join(DATA_ROOT, \"train_eegs_npy\", os.path.basename(eeg_src_path).replace(\".parquet\", \".h5\"))\n#     )\n#     for eeg_src_path in eeg_file_pathes\n# )","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2024-03-21T19:35:17.313792Z","iopub.execute_input":"2024-03-21T19:35:17.314434Z","iopub.status.idle":"2024-03-21T19:35:18.380705Z","shell.execute_reply.started":"2024-03-21T19:35:17.31439Z","shell.execute_reply":"2024-03-21T19:35:18.378858Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Let's banchmark our `h5` files:","metadata":{}},{"cell_type":"code","source":"# First: quick \"lazy\" shape access\n\n#/kaggle/input/ucu-hms-h5py/train_spectrograms_npy/train_spectrograms_npy\n\ndef lazy_read_shape(path):\n    with h5py.File(path, \"r\") as data_file:\n        shape = data_file[SPEC_TYPES[0]].shape\n    return shape\n\n# all_spec_paths = glob(os.path.join(DATA_ROOT,\"/train_spectrograms_npy/train_spectrograms_npy/\", \"*.h5\"))\nall_spec_paths = glob(os.path.join(\"/kaggle/input/ucu-hms-h5py/train_spectrograms_npy/train_spectrograms_npy\", \"*.h5\"))\nall_spec_shape = [\n    lazy_read_shape(path) for path in tqdm(all_spec_paths)\n]","metadata":{"execution":{"iopub.status.busy":"2024-03-21T20:02:24.452204Z","iopub.execute_input":"2024-03-21T20:02:24.45264Z","iopub.status.idle":"2024-03-21T20:03:58.614354Z","shell.execute_reply.started":"2024-03-21T20:02:24.452606Z","shell.execute_reply":"2024-03-21T20:03:58.613173Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"longest_spec_idx = np.argmax([el[1] for el in all_spec_shape])\nprint(\"Longest spec idx:\", longest_spec_idx, \"Shape:\", all_spec_shape[longest_spec_idx], \"Path:\", all_spec_paths[longest_spec_idx])","metadata":{"execution":{"iopub.status.busy":"2024-03-21T20:04:02.949097Z","iopub.execute_input":"2024-03-21T20:04:02.949597Z","iopub.status.idle":"2024-03-21T20:04:02.958559Z","shell.execute_reply.started":"2024-03-21T20:04:02.94956Z","shell.execute_reply":"2024-03-21T20:04:02.957465Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def read_first_600_frames(path):\n    with h5py.File(path, \"r\") as data_file:\n        specs_600 = np.stack(\n            [data_file[spec_name][:,:600] for spec_name in SPEC_TYPES], axis=-1\n        )\n    return specs_600","metadata":{"execution":{"iopub.status.busy":"2024-03-21T20:04:08.715584Z","iopub.execute_input":"2024-03-21T20:04:08.715985Z","iopub.status.idle":"2024-03-21T20:04:08.722042Z","shell.execute_reply.started":"2024-03-21T20:04:08.715955Z","shell.execute_reply":"2024-03-21T20:04:08.720701Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"seconds_for_partial_h5_read = timeit.timeit(lambda: read_first_600_frames(all_spec_paths[longest_spec_idx]), number=1000)\nseconds_for_full_h5_read = timeit.timeit(lambda: read_h5py_file(all_spec_paths[longest_spec_idx]), number=1000)\nseconds_for_full_parquet_read = timeit.timeit(lambda: pd.read_parquet(\n    all_spec_paths[longest_spec_idx].replace(\"_npy\", \"\").replace(\".h5\", \".parquet\")\n), number=1000)","metadata":{"execution":{"iopub.status.busy":"2024-03-21T20:04:20.881008Z","iopub.execute_input":"2024-03-21T20:04:20.881428Z","iopub.status.idle":"2024-03-21T20:05:51.870425Z","shell.execute_reply.started":"2024-03-21T20:04:20.881395Z","shell.execute_reply":"2024-03-21T20:05:51.868645Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"seconds_for_full_parquet_read","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\n    \"Raw .parquet file load (1000) took seconds:\", seconds_for_full_parquet_read,\n    \"\\nConverted .h5 file load (1000) took seconds:\", seconds_for_full_h5_read,\n    \"\\nFirst 600 frames of converted .h5 file load (1000) took seconds:\", seconds_for_partial_h5_read\n)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Correct segment computation was taken from [this discussion by Chris](https://www.kaggle.com/competitions/hms-harmful-brain-activity-classification/discussion/468010).","metadata":{}},{"cell_type":"code","source":"def visualise_eeg(input_df, prefix=\"\", wave_offset=None, segment_length=50.0):\n    for sensor, array in input_df.items():\n        if wave_offset is not None:\n            wave_to_vis = array[int(wave_offset * DEFAULT_SAMPLE_RATE):int((wave_offset + segment_length) * DEFAULT_SAMPLE_RATE)]\n            plt.title(prefix + \" \" + sensor)\n            plt.plot(wave_to_vis)\n            plt.show()\n\ndef visualise_eeg_spec(input_df, to_db=True, flip=True, prefix=\"\", spec_offset=None, segment_length=600):\n    for st in SPEC_TYPES:\n        freqs, times = input_df[\"freqs\"], input_df[\"times\"] \n        if spec_offset is not None:\n            mask = np.where(\n                (times>=spec_offset) & (times<spec_offset+segment_length)\n            )\n        else:\n            mask = np.ones_like(times).astype(bool)\n        plt.title(prefix + \" \" + st + f\" Time range: ({min(times[mask])}, {max(times[mask])}). Freq range: ({min(freqs)}, {max(freqs)})\")\n        spec = input_df[st]\n        spec = spec.T[mask].T\n        if to_db:\n            spec = librosa.amplitude_to_db(spec)\n        plt.imshow(spec)\n        if flip:\n            plt.gca().invert_yaxis()\n        plt.show()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df_row_id = 108\ntrain.iloc[train_df_row_id]","metadata":{"execution":{"iopub.status.busy":"2024-03-21T22:02:26.327924Z","iopub.execute_input":"2024-03-21T22:02:26.328583Z","iopub.status.idle":"2024-03-21T22:02:26.335951Z","shell.execute_reply.started":"2024-03-21T22:02:26.32855Z","shell.execute_reply":"2024-03-21T22:02:26.334488Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_wave = read_h5py_file(os.path.join(DATA_ROOT, f\"train_eegs_npy/{train['eeg_id'].iloc[train_df_row_id]}.h5\"))\nsample_spec = read_h5py_file(os.path.join(DATA_ROOT, f\"train_spectrograms_npy/{train['spectrogram_id'].iloc[train_df_row_id]}.h5\"))","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"![montage](images/montage.png)","metadata":{}},{"cell_type":"code","source":"sample_wave","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_spec","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"visualise_eeg(sample_wave, prefix=str(train['eeg_id'].iloc[train_df_row_id]), wave_offset=train['eeg_label_offset_seconds'].iloc[train_df_row_id])","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"visualise_eeg_spec(\n    sample_spec, prefix=str(train['eeg_id'].iloc[train_df_row_id]), to_db=True, \n    spec_offset=train['spectrogram_label_offset_seconds'].iloc[train_df_row_id]\n)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"> <font style=\"color:red\">**TODO**</font>: What is the `hop_size` of their spectogram? (1+ points)","metadata":{}},{"cell_type":"markdown","source":"<a id='Metric'></a>\n# Metric\n\nThe metric of the competition is [Kullback–Leibler divergence](https://en.wikipedia.org/wiki/Kullback%E2%80%93Leibler_divergence) on \"votes\".\n\nWe can find official implementation [here](https://www.kaggle.com/code/metric/kullback-leibler-divergence/notebook).\n\nBUT after the first submits we can figure out that we need to constraint all \"votes\" sum to 1.0, so how is metric computed on LB?\n\n[Answer](https://www.kaggle.com/competitions/hms-harmful-brain-activity-classification/discussion/476019): Just divide \"votes\" by their sum.","metadata":{}},{"cell_type":"code","source":"def correct_sum_to_one(input_values):\n    return input_values / input_values.sum(axis=1, keepdims=True)","metadata":{"execution":{"iopub.status.busy":"2024-03-21T22:01:27.417859Z","iopub.execute_input":"2024-03-21T22:01:27.418208Z","iopub.status.idle":"2024-03-21T22:01:27.423094Z","shell.execute_reply.started":"2024-03-21T22:01:27.418181Z","shell.execute_reply":"2024-03-21T22:01:27.421805Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport pandas.api.types\n\n# import kaggle_metric_utilities\n\nfrom typing import Optional\n\n\nclass ParticipantVisibleError(Exception):\n    pass\n\n\ndef kl_divergence(solution: pd.DataFrame, submission: pd.DataFrame, epsilon: float, micro_average: bool, sample_weights: Optional[pd.Series]):\n    # Overwrite solution for convenience\n    for col in solution.columns:\n        # Prevent issue with populating int columns with floats\n        if not pandas.api.types.is_float_dtype(solution[col]):\n            solution[col] = solution[col].astype(float)\n\n        # Clip both the min and max following Kaggle conventions for related metrics like log loss\n        # Clipping the max avoids cases where the loss would be infinite or undefined, clipping the min\n        # prevents users from playing games with the 20th decimal place of predictions.\n        submission[col] = np.clip(submission[col], epsilon, 1 - epsilon)\n\n        y_nonzero_indices = solution[col] != 0\n        solution[col] = solution[col].astype(float)\n        solution.loc[y_nonzero_indices, col] = solution.loc[y_nonzero_indices, col] * np.log(solution.loc[y_nonzero_indices, col] / submission.loc[y_nonzero_indices, col])\n        # Set the loss equal to zero where y_true equals zero following the scipy convention:\n        # https://docs.scipy.org/doc/scipy/reference/generated/scipy.special.rel_entr.html#scipy.special.rel_entr\n        solution.loc[~y_nonzero_indices, col] = 0\n\n    if micro_average:\n        return np.average(solution.sum(axis=1), weights=sample_weights)\n    else:\n        return np.average(solution.mean())\n\n\ndef score(\n        solution: pd.DataFrame,\n        submission: pd.DataFrame,\n        row_id_column_name: str,\n        epsilon: float=10**-15,\n        micro_average: bool=True,\n        sample_weights_column_name: Optional[str]=None\n    ) -> float:\n    ''' The Kullback–Leibler divergence.\n    The KL divergence is technically undefined/infinite where the target equals zero.\n\n    This implementation always assigns those cases a score of zero; effectively removing them from consideration.\n    The predictions in each row must add to one so any probability assigned to a case where y == 0 reduces\n    another prediction where y > 0, so crucially there is an important indirect effect.\n\n    https://en.wikipedia.org/wiki/Kullback%E2%80%93Leibler_divergence\n\n    solution: pd.DataFrame\n    submission: pd.DataFrame\n    epsilon: KL divergence is undefined for p=0 or p=1. If epsilon is not null, solution and submission probabilities are clipped to max(eps, min(1 - eps, p).\n    row_id_column_name: str\n    micro_average: bool. Row-wise average if True, column-wise average if False.\n\n    Examples\n    --------\n    >>> import pandas as pd\n    >>> row_id_column_name = \"id\"\n    >>> score(pd.DataFrame({'id': range(4), 'ham': [0, 1, 1, 0], 'spam': [1, 0, 0, 1]}), pd.DataFrame({'id': range(4), 'ham': [.1, .9, .8, .35], 'spam': [.9, .1, .2, .65]}), row_id_column_name=row_id_column_name)\n    0.216161...\n    >>> solution = pd.DataFrame({'id': range(3), 'ham': [0, 0.5, 0.5], 'spam': [0.1, 0.5, 0.5], 'other': [0.9, 0, 0]})\n    >>> submission = pd.DataFrame({'id': range(3), 'ham': [0, 0.5, 0.5], 'spam': [0.1, 0.5, 0.5], 'other': [0.9, 0, 0]})\n    >>> score(solution, submission, 'id')\n    0.0\n    >>> solution = pd.DataFrame({'id': range(3), 'ham': [0, 0.5, 0.5], 'spam': [0.1, 0.5, 0.5], 'other': [0.9, 0, 0]})\n    >>> submission = pd.DataFrame({'id': range(3), 'ham': [0.2, 0.3, 0.5], 'spam': [0.1, 0.5, 0.5], 'other': [0.7, 0.2, 0]})\n    >>> score(solution, submission, 'id')\n    0.160531...\n    '''\n    del solution[row_id_column_name]\n    del submission[row_id_column_name]\n\n    sample_weights = None\n    if sample_weights_column_name:\n        if sample_weights_column_name not in solution.columns:\n            raise ParticipantVisibleError(f'{sample_weights_column_name} not found in solution columns')\n        sample_weights = solution.pop(sample_weights_column_name)\n\n    if sample_weights_column_name and not micro_average:\n        raise ParticipantVisibleError('Sample weights are only valid if `micro_average` is `True`')\n\n    for col in solution.columns:\n        if col not in submission.columns:\n            raise ParticipantVisibleError(f'Missing submission column {col}')\n\n    # kaggle_metric_utilities.verify_valid_probabilities(solution, 'solution')\n    # kaggle_metric_utilities.verify_valid_probabilities(submission, 'submission')\n\n    # return kaggle_metric_utilities.safe_call_score(kl_divergence, solution, submission, epsilon=epsilon, micro_average=micro_average, sample_weights=sample_weights)\n\n    return kl_divergence(solution, submission, epsilon=epsilon, micro_average=micro_average, sample_weights=sample_weights)","metadata":{"execution":{"iopub.status.busy":"2024-03-21T22:01:30.626188Z","iopub.execute_input":"2024-03-21T22:01:30.627076Z","iopub.status.idle":"2024-03-21T22:01:30.638263Z","shell.execute_reply.started":"2024-03-21T22:01:30.627046Z","shell.execute_reply":"2024-03-21T22:01:30.637593Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"solution_temp = train[[\"label_id\"]].copy()\nsolution_temp[[el.lower() + \"_vote\" for el in TARGETS]] = correct_sum_to_one(\n    train[[el.lower() + \"_vote\" for el in TARGETS]].values\n)","metadata":{"execution":{"iopub.status.busy":"2024-03-21T22:01:37.568563Z","iopub.execute_input":"2024-03-21T22:01:37.568923Z","iopub.status.idle":"2024-03-21T22:01:37.581383Z","shell.execute_reply.started":"2024-03-21T22:01:37.568893Z","shell.execute_reply":"2024-03-21T22:01:37.58042Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_temp = train[[\"label_id\"]].copy()\nsubmission_temp[[el.lower() + \"_vote\" for el in TARGETS]] = correct_sum_to_one(\n    np.random.uniform(size=solution_temp[[el.lower() + \"_vote\" for el in TARGETS]].shape)\n)\n\nprint(\n    \"Random Uniform Prediction Score:\", score(solution_temp.copy(), submission_temp.copy(), \"label_id\")\n)","metadata":{"execution":{"iopub.status.busy":"2024-03-21T22:01:43.429784Z","iopub.execute_input":"2024-03-21T22:01:43.430304Z","iopub.status.idle":"2024-03-21T22:01:43.515192Z","shell.execute_reply.started":"2024-03-21T22:01:43.430275Z","shell.execute_reply":"2024-03-21T22:01:43.514173Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"> **TODO**: What if I do not do `*.copy()` in `score` ?","metadata":{}},{"cell_type":"code","source":"submission_temp = train[[\"label_id\"]].copy()\nsubmission_temp[[el.lower() + \"_vote\" for el in TARGETS]] = correct_sum_to_one(\n    np.ones(solution_temp[[el.lower() + \"_vote\" for el in TARGETS]].shape)\n)\n\nprint(\n    \"All Equal Vote Score:\", score(solution_temp.copy(), submission_temp.copy(), \"label_id\")\n)","metadata":{"execution":{"iopub.status.busy":"2024-03-21T20:06:56.931117Z","iopub.execute_input":"2024-03-21T20:06:56.931529Z","iopub.status.idle":"2024-03-21T20:06:57.030899Z","shell.execute_reply.started":"2024-03-21T20:06:56.931498Z","shell.execute_reply":"2024-03-21T20:06:57.029737Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_temp = train[[\"label_id\"]].copy()\nsubmission_temp[[el.lower() + \"_vote\" for el in TARGETS]] = solution_temp[[el.lower() + \"_vote\" for el in TARGETS]]\n\nprint(\n    \"All Correct Score:\", score(solution_temp.copy(), submission_temp.copy(), \"label_id\")\n)","metadata":{"execution":{"iopub.status.busy":"2024-03-21T20:07:00.416898Z","iopub.execute_input":"2024-03-21T20:07:00.41763Z","iopub.status.idle":"2024-03-21T20:07:00.512282Z","shell.execute_reply.started":"2024-03-21T20:07:00.417596Z","shell.execute_reply":"2024-03-21T20:07:00.511128Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_temp = train[[\"label_id\"]].copy()\nsubmission_temp[[el.lower() + \"_vote\" for el in TARGETS]] = np.ones_like(\n    solution_temp[[el.lower() + \"_vote\" for el in TARGETS]]\n)\n\nprint(\n    \"Invalid Score:\", score(solution_temp.copy(), submission_temp.copy(), \"label_id\")\n)","metadata":{"execution":{"iopub.status.busy":"2024-03-21T20:07:04.620946Z","iopub.execute_input":"2024-03-21T20:07:04.621342Z","iopub.status.idle":"2024-03-21T20:07:04.713832Z","shell.execute_reply.started":"2024-03-21T20:07:04.621313Z","shell.execute_reply":"2024-03-21T20:07:04.712615Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"> **TODO**: Experiment with other \"easy\" benchmarks","metadata":{}},{"cell_type":"code","source":"def numpy_kl_divergence(\n    y_true,\n    y_pred,\n    epsilon=10**-15,\n    micro_average=True\n):\n    y_pred_ = y_pred.copy().astype(np.float32)\n    y_true_ = y_true.copy().astype(np.float32)\n    \n    y_pred_ = np.clip(y_pred_, epsilon, 1 - epsilon)\n    y_nonzero_indices = y_true_ != 0\n    # print(\"Zero indices\", np.prod(y_pred_.size) - y_nonzero_indices.sum())\n    y_true_[y_nonzero_indices] = y_true_[y_nonzero_indices] * np.log(y_true_[y_nonzero_indices] / y_pred_[y_nonzero_indices])\n    y_true_[~y_nonzero_indices] = 0\n\n    if micro_average:\n        return np.average(y_true_.sum(axis=1), weights=None)\n    else:\n        return np.average(y_true_.mean(axis=0))","metadata":{"execution":{"iopub.status.busy":"2024-03-21T20:07:24.560021Z","iopub.execute_input":"2024-03-21T20:07:24.560417Z","iopub.status.idle":"2024-03-21T20:07:24.56761Z","shell.execute_reply.started":"2024-03-21T20:07:24.560385Z","shell.execute_reply":"2024-03-21T20:07:24.566746Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_temp = train[[\"label_id\"]].copy()\nsubmission_temp[[el.lower() + \"_vote\" for el in TARGETS]] = correct_sum_to_one(\n    np.random.uniform(size=solution_temp[[el.lower() + \"_vote\" for el in TARGETS]].shape)\n)","metadata":{"execution":{"iopub.status.busy":"2024-03-21T20:07:26.691141Z","iopub.execute_input":"2024-03-21T20:07:26.691559Z","iopub.status.idle":"2024-03-21T20:07:26.722265Z","shell.execute_reply.started":"2024-03-21T20:07:26.691519Z","shell.execute_reply":"2024-03-21T20:07:26.720932Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\n    \"Official Metric Score:\", score(solution_temp.copy(), submission_temp.copy(), \"label_id\")\n)\nprint(\n    \"Optimized Metric Score:\", numpy_kl_divergence(\n        solution_temp[[el.lower() + \"_vote\" for el in TARGETS]].values, \n        submission_temp[[el.lower() + \"_vote\" for el in TARGETS]].values,\n    )\n)","metadata":{"execution":{"iopub.status.busy":"2024-03-21T20:07:28.975549Z","iopub.execute_input":"2024-03-21T20:07:28.976708Z","iopub.status.idle":"2024-03-21T20:07:29.105175Z","shell.execute_reply.started":"2024-03-21T20:07:28.97664Z","shell.execute_reply":"2024-03-21T20:07:29.103972Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a id='Data_Preparation'></a>\n# Data Preparation\n\nWe will follow the idea of [this notebook](https://www.kaggle.com/code/crackle/efficientnetb0-pytorch-starter-lb-0-40),\nwith Spectograms extracted from original EEG by Chris in [his notebook](https://www.kaggle.com/code/cdeotte/how-to-make-spectrogram-from-eeg).\n\nOverall, Chris's logic is not entirely correct. He has computed one of the possible [EEG Montages](https://www.youtube.com/watch?v=AcW97nMLGEs).\n\nI strongly recommend check videos from next [YouTube channel](https://www.youtube.com/@jmoeller78/videos)","metadata":{}},{"cell_type":"markdown","source":"## Prepare Meta","metadata":{}},{"cell_type":"code","source":"aggr_train = train.groupby('eeg_id')[\n    ['spectrogram_id', 'spectrogram_label_offset_seconds']\n].agg({'spectrogram_id': 'first', 'spectrogram_label_offset_seconds': 'min'})\naggr_train.columns = ['spec_id', 'min']\n\ntmp = train.groupby('eeg_id')[\n    ['spectrogram_id','spectrogram_label_offset_seconds']\n].agg({'spectrogram_label_offset_seconds' :'max'})\naggr_train['max'] = tmp\n\ntmp = train.groupby('eeg_id')[['patient_id']].agg('first')\naggr_train['patient_id'] = tmp\n\ntmp = train.groupby('eeg_id')[[el.lower() + \"_vote\" for el in TARGETS]].agg('sum')\nfor t in [el.lower() + \"_vote\" for el in TARGETS]:\n    aggr_train[t] = tmp[t].values\n    \ny_data = aggr_train[[el.lower() + \"_vote\" for el in TARGETS]].values\n\ny_data = y_data / y_data.sum(axis=1, keepdims=True)\naggr_train[[el.lower() + \"_vote\" for el in TARGETS]] = y_data\n\ntmp = train.groupby('eeg_id')[[el.lower() + \"_vote\" for el in TARGETS]].apply(lambda df: np.median(df.values.sum(axis=1)))\n\naggr_train['votes_sum'] = tmp\n\ntmp = train.groupby('eeg_id')[['expert_consensus']].agg('first')\naggr_train['target'] = tmp\n\naggr_train[\"middle\"] = (aggr_train[\"min\"] + aggr_train[\"max\"]) // 4 \n\naggr_train = aggr_train.reset_index()\nprint('Train non-overlapp eeg_id shape:', aggr_train.shape )\naggr_train.head()","metadata":{"execution":{"iopub.status.busy":"2024-03-21T22:01:57.182579Z","iopub.execute_input":"2024-03-21T22:01:57.183272Z","iopub.status.idle":"2024-03-21T22:01:58.046155Z","shell.execute_reply.started":"2024-03-21T22:01:57.183233Z","shell.execute_reply":"2024-03-21T22:01:58.045201Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"aggr_train_row_id = np.where(aggr_train[\"eeg_id\"] == train[\"eeg_id\"].iloc[train_df_row_id])[0][0]","metadata":{"execution":{"iopub.status.busy":"2024-03-21T22:02:42.076415Z","iopub.execute_input":"2024-03-21T22:02:42.076778Z","iopub.status.idle":"2024-03-21T22:02:42.082712Z","shell.execute_reply.started":"2024-03-21T22:02:42.076751Z","shell.execute_reply":"2024-03-21T22:02:42.081548Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"aggr_train.iloc[aggr_train_row_id]","metadata":{"execution":{"iopub.status.busy":"2024-03-21T22:02:47.984745Z","iopub.execute_input":"2024-03-21T22:02:47.985114Z","iopub.status.idle":"2024-03-21T22:02:47.993135Z","shell.execute_reply.started":"2024-03-21T22:02:47.985059Z","shell.execute_reply":"2024-03-21T22:02:47.991881Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def process_kaggle_spec(\n    h5py_path, middle, display=False\n):\n    middle = int(middle)\n    X = np.zeros((128, 256, 4),dtype='float32')\n    specs = read_h5py_file(h5py_path)\n    if display:\n        plt.figure(figsize=(10,10))\n    for k_id, k in enumerate(SPEC_TYPES):\n        # EXTRACT 300 ROWS OF SPECTROGRAM\n        img = specs[k][:, middle:middle+300]\n        \n        # LOG TRANSFORM SPECTROGRAM\n        img = np.clip(img, np.exp(-4), np.exp(8))\n        img = np.log(img)\n        \n        # STANDARDIZE PER IMAGE\n        ep = 1e-6\n        m = np.nanmean(img.flatten())\n        s = np.nanstd(img.flatten())\n        img = (img - m) / (s + ep)\n        img = np.nan_to_num(img, nan=0.0)\n        \n        # CROP TO 256 TIME STEPS\n        X[14:-14, :, k_id] = img[:, 22:-22] / 2.0\n\n        if display:\n            eeg_id = os.path.splitext(os.path.basename(h5py_path))[0]\n            \n            plt.subplot(2,2,k_id+1)\n            plt.imshow(X[:,:,k_id],aspect='auto',origin='lower')\n            plt.title(f'EEG {eeg_id} - Spectrogram {k}')\n    return X","metadata":{"execution":{"iopub.status.busy":"2024-03-21T22:03:03.174057Z","iopub.execute_input":"2024-03-21T22:03:03.174766Z","iopub.status.idle":"2024-03-21T22:03:03.183935Z","shell.execute_reply.started":"2024-03-21T22:03:03.174727Z","shell.execute_reply":"2024-03-21T22:03:03.182651Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DATA_ROOT_N = \"/kaggle/input/ucu-hms-h5py\"\n\nprocess_kaggle_spec(\n    os.path.join(DATA_ROOT_N, f\"/train_spectrograms_npy/{aggr_train['spec_id'].iloc[aggr_train_row_id]}.h5\"),\n    aggr_train['middle'].iloc[aggr_train_row_id],\n    display=True\n);","metadata":{"execution":{"iopub.status.busy":"2024-03-21T22:04:39.007125Z","iopub.execute_input":"2024-03-21T22:04:39.007737Z","iopub.status.idle":"2024-03-21T22:04:39.161397Z","shell.execute_reply.started":"2024-03-21T22:04:39.007707Z","shell.execute_reply":"2024-03-21T22:04:39.159693Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def spectrogram_from_eeg(h5py_path, display=False):\n    \n    # LOAD MIDDLE 50 SECONDS OF EEG SERIES\n    eeg = read_h5py_file(h5py_path)\n    middle = (len(eeg[NEURAL_SENSORS[0]])-10_000)//2\n    for k in NEURAL_SENSORS:\n        eeg[k] = eeg[k][middle:middle+10_000]\n        \n    # VARIABLE TO HOLD SPECTROGRAM\n    img = np.zeros((128,256,4), dtype='float32')\n    \n    if display: \n        plt.figure(figsize=(12,12))\n    signals = []\n    for k in range(4):\n        COLS = REVERSED_FEATS[k]\n        \n        for kk in range(4):\n        \n            # COMPUTE PAIR DIFFERENCES\n            x = eeg[COLS[kk]] - eeg[COLS[kk+1]]\n\n            # FILL NANS\n            m = np.nanmean(x)\n            if np.isnan(x).mean() < 1:\n                x = np.nan_to_num(x, nan=m)\n            else:\n                x[:] = 0\n\n            # DENOISE\n            # if USE_WAVELET:\n            #     x = denoise(x, wavelet=USE_WAVELET)\n            signals.append(x)\n\n            # RAW SPECTROGRAM\n            mel_spec = librosa.feature.melspectrogram(y=x, sr=200, hop_length=len(x)//256, \n                  n_fft=1024, n_mels=128, fmin=0, fmax=20, win_length=128)\n\n            # LOG TRANSFORM\n            width = (mel_spec.shape[1]//32) * 32\n            mel_spec_db = librosa.power_to_db(mel_spec, ref=np.max).astype(np.float32)[:,:width]\n\n            # STANDARDIZE TO -1 TO 1\n            mel_spec_db = (mel_spec_db + 40) / 40 \n            img[:,:,k] += mel_spec_db\n                \n        # AVERAGE THE 4 MONTAGE DIFFERENCES\n        img[:,:,k] /= 4.0\n        \n        if display:\n            eeg_id = os.path.splitext(os.path.basename(h5py_path))[0]\n            plt.subplot(2,2,k+1)\n            plt.imshow(img[:,:,k], aspect='auto', origin='lower')\n            plt.title(f'EEG {eeg_id} - Spectrogram {REVERSED_FEATS[k]}')\n    \n    if display: \n        plt.show()\n        plt.figure(figsize=(10,10))\n        offset = 0\n        for k in range(4):\n            if k > 0:\n                offset -= signals[3-k].min()\n            plt.plot(range(10_000), signals[k] + offset, label=REVERSED_FEATS[3-k])\n            offset += signals[3-k].max()\n        plt.legend()\n        plt.title(f'EEG {eeg_id} Signals')\n        plt.show()\n        \n    return img\n","metadata":{"execution":{"iopub.status.busy":"2024-03-21T22:21:26.731763Z","iopub.execute_input":"2024-03-21T22:21:26.732459Z","iopub.status.idle":"2024-03-21T22:21:26.745009Z","shell.execute_reply.started":"2024-03-21T22:21:26.73241Z","shell.execute_reply":"2024-03-21T22:21:26.744249Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"> <font style=\"color:red\">**TODO**</font>: Re-write `spectrogram_from_eeg`, so it extracts only `signal` and than all specogtams are extracted using torchaudio/nnAudio. (2+ points)","metadata":{}},{"cell_type":"code","source":"def spectrogram_from_eeg_new(h5py_path, display=False):\n    \n    # LOAD MIDDLE 50 SECONDS OF EEG SERIES\n    eeg = read_h5py_file(h5py_path)\n    middle = (len(eeg[NEURAL_SENSORS[0]])-10_000)//2\n    signals = []\n    for k in NEURAL_SENSORS:\n        signal = torch.tensor(eeg[k][middle:middle+10_000], dtype=torch.float32)\n        signals.append(signal)\n        \n    # VARIABLE TO HOLD SPECTROGRAM\n    img = torch.zeros((128, 256, 4), dtype=torch.float32)\n    \n    if display: \n        plt.figure(figsize=(12, 12))\n    \n    for k in range(4):\n        COLS = REVERSED_FEATS[k]\n        signal_diff = signals[k] - signals[(k+1) % 4]\n        \n        # COMPUTE SPECTROGRAM\n        spec = torchaudio.transforms.MelSpectrogram(\n            sample_rate=200, n_fft=1024, hop_length=len(signal_diff)//256,\n            n_mels=128, f_min=0, f_max=20)(signal_diff.unsqueeze(0))\n        \n        # LOG TRANSFORM\n        spec_db = torchaudio.transforms.AmplitudeToDB()(spec)\n        \n        # STANDARDIZE TO -1 TO 1\n        spec_db = (spec_db + 40) / 40 \n        \n        # Resize to match the target shape\n        spec_db = torch.nn.functional.interpolate(spec_db.unsqueeze(0), size=(128, 256))\n        \n        # Add the spectrogram to the corresponding slice of the image\n        img[:, :, k] += spec_db.squeeze(0).squeeze(0)\n        \n        if display:\n            eeg_id = os.path.splitext(os.path.basename(h5py_path))[0]\n            plt.subplot(2, 2, k+1)\n            plt.imshow(img[:, :, k], aspect='auto', origin='lower')\n            plt.title(f'EEG {eeg_id} - Spectrogram {REVERSED_FEATS[k]}')\n    \n    if display: \n        plt.show()\n        plt.figure(figsize=(10, 10))\n        offset = 0\n        for k in range(4):\n            signal_data = signals[k].numpy()\n            if k > 0:\n                offset -= signal_data.min()\n            plt.plot(range(10_000), signal_data + offset, label=REVERSED_FEATS[3-k])\n            offset += signal_data.max()\n        plt.legend()\n        plt.title(f'EEG {eeg_id} Signals')\n        plt.show()\n        \n    return img.numpy()\n","metadata":{"execution":{"iopub.status.busy":"2024-03-21T22:26:07.748432Z","iopub.execute_input":"2024-03-21T22:26:07.74891Z","iopub.status.idle":"2024-03-21T22:26:07.762098Z","shell.execute_reply.started":"2024-03-21T22:26:07.748875Z","shell.execute_reply":"2024-03-21T22:26:07.760591Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"path = \"/kaggle/input/ucu-hms-h5py/train_eegs_npy/train_eegs_npy\"\nspectrogram_from_eeg_new(\n    os.path.join(path, f\"{aggr_train['eeg_id'].iloc[aggr_train_row_id]}.h5\"),\n#     os.path.join(DATA_ROOT, f\"/train_eegs_npy/train_eegs_npy/{aggr_train['eeg_id'].iloc[aggr_train_row_id]}.h5\"),\n    display=True\n)","metadata":{"execution":{"iopub.status.busy":"2024-03-21T22:26:11.872969Z","iopub.execute_input":"2024-03-21T22:26:11.873334Z","iopub.status.idle":"2024-03-21T22:26:14.338738Z","shell.execute_reply.started":"2024-03-21T22:26:11.873306Z","shell.execute_reply":"2024-03-21T22:26:14.337171Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"spectrogram_from_eeg(\n    os.path.join(path, f\"{aggr_train['eeg_id'].iloc[aggr_train_row_id]}.h5\"),\n#     os.path.join(DATA_ROOT, f\"/train_eegs_npy/train_eegs_npy/{aggr_train['eeg_id'].iloc[aggr_train_row_id]}.h5\"),\n    display=True\n)","metadata":{"execution":{"iopub.status.busy":"2024-03-21T22:21:37.604711Z","iopub.execute_input":"2024-03-21T22:21:37.605372Z","iopub.status.idle":"2024-03-21T22:21:39.257116Z","shell.execute_reply.started":"2024-03-21T22:21:37.605325Z","shell.execute_reply":"2024-03-21T22:21:39.256132Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"seconds_for_process_kaggle_spec = timeit.timeit(lambda: process_kaggle_spec(\n        os.path.join(DATA_ROOT, f\"train_spectrograms_npy/{aggr_train['spec_id'].iloc[aggr_train_row_id]}.h5\"),\n        aggr_train['middle'].iloc[aggr_train_row_id],\n        display=False\n    ),\n    number=100\n)\nseconds_for_process_spectrogram_from_eeg =  timeit.timeit(lambda: spectrogram_from_eeg(\n        os.path.join(DATA_ROOT, f\"train_eegs_npy/{aggr_train['eeg_id'].iloc[aggr_train_row_id]}.h5\"),\n        display=False\n    ),\n    number=100\n)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\n    \"process_kaggle_spec (100) took seconds:\", seconds_for_process_kaggle_spec,\n    \"\\nspectrogram_from_eeg (100) took seconds:\", seconds_for_process_spectrogram_from_eeg,\n)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def spectrogram_from_eeg_and_save(eeg_src_path, eeg_tgt_path):\n    \n    img = spectrogram_from_eeg(eeg_src_path, display=False)\n\n    with h5py.File(eeg_tgt_path, \"w\") as data_file:\n        for i, spec_type in enumerate(REVERSED_SPEC_TYPES):\n            data_file.create_dataset(spec_type, data=img[:,:,i]) ","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<font style=\"color:red\">Uncomment the next block:</font>","metadata":{}},{"cell_type":"code","source":"# eeg_file_pathes = glob(\n#     os.path.join(DATA_ROOT,\"train_eegs_npy/*.h5\")\n# )\n# print(f\"Found {len(eeg_file_pathes)} EEG files\")\n\n# os.makedirs(\n#     os.path.join(DATA_ROOT,\"train_spectrograms_reversed_npy\"), exist_ok=True\n# )\n# ProgressParallel(n_jobs=4, total=len(eeg_file_pathes))(\n#     delayed(spectrogram_from_eeg_and_save)(\n#         eeg_src_path=eeg_src_path,\n#         eeg_tgt_path=eeg_src_path.replace(\"train_eegs_npy\", \"train_spectrograms_reversed_npy\")\n#     )\n#     for eeg_src_path in eeg_file_pathes\n# );","metadata":{"scrolled":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def visualise_eeg_specs(input_df, prefix=\"\"):\n    plt.figure(figsize=(12,12))\n    for k, (feat_name, array) in enumerate(input_df.items()):\n        plt.subplot(2,2,k+1)\n        plt.imshow(array, aspect='auto', origin='lower')\n        plt.title(f'EEG {prefix} - Spectrogram {feat_name}')","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_eeg_spec = read_h5py_file(os.path.join(DATA_ROOT, f\"train_spectrograms_reversed_npy/{aggr_train['eeg_id'].iloc[aggr_train_row_id]}.h5\"))","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_eeg_spec","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"visualise_eeg_specs(sample_eeg_spec, aggr_train['eeg_id'].iloc[aggr_train_row_id])","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"> <font style=\"color:red\">**TODO**</font>: Explain why such \"aggregated\" approach leads to better LB score. (1+ points)","metadata":{}},{"cell_type":"markdown","source":"The “aggregated” approach leads to better LB score due to following reasons:\n\n* Aggregating data at the EEG level may capture more comprehensive information and more discriminative features for classification as a result\n* It can help reduce noise or variability of individual samples\n* It may provide a more generalized view of EEG characteristics across different patients - better generalization on unseen data\n* It provides a larger and more balanced dataset, which makes training more effective\n","metadata":{}},{"cell_type":"markdown","source":"<a id='Validation'></a>\n# Validation","metadata":{}},{"cell_type":"code","source":"split = list(StratifiedGroupKFold(n_splits=5, shuffle=True, random_state=42).split(\n    aggr_train,\n    aggr_train[\"target\"], \n    aggr_train[\"patient_id\"]\n))\nsplit","metadata":{"execution":{"iopub.status.busy":"2024-03-21T20:09:37.992266Z","iopub.execute_input":"2024-03-21T20:09:37.993315Z","iopub.status.idle":"2024-03-21T20:09:38.854987Z","shell.execute_reply.started":"2024-03-21T20:09:37.993276Z","shell.execute_reply":"2024-03-21T20:09:38.853832Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for train_idx, val_idx in split:\n    assert not set(aggr_train.iloc[train_idx].index) & set(aggr_train.iloc[val_idx].index)\n    assert not set(aggr_train.iloc[train_idx].patient_id) & set(aggr_train.iloc[val_idx].patient_id)","metadata":{"execution":{"iopub.status.busy":"2024-03-21T20:09:46.021198Z","iopub.execute_input":"2024-03-21T20:09:46.021596Z","iopub.status.idle":"2024-03-21T20:09:46.065584Z","shell.execute_reply.started":"2024-03-21T20:09:46.021564Z","shell.execute_reply":"2024-03-21T20:09:46.064722Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"OUTPUT_DIR = \"/kaggle/working/\"\n\nsplit = np.array(split, dtype=object)\nnp.save(os.path.join(OUTPUT_DIR, \"cv_split_v1.npy\"), split)","metadata":{"execution":{"iopub.status.busy":"2024-03-21T20:10:40.892948Z","iopub.execute_input":"2024-03-21T20:10:40.893657Z","iopub.status.idle":"2024-03-21T20:10:40.899906Z","shell.execute_reply.started":"2024-03-21T20:10:40.893607Z","shell.execute_reply":"2024-03-21T20:10:40.898982Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a id='Code_Preparation'></a>\n# Code Preparation\n\nIn this part, we will prepare all Modules needed for future training:\n\n- Dataset.\n- Loss.\n- [LightningModule](https://lightning.ai/docs/pytorch/stable/common/lightning_module.html).\n- Callbacks.","metadata":{}},{"cell_type":"markdown","source":"## [Dataset and Dataloader](https://pytorch.org/docs/stable/data.html)","metadata":{}},{"cell_type":"code","source":"class CombinedDataset(torch.utils.data.Dataset):\n    def __init__(\n        self,\n        root_eeg,\n        root_spec,\n        df,\n        target_col=\"target\",\n        target_cols=None,\n        spec_id_col=\"spec_id\",\n        eeg_id_col=\"eeg_id\",\n        middle_second_col=\"middle\",\n        transform=None,\n        test_mode=False,\n        specs_to_use=\"all\"\n    ):\n        assert specs_to_use in [\"all\", \"eeg_spec\", \"original_spec\"]\n        \n        self.df = df.reset_index(drop=True)\n\n        self.target_col = target_col\n        if target_cols is None:\n            self.target_cols = [el.lower() + \"_vote\" for el in TARGETS]\n        else:\n            self.target_cols = target_cols\n        self.name_col = spec_id_col\n        self.eeg_id_col = eeg_id_col\n        self.test_mode = test_mode\n        self.middle_second_col = middle_second_col\n        self.specs_to_use = specs_to_use\n\n        self.root_eeg = root_eeg\n        self.root_spec = root_spec\n\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def _prepare_sample(\n        self,\n        spec_id,\n        eeg_id,\n        middle\n    ):\n        eeg_path = os.path.join(self.root_eeg, f\"{eeg_id}.h5\")\n        spec_path = os.path.join(self.root_spec, f\"{spec_id}.h5\")\n        middle = int(middle)\n\n        if self.specs_to_use == \"all\":\n            n_specs = 8\n        else:\n            n_specs = 4\n        \n        X = np.zeros((128, 256, n_specs),dtype='float32')\n\n        if self.specs_to_use in [\"all\", \"original_spec\"]:\n            specs = read_h5py_file(spec_path)\n\n            for k_id, k in enumerate(SPEC_TYPES):\n                # EXTRACT 300 ROWS OF SPECTROGRAM\n                # import ipdb; ipdb.set_trace()\n                img = specs[k][:, middle:middle+300]\n                \n                # LOG TRANSFORM SPECTROGRAM\n                img = np.clip(img, np.exp(-4), np.exp(8))\n                img = np.log(img)\n                \n                # STANDARDIZE PER IMAGE\n                ep = 1e-6\n                m = np.nanmean(img.flatten())\n                s = np.nanstd(img.flatten())\n                img = (img - m) / (s + ep)\n                img = np.nan_to_num(img, nan=0.0)\n                \n                # CROP TO 256 TIME STEPS\n                X[14:-14, :, k_id] = img[:, 22:-22] / 2.0\n\n        if self.specs_to_use in [\"all\", \"eeg_spec\"]:\n            eeg_spec = read_h5py_file(eeg_path)\n\n            if self.specs_to_use == \"all\":\n                X[:, :, 4:] = np.stack([\n                    eeg_spec[k] for k in REVERSED_SPEC_TYPES\n                ], axis=-1)\n            else:\n                X[:, :, :] = np.stack([\n                    eeg_spec[k] for k in REVERSED_SPEC_TYPES\n                ], axis=-1)\n\n        \n        return X\n\n    def __getitem__(self, idx: int):\n        middle_second = self.df[self.middle_second_col].iloc[idx]\n        eeg_id = self.df[self.eeg_id_col].iloc[idx]\n        spec_id = self.df[self.name_col].iloc[idx]\n\n        if self.test_mode:\n            main_target = -1\n            all_targets = np.full(len(self.target_cols), -1.0)\n        else:\n            main_target = self.df[self.target_col].iloc[idx]\n            main_target = TARGET2ID[main_target]\n            all_targets = self.df[self.target_cols].iloc[idx].values\n\n        all_targets = torch.from_numpy(all_targets.astype(np.float32))\n        main_target = torch.tensor(main_target).long()\n        specs = self._prepare_sample(\n            spec_id=spec_id,\n            eeg_id=eeg_id,\n            middle=middle_second\n        )\n\n        if self.transform is not None:\n            specs = self.transform(image=specs)[\"image\"]\n\n        specs = specs.float()\n\n        return specs, main_target, all_targets, eeg_id","metadata":{"execution":{"iopub.status.busy":"2024-03-21T20:11:26.517034Z","iopub.execute_input":"2024-03-21T20:11:26.517404Z","iopub.status.idle":"2024-03-21T20:11:26.537485Z","shell.execute_reply.started":"2024-03-21T20:11:26.517375Z","shell.execute_reply":"2024-03-21T20:11:26.53637Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"> **TODO**: Implement MixUp and try other augmentations.","metadata":{}},{"cell_type":"code","source":"DATA_ROOT_N = \"/kaggle/input/ucu-hms-h5py/\"\n\ntesting_dataset = CombinedDataset(\n    root_eeg=os.path.join(DATA_ROOT_N, \"train_spectrograms_reversed_npy/train_spectrograms_reversed_npy\"),\n    root_spec=os.path.join(DATA_ROOT_N, \"train_spectrograms_npy/train_spectrograms_npy\"),\n    df=aggr_train,\n    transform=A.Compose([\n        ToTensorV2(transpose_mask=True),\n    ]),\n)","metadata":{"execution":{"iopub.status.busy":"2024-03-21T20:13:34.236019Z","iopub.execute_input":"2024-03-21T20:13:34.236462Z","iopub.status.idle":"2024-03-21T20:13:34.244349Z","shell.execute_reply.started":"2024-03-21T20:13:34.236421Z","shell.execute_reply":"2024-03-21T20:13:34.242876Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"testing_dataset[0]","metadata":{"execution":{"iopub.status.busy":"2024-03-21T20:13:35.946278Z","iopub.execute_input":"2024-03-21T20:13:35.946726Z","iopub.status.idle":"2024-03-21T20:13:36.112011Z","shell.execute_reply.started":"2024-03-21T20:13:35.946691Z","shell.execute_reply":"2024-03-21T20:13:36.111111Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"testing_specs, testing_main_target, testing_all_targets, testing_eeg_id = testing_dataset[120]\n\nprint(\"Main target\", testing_main_target)\nprint(\"All targets\", testing_all_targets)\nprint(\"EEG Id\", testing_eeg_id)\n\nprint(\"Shape\", testing_specs.shape)\n\nfor idx, spec_type in enumerate(SPEC_TYPES+REVERSED_SPEC_TYPES):\n    plt.title(spec_type)\n    plt.imshow(testing_specs[idx].numpy())\n    plt.gca().invert_yaxis()\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-03-21T20:14:21.992276Z","iopub.execute_input":"2024-03-21T20:14:21.993595Z","iopub.status.idle":"2024-03-21T20:14:24.532456Z","shell.execute_reply.started":"2024-03-21T20:14:21.993552Z","shell.execute_reply":"2024-03-21T20:14:24.530999Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"testing_dataloader = torch.utils.data.DataLoader(\n    testing_dataset,\n    batch_size=16,\n    shuffle=True,\n    drop_last=True,\n    num_workers=8,\n    pin_memory=True\n)\nfor batch in testing_dataloader:\n    break\n    \nprint(\n    \"Spec Batch Shape:\", batch[0].shape,\n    \"\\nMain Target Batch Shape:\", batch[1].shape,\n    \"\\nAll Target Batch Shape:\", batch[2].shape,\n    \"\\nEEG IDs Batch Shape:\", batch[3].shape,\n)","metadata":{"execution":{"iopub.status.busy":"2024-03-21T20:14:29.680633Z","iopub.execute_input":"2024-03-21T20:14:29.681096Z","iopub.status.idle":"2024-03-21T20:14:32.400033Z","shell.execute_reply.started":"2024-03-21T20:14:29.68106Z","shell.execute_reply":"2024-03-21T20:14:32.398453Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Loss\n\nWe will use a wrapper on top of [KLDivLoss](https://pytorch.org/docs/stable/generated/torch.nn.KLDivLoss.html).","metadata":{}},{"cell_type":"code","source":"class TorchKLDivergenceLoss(nn.Module):\n    def __init__(\n        self, \n        epsilon=1e-6,\n    ):\n        super(TorchKLDivergenceLoss, self).__init__()\n        self.epsilon = epsilon\n        self.kl_loss = nn.KLDivLoss(reduction=\"batchmean\", log_target=True)\n\n    def forward(self, y_pred_, y_true_):\n        y_true_ = y_true_ + self.epsilon\n        y_true_ = torch.log(y_true_ / y_true_.sum(dim=1, keepdim=True))\n        y_pred_ = F.log_softmax(y_pred_, dim=1)\n\n        return self.kl_loss(y_pred_, y_true_)","metadata":{"execution":{"iopub.status.busy":"2024-03-21T20:15:28.059846Z","iopub.execute_input":"2024-03-21T20:15:28.060431Z","iopub.status.idle":"2024-03-21T20:15:28.069722Z","shell.execute_reply.started":"2024-03-21T20:15:28.060391Z","shell.execute_reply":"2024-03-21T20:15:28.068418Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_target_batch = batch[2]\nsample_target_batch","metadata":{"execution":{"iopub.status.busy":"2024-03-21T20:15:31.565951Z","iopub.execute_input":"2024-03-21T20:15:31.566397Z","iopub.status.idle":"2024-03-21T20:15:31.577794Z","shell.execute_reply.started":"2024-03-21T20:15:31.566366Z","shell.execute_reply":"2024-03-21T20:15:31.576332Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_pred_batch = torch.randn_like(sample_target_batch)\nsample_pred_batch","metadata":{"execution":{"iopub.status.busy":"2024-03-21T20:15:34.917757Z","iopub.execute_input":"2024-03-21T20:15:34.918153Z","iopub.status.idle":"2024-03-21T20:15:34.92883Z","shell.execute_reply.started":"2024-03-21T20:15:34.918126Z","shell.execute_reply":"2024-03-21T20:15:34.927398Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"loss_func = TorchKLDivergenceLoss()\nprint(\"Random Prediction:\", loss_func(\n        sample_pred_batch, sample_target_batch\n    )\n)","metadata":{"execution":{"iopub.status.busy":"2024-03-21T20:15:37.423779Z","iopub.execute_input":"2024-03-21T20:15:37.424239Z","iopub.status.idle":"2024-03-21T20:15:37.452793Z","shell.execute_reply.started":"2024-03-21T20:15:37.424208Z","shell.execute_reply":"2024-03-21T20:15:37.451926Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"> **TODO**: Try other losses","metadata":{}},{"cell_type":"markdown","source":"## [LightningModule](https://lightning.ai/docs/pytorch/stable/common/lightning_module.html)","metadata":{}},{"cell_type":"code","source":"class LitTrainer(lightning.LightningModule):\n    def __init__(\n        self,\n        model,\n        forward,\n        optimizer,\n        scheduler,\n        scheduler_params,\n        batch_key,\n    ):\n        super().__init__()\n\n        self.model = model\n        self._forward = forward\n        self._optimizer = optimizer\n        self._scheduler = scheduler\n        self._scheduler_params = scheduler_params\n        self._batch_key = batch_key\n\n    def _aggregate_outputs(self, losses, inputs, outputs):\n        united = losses\n        united.update({\"input_\" + k: v for k, v in inputs.items()})\n        united.update({\"output_\" + k: v for k, v in outputs.items()})\n        return united\n\n    def training_step(self, batch):\n\n        start_time = time()\n        losses, inputs, outputs = self._forward(self, batch, epoch=self.current_epoch)\n        model_time = time() - start_time\n\n        for k, v in losses.items():\n            self.log(\n                \"train_\" + k,\n                v,\n                on_step=True,\n                on_epoch=False,\n                prog_bar=True,\n                logger=True,\n                batch_size=inputs[self._batch_key].shape[0],\n                sync_dist=True,\n            )\n            self.log(\n                \"train_avg_\" + k,\n                v,\n                on_step=False,\n                on_epoch=True,\n                prog_bar=True,\n                logger=True,\n                batch_size=inputs[self._batch_key].shape[0],\n                sync_dist=True,\n            )\n        self.log(\n            \"train_model_time\",\n            model_time,\n            on_step=True,\n            on_epoch=False,\n            prog_bar=True,\n            logger=True,\n            batch_size=1,\n            sync_dist=True,\n        )\n        self.log(\n            \"train_avg_model_time\",\n            model_time,\n            on_step=False,\n            on_epoch=True,\n            prog_bar=True,\n            logger=True,\n            batch_size=1,\n            sync_dist=True,\n        )\n\n        return self._aggregate_outputs(losses, inputs, outputs)\n\n    def validation_step(self, batch, batch_idx):\n\n        start_time = time()\n        losses, inputs, outputs = self._forward(self, batch, epoch=self.current_epoch)\n        model_time = time() - start_time\n\n        for k, v in losses.items():\n            self.log(\n                \"valid_\" + k,\n                v,\n                on_step=True,\n                on_epoch=False,\n                prog_bar=True,\n                logger=True,\n                batch_size=inputs[self._batch_key].shape[0],\n                sync_dist=True,\n            )\n            self.log(\n                \"valid_avg_\" + k,\n                v,\n                on_step=False,\n                on_epoch=True,\n                prog_bar=True,\n                logger=True,\n                batch_size=inputs[self._batch_key].shape[0],\n                sync_dist=True,\n            )\n        self.log(\n            \"valid_model_time\",\n            model_time,\n            on_step=True,\n            on_epoch=False,\n            prog_bar=True,\n            logger=True,\n            batch_size=1,\n            sync_dist=True,\n        )\n        self.log(\n            \"valid_avg_model_time\",\n            model_time,\n            on_step=False,\n            on_epoch=True,\n            prog_bar=True,\n            logger=True,\n            batch_size=1,\n            sync_dist=True,\n        )\n\n        return self._aggregate_outputs(losses, inputs, outputs)\n\n    def configure_optimizers(self):\n        scheduler = {\"scheduler\": self._scheduler}\n        scheduler.update(self._scheduler_params)\n        return (\n            [self._optimizer],\n            [scheduler],\n        )","metadata":{"execution":{"iopub.status.busy":"2024-03-21T20:15:50.311912Z","iopub.execute_input":"2024-03-21T20:15:50.312356Z","iopub.status.idle":"2024-03-21T20:15:50.330198Z","shell.execute_reply.started":"2024-03-21T20:15:50.312319Z","shell.execute_reply":"2024-03-21T20:15:50.329182Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MultilabelSpecClsForward(nn.Module):\n    def __init__(\n        self,\n        loss_function,\n        output_key=\"logits\",\n        input_key=\"all_targets\",\n    ):\n        super().__init__()\n        self.loss_function = loss_function\n        self.output_key = output_key\n        self.input_key = input_key\n\n    def forward(self, runner, batch, epoch=None):\n\n        specs, main_target, all_targets, _ = batch\n\n        output = runner.model(specs)\n\n        inputs = {\n            \"specs\": specs,\n            \"main_target\": main_target,\n            \"all_targets\": all_targets,\n        }\n\n        losses = {\n            \"loss\": self.loss_function(\n                output[self.output_key],\n                inputs[self.input_key],\n            )\n        }\n\n        return losses, inputs, output","metadata":{"execution":{"iopub.status.busy":"2024-03-21T20:15:55.198994Z","iopub.execute_input":"2024-03-21T20:15:55.199443Z","iopub.status.idle":"2024-03-21T20:15:55.206979Z","shell.execute_reply.started":"2024-03-21T20:15:55.199408Z","shell.execute_reply":"2024-03-21T20:15:55.205885Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## [Callbacks](https://lightning.ai/docs/pytorch/stable/extensions/callbacks.html)","metadata":{}},{"cell_type":"code","source":"class KLCallback(Callback):\n    def __init__(\n        self,\n        pred_key: str = \"preds\",\n        gt_label_key: str = \"all_targets\",\n        loader_names: Tuple = (\"valid\"),\n        verbose: bool = False,\n    ):\n        self.pred_key = pred_key\n        self.gt_label_key = gt_label_key\n        self.loader_names = loader_names\n        self.accums = {\n            loader_name: {\n                \"preds\": [],\n                \"gt_labels\": [],\n            }\n            for loader_name in loader_names\n        }\n        self.verbose = verbose\n\n    def initialize_accums(self, loader_name):\n        self.accums[loader_name] = {\n            \"preds\": [],\n            \"gt_labels\": [],\n        }\n\n    def update_accums(self, outputs, loader_name):\n        pred = outputs[\"output_\" + self.pred_key].detach()\n        pred = torch.softmax(pred, dim=1)\n        pred = pred.cpu().numpy()\n\n        gt_label = outputs[\"input_\" + self.gt_label_key].detach()\n        gt_label = gt_label.cpu().numpy()\n\n        self.accums[loader_name][\"preds\"].append(pred)\n        self.accums[loader_name][\"gt_labels\"].append(gt_label)\n\n    def compute_kl(self, pl_module, loader_name):\n        preds = np.concatenate(self.accums[loader_name][\"preds\"], axis=0) \n        gt_labels = np.concatenate(self.accums[loader_name][\"gt_labels\"], axis=0)\n        gt_labels = gt_labels / gt_labels.sum(axis=1, keepdims=True)\n\n        kl = numpy_kl_divergence(\n            y_true=gt_labels,\n            y_pred=preds,\n            micro_average=True,\n        )\n\n        pl_module.log(\n            loader_name + \"_kl\",\n            kl,\n        )\n\n    def on_train_epoch_start(self, trainer, pl_module):\n        if \"train\" in self.loader_names:\n            self.initialize_accums(\"train\")\n\n    def on_train_batch_end(self, trainer, pl_module, outputs, batch, batch_idx, dataloader_idx=0):\n        if \"train\" in self.loader_names:\n            self.update_accums(outputs, \"train\")\n\n    def on_train_epoch_end(self, trainer, pl_module):\n        if \"train\" in self.loader_names:\n            self.compute_kl(pl_module, \"train\")\n            self.initialize_accums(\"train\")\n\n    def on_validation_epoch_start(self, trainer, pl_module):\n        if \"valid\" in self.loader_names:\n            self.initialize_accums(\"valid\")\n\n    def on_validation_batch_end(self, trainer, pl_module, outputs, batch, batch_idx, dataloader_idx=0):\n        if \"valid\" in self.loader_names:\n            self.update_accums(outputs, \"valid\")\n\n    def on_validation_epoch_end(self, trainer, pl_module):\n        if \"valid\" in self.loader_names:\n            self.compute_kl(pl_module, \"valid\")\n            self.initialize_accums(\"valid\")","metadata":{"execution":{"iopub.status.busy":"2024-03-21T20:15:59.406353Z","iopub.execute_input":"2024-03-21T20:15:59.40765Z","iopub.status.idle":"2024-03-21T20:15:59.424141Z","shell.execute_reply.started":"2024-03-21T20:15:59.407608Z","shell.execute_reply":"2024-03-21T20:15:59.422498Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a id='Model'></a>\n# Model\n\nWe will use CNN based approach and build next Pipeline:\n1. Normalization Layer.\n2. (Optional) Augmentation Layer.\n3. CNN Encoder.\n4. Pooling Layer.\n5. Classification Layer.\n\nSomething near the architecture, which I have used on [BirdCLEF 2023](https://www.kaggle.com/competitions/birdclef-2023/discussion/412808).\n\n![nn_arch](images/nn_arch.png)\n\nAs CNN backbone we will use [EfficientNet](https://arxiv.org/abs/1905.11946) family with pretrains from [timm repo](https://github.com/huggingface/pytorch-image-models).\n\n> **TODO**: Try other CNN architecture \"families\".","metadata":{}},{"cell_type":"code","source":"class NormalizeMelSpec(nn.Module):\n    def __init__(\n        self,\n        eps=1e-6,\n        normalize_standart=True,\n        normalize_minmax=True,\n    ):\n        super().__init__()\n        self.eps = eps\n        self.normalize_standart = normalize_standart\n        self.normalize_minmax = normalize_minmax\n\n    def forward(self, X):\n        if self.normalize_standart:\n            mean = X.mean((2, 3), keepdim=True)\n            std = X.std((2, 3), keepdim=True)\n            X = (X - mean) / (std + self.eps)\n            \n        if self.normalize_minmax:\n            norm_max = torch.amax(X, dim=(2, 3), keepdim=True)\n            norm_min = torch.amin(X, dim=(2, 3), keepdim=True)\n            X = (X - norm_min) / (norm_max - norm_min + self.eps)\n\n        return X\n\nclass CustomMasking(nn.Module):\n    def __init__(self, mask_max_length: int, mask_max_masks: int, p=1.0, inplace=True):\n        super().__init__()\n        assert isinstance(mask_max_masks, int) and mask_max_masks > 0\n        self.mask_max_masks = mask_max_masks\n        self.mask_max_length = mask_max_length\n        self.mask_module = None\n        self.p = p\n        self.inplace = inplace\n\n    def forward(self, x):\n        if not self.inplace:\n            output = x.clone()\n        for i in range(x.shape[0]):\n            if np.random.binomial(n=1, p=self.p):\n                n_applies = np.random.randint(low=1, high=self.mask_max_masks + 1)\n                for _ in range(n_applies):\n                    if self.inplace:\n                        x[i : i + 1] = self.mask_module(x[i : i + 1])\n                    else:\n                        output[i : i + 1] = self.mask_module(output[i : i + 1])\n        if self.inplace:\n            return x\n        else:\n            return output\n\n\nclass CustomTimeMasking(CustomMasking):\n    def __init__(self, mask_max_length: int, mask_max_masks: int, p=1.0, inplace=True):\n        super().__init__(mask_max_length=mask_max_length, mask_max_masks=mask_max_masks, p=p, inplace=inplace)\n        self.mask_module = TimeMasking(time_mask_param=mask_max_length)\n\n\nclass CustomFreqMasking(CustomMasking):\n    def __init__(self, mask_max_length: int, mask_max_masks: int, p=1.0, inplace=True):\n        super().__init__(mask_max_length=mask_max_length, mask_max_masks=mask_max_masks, p=p, inplace=inplace)\n        self.mask_module = FrequencyMasking(freq_mask_param=mask_max_length)\n\nclass SpecCNNClasifier(nn.Module):\n    def __init__(\n        self,\n        backbone: str,\n        device: str,\n        n_specs: int,\n        n_classes: int,\n        classifier_dropout: float = 0.5,\n        normalize_config: Dict[str, bool] = {\n            \"normalize_standart\": True,\n            \"normalize_minmax\": True,\n        },\n        pretrained: bool = True,\n        timm_kwargs: Optional[Dict] = None,\n        spec_augment_config: Optional[Dict[str, Any]] = None,\n    ):\n        super().__init__()\n        timm_kwargs = {} if timm_kwargs is None else timm_kwargs\n        self.device = device\n\n        self.instance_norm = NormalizeMelSpec(\n            **normalize_config\n        )\n        \n        if spec_augment_config is not None:\n            self.spec_augment = []\n            if \"freq_mask\" in spec_augment_config:\n                self.spec_augment.append(CustomFreqMasking(**spec_augment_config[\"freq_mask\"]))\n            if \"time_mask\" in spec_augment_config:\n                self.spec_augment.append(CustomTimeMasking(**spec_augment_config[\"time_mask\"]))\n            self.spec_augment = nn.Sequential(*self.spec_augment)\n        else:\n            self.spec_augment = None\n\n        \n        self.backbone = timm.create_model(\n            backbone,\n            features_only=True,\n            pretrained=pretrained,\n            in_chans=n_specs,\n            exportable=True,\n            **timm_kwargs,\n        )\n\n        self.pool = nn.AdaptiveAvgPool2d((1, 1))\n        self.classifier = nn.Sequential(\n            nn.Dropout(p=classifier_dropout),\n            nn.Linear(self.backbone.feature_info.channels()[-1], n_classes),\n        )\n        \n        self.to(self.device)\n\n    def forward(self, input, return_spec_feature=False, return_cnn_emb=False):\n        processed_spec = self.instance_norm(input)\n        if self.spec_augment is not None and self.training:\n            processed_spec = self.spec_augment(processed_spec)\n        if return_spec_feature:\n            return processed_spec\n            \n        emb = self.backbone(processed_spec)[-1]\n        if return_cnn_emb:\n            return emb\n\n        bs, ch, h, w = emb.shape\n        emb = self.pool(emb)\n        emb = emb.view(bs, ch)\n\n        logits = self.classifier(emb)\n\n        return {\"logits\": logits}","metadata":{"execution":{"iopub.status.busy":"2024-03-21T20:16:06.645899Z","iopub.execute_input":"2024-03-21T20:16:06.646533Z","iopub.status.idle":"2024-03-21T20:16:06.670985Z","shell.execute_reply.started":"2024-03-21T20:16:06.6465Z","shell.execute_reply":"2024-03-21T20:16:06.669491Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"testing_model = SpecCNNClasifier(\n    backbone=\"tf_efficientnet_b0.in1k\",\n    device=\"cpu\",\n    n_specs=8,\n    n_classes=6,\n    # spec_augment_config={\n    #     \"freq_mask\": {\n    #         \"mask_max_length\": 20,\n    #         \"mask_max_masks\": 5,\n    #         \"p\": 0.5,\n    #         \"inplace\": True,\n    #     },\n    #     \"time_mask\": {\n    #         \"mask_max_length\": 30,\n    #         \"mask_max_masks\": 5,\n    #         \"p\": 0.5,\n    #         \"inplace\": True,\n    #     },\n    # }\n)","metadata":{"execution":{"iopub.status.busy":"2024-03-21T20:16:15.591447Z","iopub.execute_input":"2024-03-21T20:16:15.591902Z","iopub.status.idle":"2024-03-21T20:16:16.976146Z","shell.execute_reply.started":"2024-03-21T20:16:15.591867Z","shell.execute_reply":"2024-03-21T20:16:16.974925Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Original input spectograms:","metadata":{}},{"cell_type":"code","source":"for idx, spec_type in enumerate(SPEC_TYPES+REVERSED_SPEC_TYPES):\n    plt.title(spec_type)\n    plt.imshow(batch[0][0,idx].numpy())\n    plt.gca().invert_yaxis()\n    plt.show()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Transformed Spectograms:","metadata":{}},{"cell_type":"code","source":"transformed_spec = testing_model(\n    batch[0],\n    return_spec_feature=True\n)","metadata":{"execution":{"iopub.status.busy":"2024-03-21T20:16:34.223164Z","iopub.execute_input":"2024-03-21T20:16:34.223552Z","iopub.status.idle":"2024-03-21T20:16:34.270726Z","shell.execute_reply.started":"2024-03-21T20:16:34.223524Z","shell.execute_reply":"2024-03-21T20:16:34.269716Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for idx, spec_type in enumerate(SPEC_TYPES+REVERSED_SPEC_TYPES):\n    plt.title(spec_type)\n    plt.imshow(transformed_spec[0,idx].numpy())\n    plt.gca().invert_yaxis()\n    plt.show()","metadata":{"collapsed":true,"jupyter":{"outputs_hidden":true}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Transformed + Augmneted Spectograms:","metadata":{}},{"cell_type":"code","source":"testing_model = SpecCNNClasifier(\n    backbone=\"tf_efficientnet_b0.in1k\",\n    device=\"cpu\",\n    n_specs=8,\n    n_classes=6,\n    spec_augment_config={\n        \"freq_mask\": {\n            \"mask_max_length\": 20,\n            \"mask_max_masks\": 5,\n            \"p\": 1.0,\n            \"inplace\": True,\n        },\n        \"time_mask\": {\n            \"mask_max_length\": 30,\n            \"mask_max_masks\": 5,\n            \"p\": 1.0,\n            \"inplace\": True,\n        },\n    }\n)","metadata":{"execution":{"iopub.status.busy":"2024-03-21T20:16:51.095646Z","iopub.execute_input":"2024-03-21T20:16:51.096567Z","iopub.status.idle":"2024-03-21T20:16:51.303589Z","shell.execute_reply.started":"2024-03-21T20:16:51.09653Z","shell.execute_reply":"2024-03-21T20:16:51.30204Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transformed_spec = testing_model(\n    batch[0],\n    return_spec_feature=True\n)","metadata":{"execution":{"iopub.status.busy":"2024-03-21T20:16:55.686226Z","iopub.execute_input":"2024-03-21T20:16:55.686643Z","iopub.status.idle":"2024-03-21T20:16:55.801213Z","shell.execute_reply.started":"2024-03-21T20:16:55.68661Z","shell.execute_reply":"2024-03-21T20:16:55.800265Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for idx, spec_type in enumerate(SPEC_TYPES+REVERSED_SPEC_TYPES):\n    plt.title(spec_type)\n    plt.imshow(transformed_spec[0,idx].numpy())\n    plt.gca().invert_yaxis()\n    plt.show()","metadata":{"collapsed":true,"jupyter":{"outputs_hidden":true}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cnn_embed = testing_model(\n    batch[0],\n    return_cnn_emb=True\n)","metadata":{"execution":{"iopub.status.busy":"2024-03-21T20:17:09.820296Z","iopub.execute_input":"2024-03-21T20:17:09.823085Z","iopub.status.idle":"2024-03-21T20:17:11.330063Z","shell.execute_reply.started":"2024-03-21T20:17:09.82302Z","shell.execute_reply":"2024-03-21T20:17:11.32888Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for idx, spec_type in enumerate(SPEC_TYPES+REVERSED_SPEC_TYPES):\n    plt.title(spec_type)\n    plt.imshow(cnn_embed[0,idx].detach().numpy())\n    plt.gca().invert_yaxis()\n    plt.show()","metadata":{"collapsed":true,"jupyter":{"outputs_hidden":true}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"testing_model(batch[0])","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a id='Training'></a>\n# Training","metadata":{}},{"cell_type":"code","source":"def lightning_training(\n    train_df: Optional[pd.DataFrame],\n    val_df: Optional[pd.DataFrame],\n    exp_name: str,\n    fold_id: Optional[int],\n    forward_batch_key: str,\n    train_dataset_class: torch.utils.data.Dataset,\n    val_dataset_class: Optional[torch.utils.data.Dataset],\n    train_dataset_config: dict,\n    val_dataset_config: Optional[dict],\n    train_dataloader_config: dict,\n    val_dataloader_config: Optional[dict],\n    nn_model_class: torch.nn.Module,\n    nn_model_config: dict,\n    optimizer_init: Callable,\n    scheduler_init: Callable,\n    scheduler_params: dict,\n    forward: Union[torch.nn.Module, Callable],\n    # It is not really Callable. It just lambda that will init List of callbacks\n    # each time. It is just done for safe CV training.\n    callbacks: Optional[Callable],\n    n_epochs: int,\n    main_metric: str,\n    metric_mode: str,\n    checkpoint_callback_params: dict = {},\n    tensorboard_logger_params: dict = {},\n    trainer_params: dict = {},\n    precision_mode: str = \"32-true\",\n    n_checkpoints_to_save: int = 3,\n    log_every_n_steps: int = 100,\n    train_strategy: str = \"auto\",\n):\n    # Set device\n    if torch.cuda.is_available():\n        device = \"cuda\"\n    else:\n        device = \"cpu\"\n    print(f\"Training Device : {device}\")\n\n    train_dataset = train_dataset_class(\n        df=train_df,\n        **train_dataset_config,\n    )\n    val_dataset = val_dataset_class(\n        df=val_df,\n        **val_dataset_config,\n    )\n\n    loaders = {\n        \"train\": torch.utils.data.DataLoader(\n            train_dataset,\n            **train_dataloader_config,\n        ),\n        \"valid\": torch.utils.data.DataLoader(\n            val_dataset, \n            **val_dataloader_config\n        )\n    }\n\n    model = nn_model_class(device=device, **nn_model_config)\n\n    for k in loaders.keys():\n        print(f\"{k} Loader Len = {len(loaders[k])}\")\n\n    optimizer = optimizer_init(model)\n    scheduler = scheduler_init(optimizer, len(loaders[\"train\"]))\n\n    if not isinstance(forward, torch.nn.Module):\n        forward = forward()\n\n    lightning_model = LitTrainer(\n        model,\n        forward=forward,\n        optimizer=optimizer,\n        scheduler=scheduler,\n        scheduler_params=scheduler_params,\n        batch_key=forward_batch_key\n    )\n\n    all_callbacks = [\n        ModelCheckpoint(\n            dirpath=os.path.join(exp_name, \"checkpoints\"),\n            save_top_k=n_checkpoints_to_save,\n            mode=metric_mode,\n            monitor=main_metric,\n            **checkpoint_callback_params,\n        ),\n        LearningRateMonitor(logging_interval=\"step\"),\n    ]\n    if callbacks is not None:\n        all_callbacks += callbacks()\n\n    tensorboard_logger = pl_loggers.TensorBoardLogger(\n        save_dir=os.path.join(exp_name, \"tensorboard\"),\n        **tensorboard_logger_params,\n    )\n    trainer = lightning.Trainer(\n        devices=4,\n        precision=precision_mode,\n        strategy=train_strategy,\n        max_epochs=n_epochs,\n        logger=tensorboard_logger,\n        log_every_n_steps=log_every_n_steps,\n        callbacks=all_callbacks,\n        **trainer_params,\n    )\n    trainer.fit(model=lightning_model, train_dataloaders=loaders[\"train\"], val_dataloaders=loaders[\"valid\"])","metadata":{"execution":{"iopub.status.busy":"2024-03-21T20:21:39.405763Z","iopub.execute_input":"2024-03-21T20:21:39.406856Z","iopub.status.idle":"2024-03-21T20:21:39.426369Z","shell.execute_reply.started":"2024-03-21T20:21:39.406811Z","shell.execute_reply":"2024-03-21T20:21:39.424941Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"aggr_train.iloc[:5000]","metadata":{"execution":{"iopub.status.busy":"2024-03-21T20:17:33.262899Z","iopub.execute_input":"2024-03-21T20:17:33.263906Z","iopub.status.idle":"2024-03-21T20:17:33.29268Z","shell.execute_reply.started":"2024-03-21T20:17:33.263866Z","shell.execute_reply":"2024-03-21T20:17:33.291259Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Let's start a test run - train and test on the same data. The idea is to check if model CAN be optimized","metadata":{}},{"cell_type":"code","source":"lightning_training(\n    train_df=aggr_train.iloc[:5000],\n    val_df=aggr_train.iloc[:5000],\n    exp_name=\"logdirs/debug_run\",\n    fold_id=0,\n    forward_batch_key=\"specs\",\n    train_dataset_class=CombinedDataset,\n    val_dataset_class=CombinedDataset,\n    train_dataset_config=dict(\n        root_eeg=os.path.join(DATA_ROOT_N, \"train_spectrograms_reversed_npy/train_spectrograms_reversed_npy\"),\n        root_spec=os.path.join(DATA_ROOT_N, \"train_spectrograms_npy/train_spectrograms_npy\"),\n        transform=A.Compose([\n            ToTensorV2(transpose_mask=True),\n        ]),\n    ),\n    val_dataset_config=dict(\n        root_eeg=os.path.join(DATA_ROOT_N, \"train_spectrograms_reversed_npy/train_spectrograms_reversed_npy\"),\n        root_spec=os.path.join(DATA_ROOT_N, \"train_spectrograms_npy/train_spectrograms_npy\"),\n        transform=A.Compose([\n            ToTensorV2(transpose_mask=True),\n        ]),\n    ),\n    train_dataloader_config={\n        \"batch_size\": 64,\n        \"shuffle\": False,\n        \"drop_last\": True,\n        \"num_workers\": 8,\n        \"pin_memory\": True,\n    },\n    val_dataloader_config={\n        \"batch_size\": 64,\n        \"shuffle\": False,\n        \"drop_last\": False,\n        \"num_workers\": 8,\n        \"pin_memory\": True,\n    },\n    nn_model_class=SpecCNNClasifier,\n    nn_model_config=dict(\n        backbone=\"tf_efficientnet_b0.in1k\",\n        n_specs=8,\n        n_classes=6,\n        spec_augment_config={\n            \"freq_mask\": {\n                \"mask_max_length\": 20,\n                \"mask_max_masks\": 5,\n                \"p\": 1.0,\n                \"inplace\": True,\n            },\n            \"time_mask\": {\n                \"mask_max_length\": 30,\n                \"mask_max_masks\": 5,\n                \"p\": 1.0,\n                \"inplace\": True,\n            },\n        }\n    ),\n    optimizer_init=lambda model: torch.optim.Adam(model.parameters(), lr=1e-3),\n    scheduler_init=lambda optimizer, len_train: torch.optim.lr_scheduler.CosineAnnealingLR(\n        optimizer, \n        eta_min=1e-6,\n        T_max=len_train * 3,\n    ),\n    scheduler_params={\"interval\": \"step\", \"monitor\": \"valid_kl\"},\n    forward=lambda: MultilabelSpecClsForward(\n        TorchKLDivergenceLoss(),\n        input_key=\"all_targets\",\n        output_key=\"logits\",\n    ),\n    callbacks=lambda: [\n        KLCallback(\n            pred_key=\"logits\",\n            gt_label_key=\"all_targets\",\n            loader_names=(\"valid\", \"train\"),\n            verbose=True,\n        )\n    ],\n    n_epochs=3,\n    main_metric=\"valid_kl\",\n    metric_mode=\"min\",\n    checkpoint_callback_params=dict(\n        save_last=True,\n        auto_insert_metric_name=True,\n        save_weights_only=True,\n        save_on_train_epoch_end=True,\n        filename=\"{epoch}-{step}-{valid_kl:.3f}\",\n    ),\n    precision_mode=\"16-mixed\",\n)","metadata":{"execution":{"iopub.status.busy":"2024-03-21T20:21:49.34694Z","iopub.execute_input":"2024-03-21T20:21:49.347373Z","iopub.status.idle":"2024-03-21T20:22:41.384978Z","shell.execute_reply.started":"2024-03-21T20:21:49.347338Z","shell.execute_reply":"2024-03-21T20:22:41.336709Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Okay, looks like it works!\n\nNow let's run on 5 folds with more epochs:","metadata":{}},{"cell_type":"code","source":"for fold_idx in range(5):\n    print(f\"Start Fold {fold_idx} ...\")\n    lightning_training(\n        train_df=aggr_train.iloc[split[fold_idx][0]],\n        val_df=aggr_train.iloc[split[fold_idx][1]],\n        exp_name=os.path.join(\"logdirs/hms_baseline\", f\"fold_{fold_idx}\"),\n        fold_id=fold_idx,\n        forward_batch_key=\"specs\",\n        train_dataset_class=CombinedDataset,\n        val_dataset_class=CombinedDataset,\n        train_dataset_config=dict(\n            root_eeg=os.path.join(DATA_ROOT, \"train_spectrograms_reversed_npy\"),\n            root_spec=os.path.join(DATA_ROOT, \"train_spectrograms_npy\"),\n            transform=A.Compose([\n                ToTensorV2(transpose_mask=True),\n            ]),\n        ),\n        val_dataset_config=dict(\n            root_eeg=os.path.join(DATA_ROOT, \"train_spectrograms_reversed_npy\"),\n            root_spec=os.path.join(DATA_ROOT, \"train_spectrograms_npy\"),\n            transform=A.Compose([\n                ToTensorV2(transpose_mask=True),\n            ]),\n        ),\n        train_dataloader_config={\n            \"batch_size\": 64,\n            \"shuffle\": False,\n            \"drop_last\": True,\n            \"num_workers\": 8,\n            \"pin_memory\": True,\n        },\n        val_dataloader_config={\n            \"batch_size\": 64,\n            \"shuffle\": False,\n            \"drop_last\": False,\n            \"num_workers\": 8,\n            \"pin_memory\": True,\n        },\n        nn_model_class=SpecCNNClasifier,\n        nn_model_config=dict(\n            backbone=\"tf_efficientnet_b0.in1k\",\n            n_specs=8,\n            n_classes=6,\n            spec_augment_config={\n                \"freq_mask\": {\n                    \"mask_max_length\": 20,\n                    \"mask_max_masks\": 5,\n                    \"p\": 1.0,\n                    \"inplace\": True,\n                },\n                \"time_mask\": {\n                    \"mask_max_length\": 30,\n                    \"mask_max_masks\": 5,\n                    \"p\": 1.0,\n                    \"inplace\": True,\n                },\n            }\n        ),\n        optimizer_init=lambda model: torch.optim.Adam(model.parameters(), lr=1e-3),\n        scheduler_init=lambda optimizer, len_train: torch.optim.lr_scheduler.CosineAnnealingLR(\n            optimizer, \n            eta_min=1e-6,\n            T_max=len_train * 12,\n        ),\n        scheduler_params={\"interval\": \"step\", \"monitor\": \"valid_kl\"},\n        forward=lambda: MultilabelSpecClsForward(\n            TorchKLDivergenceLoss(),\n            input_key=\"all_targets\",\n            output_key=\"logits\",\n        ),\n        callbacks=lambda: [\n            KLCallback(\n                pred_key=\"logits\",\n                gt_label_key=\"all_targets\",\n                loader_names=(\"valid\", \"train\"),\n                verbose=True,\n            )\n        ],\n        n_epochs=12,\n        main_metric=\"valid_kl\",\n        metric_mode=\"min\",\n        checkpoint_callback_params=dict(\n            save_last=True,\n            auto_insert_metric_name=True,\n            save_weights_only=True,\n            save_on_train_epoch_end=True,\n            filename=\"{epoch}-{step}-{valid_kl:.3f}\",\n        ),\n        precision_mode=\"16-mixed\",\n    )\n    print(f\"End Fold {fold_idx}!\")","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls logdirs/hms_baseline/","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls logdirs/hms_baseline/fold_0/","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls logdirs/hms_baseline/fold_0/checkpoints/","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a id='Evaluation'></a>\n# Evaluation","metadata":{}},{"cell_type":"markdown","source":"## Preparation for inference","metadata":{}},{"cell_type":"code","source":"val_dfs = [\n    aggr_train.iloc[split[fold_idx][1]] for fold_idx in range(5)\n]","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_datasets = [\n    CombinedDataset(\n        root_eeg=os.path.join(DATA_ROOT, \"train_spectrograms_reversed_npy\"),\n        root_spec=os.path.join(DATA_ROOT, \"train_spectrograms_npy\"),\n        test_mode=True,\n        transform=A.Compose([\n            ToTensorV2(transpose_mask=True),\n        ]),\n        df=val_df\n    )\n    for val_df in val_dfs\n]","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_dataloaders = [\n    torch.utils.data.DataLoader(\n        dataset,\n        batch_size=64,\n        drop_last=False,\n        shuffle=False,\n        num_workers=4\n    )\n    for dataset in val_datasets\n]","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def delete_prefix_from_chkp(chkp_dict: OrderedDict, prefix: str):\n    new_dict = OrderedDict()\n    for k in chkp_dict.keys():\n        if k.startswith(prefix):\n            new_dict[k[len(prefix) :]] = chkp_dict[k]\n        else:\n            new_dict[k] = chkp_dict[k]\n\n    return new_dict\n\ndef create_model_and_load_best_checkpoint(\n    model_class,\n    model_config,\n    model_device,\n    model_chkp_root,\n    model_chkp_regex,\n    sort_rule,\n    delete_prefix=None,\n):\n    basenames = os.listdir(model_chkp_root)\n    checkpoints = []\n    for el in basenames:\n        matches = re.findall(model_chkp_regex, el)\n        if not matches:\n            continue\n        parsed_dict = {key: value for key, value in matches}\n        parsed_dict[\"name\"] = el\n        checkpoints.append(parsed_dict)\n    print(\"All checkpoints\")\n    pprint(checkpoints)\n    checkpoints = sorted(checkpoints, key=sort_rule)\n    print(\"Sorted checkpoints\")\n    pprint(checkpoints)\n    best_checkpoint = os.path.join(model_chkp_root, checkpoints[0][\"name\"])\n    print(\"Best checkpoint\")\n    print(best_checkpoint)\n    t_chkp = torch.load(\n        best_checkpoint, \n        map_location=\"cpu\"\n    )[\"state_dict\"]\n    if delete_prefix is not None:\n        t_chkp = delete_prefix_from_chkp(t_chkp, delete_prefix)\n    t_model = model_class(**model_config, device=model_device)\n    t_model.load_state_dict(t_chkp)\n    t_model.eval()\n\n    return t_model","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = [create_model_and_load_best_checkpoint(\n        model_class=SpecCNNClasifier,\n        model_config=dict(\n            backbone=\"tf_efficientnet_b0.in1k\",\n            n_specs=8,\n            n_classes=6,\n            spec_augment_config={\n                \"freq_mask\": {\n                    \"mask_max_length\": 20,\n                    \"mask_max_masks\": 5,\n                    \"p\": 1.0,\n                    \"inplace\": True,\n                },\n                \"time_mask\": {\n                    \"mask_max_length\": 30,\n                    \"mask_max_masks\": 5,\n                    \"p\": 1.0,\n                    \"inplace\": True,\n                },\n            }\n        ),\n        model_device=\"cuda\",\n        model_chkp_root=f\"logdirs/hms_baseline/fold_{m_i}/checkpoints\",\n        model_chkp_regex=r'(?P<key>\\w+)=(?P<value>[\\d.]+)(?=\\.ckpt|$)',\n        sort_rule=lambda x: float(x[\"valid_kl\"]),\n        delete_prefix=\"model.\"\n) for m_i in range(5)]","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Run inference","metadata":{}},{"cell_type":"code","source":"@torch.inference_mode()\ndef inference_function(\n    loader,\n    nn_model,\n    output_key,\n    device\n):\n    predicted_df = {\"eeg_id\": []}\n    for target in TARGETS:\n        predicted_df[target.lower() + \"_vote\"] = []\n\n    for batch in tqdm(loader):\n        specs, _, _, eeg_id = batch\n        pred_probs = nn_model(specs.to(device))[output_key]\n        pred_probs = torch.softmax(pred_probs, dim=1)\n        pred_probs = pred_probs.detach().cpu().numpy()\n        predicted_df[\"eeg_id\"].append(eeg_id.cpu().numpy())\n        for i, target in enumerate(TARGETS):\n            predicted_df[target.lower() + \"_vote\"].append(pred_probs[:, i])\n\n    for key in predicted_df:\n        predicted_df[key] = np.concatenate(predicted_df[key])\n\n    predicted_df = pd.DataFrame(predicted_df)\n\n    return predicted_df","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"all_predicted_dfs = []\nfor fold_idx in range(5):\n    all_predicted_dfs.append(inference_function(\n        loader=val_dataloaders[fold_idx],\n        nn_model=model[fold_idx],\n        output_key=\"logits\",\n        device=\"cuda\"\n    ))","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Evaluate results","metadata":{}},{"cell_type":"code","source":"for fold_idx in range(5):\n    val_dfs[fold_idx][\"fold\"] = fold_idx\n\nall_val_dfs = pd.concat(val_dfs).reset_index(drop=True)\nall_predicted_dfs = pd.concat(all_predicted_dfs).reset_index(drop=True)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"assert (all_val_dfs[\"eeg_id\"] == all_predicted_dfs[\"eeg_id\"]).all()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"all_val_dfs.head()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"all_predicted_dfs.head()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"oof_score = score(\n    solution=all_val_dfs[[el.lower() + \"_vote\" for el in TARGETS] + [\"eeg_id\"]].copy(),\n    submission=all_predicted_dfs.copy(),\n    row_id_column_name=\"eeg_id\"\n)\n\nscores = []\nfor fold_idx in range(5):\n    df_mask = all_val_dfs[\"fold\"] == fold_idx\n    scores.append(score(\n        solution=all_val_dfs.loc[df_mask, [el.lower() + \"_vote\" for el in TARGETS] + [\"eeg_id\"]].copy(),\n        submission=all_predicted_dfs[df_mask].copy(),\n        row_id_column_name=\"eeg_id\"\n    ))\n\nscores = np.array(scores)\navg_score = round(scores.mean(), 4)\nstd_score = round(scores.std(), 4)\nmin_score = round(scores.min(), 4)\nmax_score = round(scores.max(), 4)\nscore_lower_bound = np.clip(round(avg_score - 3*std_score, 4), 0, np.inf)\nscore_higher_bound = round(avg_score + 3*std_score, 4)\nprint(\"All Folds KL scores\")\nprint(scores)\npprint({\n    'kl_oof': oof_score,\n    'kl_avg': avg_score,\n    'kl_std': std_score,\n    'kl_min': min_score,\n    'kl_max': max_score,\n    'kl_lower_bound_estimation': score_lower_bound,\n    'kl_higher_bound_estimation': score_higher_bound\n})","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"> **TODO**: Explore Model mistakes. Maybe you can find the reason for severe overfitting.","metadata":{}},{"cell_type":"markdown","source":"## Inference on Kaggle\n\nAs HMS is a [Code Competition](https://www.kaggle.com/docs/competitions#notebooks-only-competitions), you have to inference with the help of the [Inference Notebook](https://www.kaggle.com/code/vladimirsydor/ucu-hms-inference/notebook).\n\nHere is our baseline LB Score: 0.52\n\n![lb_score](images/lb_score.png)","metadata":{}},{"cell_type":"markdown","source":"<a id='Homework'></a>\n# Homework\n\nTheory (5 points):\n- Follow links.\n- Try to fill/do **TODO** comments.\n- Answer theory questions in the Google Form.\n\nPractice (10 points):\n\nImprove HMS baseline. Some ideas to try:\n\n- Work more with EDA. Maybe you will find important insights. Pay attention to class balance/imbalance.\n- Data preparation is very specific and might be sub-optimal. Think about alternative approaches.\n- Experiment with other Neural Net Models.\n    - Try 1D approaches, which work with pure signal.\n- Twick other hyperparameters like Optimizers, Schedulers, Batch Size, Number of Epochs, and so on.\n- Solution to \"Overfitting\" maybe a \"Killing Feature\".\n\n**IMPORTANT** : Make sure to submit your baseline to Kaggle and get Leaderboard score\n\n> After the competition ends, high-ranking solutions (starting from the bronze zone) will receive additional points. ","metadata":{}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}