{"metadata":{"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"},{"sourceId":7392775,"sourceType":"datasetVersion","datasetId":4297782},{"sourceId":7447509,"sourceType":"datasetVersion","datasetId":4334995},{"sourceId":7818085,"sourceType":"datasetVersion","datasetId":4511932}],"dockerImageVersionId":30664,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true},"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.10.13"},"papermill":{"default_parameters":{},"duration":11188.402769,"end_time":"2024-03-11T17:52:49.511492","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2024-03-11T14:46:21.108723","version":"2.5.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Problem specification","metadata":{"papermill":{"duration":0.015696,"end_time":"2024-03-11T14:46:23.771516","exception":false,"start_time":"2024-03-11T14:46:23.755820","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"- Using the notebook **[HMS | EfficientNetB0 PyTorch [Train]](https://www.kaggle.com/code/alejopaullier/hms-efficientnetb0-pytorch-train/)** as reference. The data processing pipeline is essentialy the same.","metadata":{"papermill":{"duration":0.014717,"end_time":"2024-03-11T14:46:23.801545","exception":false,"start_time":"2024-03-11T14:46:23.786828","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"# Importing libraries","metadata":{"papermill":{"duration":0.016349,"end_time":"2024-03-11T14:46:23.833218","exception":false,"start_time":"2024-03-11T14:46:23.816869","status":"completed"},"tags":[]}},{"cell_type":"code","source":"!pip install torchsummary","metadata":{"papermill":{"duration":13.489077,"end_time":"2024-03-11T14:46:37.337319","exception":false,"start_time":"2024-03-11T14:46:23.848242","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# data and moddeling\nimport torch\nfrom torchvision import transforms, datasets, models\nfrom torch import nn, optim\nfrom torch.utils import data\nimport torch.nn.functional as F\n\n# model evaluation\nfrom sklearn.model_selection import train_test_split\n\n# plotting and visualization\nfrom torchsummary import summary\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nsns.set()\n\n# others\nimport zipfile\nimport os\nfrom glob import glob\nimport numpy as np\nimport pandas as pd\nimport albumentations as A\nimport random\nfrom typing import Dict, List\nimport time\nimport gc","metadata":{"papermill":{"duration":7.901501,"end_time":"2024-03-11T14:46:45.254512","exception":false,"start_time":"2024-03-11T14:46:37.353011","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# setting the device\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nprint('Using', torch.cuda.device_count(), 'GPU(s)')","metadata":{"papermill":{"duration":0.090933,"end_time":"2024-03-11T14:46:45.361245","exception":false,"start_time":"2024-03-11T14:46:45.270312","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Utils","metadata":{"papermill":{"duration":0.015019,"end_time":"2024-03-11T14:46:45.391712","exception":false,"start_time":"2024-03-11T14:46:45.376693","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class config:\n    NUM_WORKERS = 2\n    BATCH_SIZE_TRAIN = 16\n    BATCH_SIZE_VALID = 16\n    EPOCHS = 10\n    LR = 1e-6\n    WEIGHT_DECAY = 1e-2\n    DEVICE = DEVICE = device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    ","metadata":{"papermill":{"duration":0.02294,"end_time":"2024-03-11T14:46:45.429971","exception":false,"start_time":"2024-03-11T14:46:45.407031","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class paths:\n    MODEL_PATH_PREFIX = \"/kaggle/input/pytorch-models-hms-competition/sub23_EfficientNet_v2_s_best_\"\n    OUTPUT_DIR = \"/kaggle/working/\"\n    PRE_LOADED_EEGS = '/kaggle/input/brain-eeg-spectrograms/eeg_specs.npy'\n    PRE_LOADED_SPECTROGRAMS = '/kaggle/input/brain-spectrograms/specs.npy'\n    TRAIN_CSV = \"/kaggle/input/hms-harmful-brain-activity-classification/train.csv\"\n    TRAIN_EEGS = \"/kaggle/input/brain-eeg-spectrograms/EEG_Spectrograms/\"\n    TRAIN_SPECTROGRAMS = \"/kaggle/input/hms-harmful-brain-activity-classification/train_spectrograms/\"","metadata":{"papermill":{"duration":0.023107,"end_time":"2024-03-11T14:46:45.468375","exception":false,"start_time":"2024-03-11T14:46:45.445268","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class AverageMeter:\n    \"\"\"Computes and stores the average and current value\"\"\"\n    # source: https://kaiyangzhou.github.io/deep-person-reid/_modules/torchreid/utils/avgmeter.html#AverageMeter\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\n\n    def update(self, val, n=1):\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count\n\n\ndef sep():\n    print(\"-\"*35)\n    \n\ndef plot_spectrogram(spectrogram_path: str):\n    \"\"\"\n    Source: https://www.kaggle.com/code/mvvppp/hms-eda-and-domain-journey\n    Visualize spectrogram recordings from a parquet file.\n    :param spectrogram_path: path to the spectrogram parquet.\n    \"\"\"\n    sample_spect = pd.read_parquet(spectrogram_path)\n\n    split_spect = {\n        \"LL\": sample_spect.filter(regex='^LL', axis=1),\n        \"RL\": sample_spect.filter(regex='^RL', axis=1),\n        \"RP\": sample_spect.filter(regex='^RP', axis=1),\n        \"LP\": sample_spect.filter(regex='^LP', axis=1),\n    }\n\n    fig, axes = plt.subplots(nrows=2, ncols=2, figsize=(15, 12))\n    axes = axes.flatten()\n    label_interval = 5\n    for i, split_name in enumerate(split_spect.keys()):\n        ax = axes[i]\n        img = ax.imshow(np.log(split_spect[split_name]).T, cmap='viridis', aspect='auto', origin='lower')\n        cbar = fig.colorbar(img, ax=ax)\n        cbar.set_label('Log(Value)')\n        ax.set_title(split_name)\n        ax.set_ylabel(\"Frequency (Hz)\")\n        ax.set_xlabel(\"Time\")\n\n        ax.set_yticks(np.arange(len(split_spect[split_name].columns)))\n        ax.set_yticklabels([column_name[3:] for column_name in split_spect[split_name].columns])\n        frequencies = [column_name[3:] for column_name in split_spect[split_name].columns]\n        ax.set_yticks(np.arange(0, len(split_spect[split_name].columns), label_interval))\n        ax.set_yticklabels(frequencies[::label_interval])\n    plt.tight_layout()\n    plt.show()\n\n\ndef train_val_split(dataset, train_percentage=0.85):\n    global transform, transform_eval\n    length = len(dataset)\n    indices = list(range(length))\n    np.random.shuffle(indices)\n    len_train = int(np.floor(train_percentage * length))\n    train_idx, val_idx = indices[:len_train], indices[len_train:]\n    return train_idx, val_idx","metadata":{"papermill":{"duration":0.033686,"end_time":"2024-03-11T14:46:45.517209","exception":false,"start_time":"2024-03-11T14:46:45.483523","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data processing","metadata":{"papermill":{"duration":0.014836,"end_time":"2024-03-11T14:46:45.547570","exception":false,"start_time":"2024-03-11T14:46:45.532734","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"## Load data","metadata":{"papermill":{"duration":0.014788,"end_time":"2024-03-11T14:46:45.577423","exception":false,"start_time":"2024-03-11T14:46:45.562635","status":"completed"},"tags":[]}},{"cell_type":"code","source":"df = pd.read_csv(paths.TRAIN_CSV)\nprint(f'Shape of the training dataframe: {df.shape}')\ndf.head()","metadata":{"papermill":{"duration":0.291366,"end_time":"2024-03-11T14:46:45.883771","exception":false,"start_time":"2024-03-11T14:46:45.592405","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"label_columns = df.columns[-6:]\nprint(f'We have {len(label_columns)} classes: {list(label_columns)}')","metadata":{"papermill":{"duration":0.023053,"end_time":"2024-03-11T14:46:45.923121","exception":false,"start_time":"2024-03-11T14:46:45.900068","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"target_preds = [c + \"_pred\" for c in list(label_columns)]\nlabel_to_num = {'Seizure': 0, 'LPD': 1, 'GPD': 2, 'LRDA': 3, 'GRDA': 4, 'Other':5}\nnum_to_label = {v: k for k, v in label_to_num.items()}","metadata":{"papermill":{"duration":0.023135,"end_time":"2024-03-11T14:46:45.961794","exception":false,"start_time":"2024-03-11T14:46:45.938659","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data pre-processing","metadata":{"papermill":{"duration":0.015464,"end_time":"2024-03-11T14:46:45.993074","exception":false,"start_time":"2024-03-11T14:46:45.977610","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"### Create Non-Overlapping eeg_id training data\nFor this part, we are using the processing based in this [this](https://www.kaggle.com/competitions/hms-harmful-brain-activity-classification/discussion/468010) discussion.\n\nSummarizing: The information provided in the competition data description states that the test data does not include multiple crops from the same eeg_id. Consequently, our training and validation processes will exclusively utilize one crop per eeg_id.","metadata":{"papermill":{"duration":0.015457,"end_time":"2024-03-11T14:46:46.024219","exception":false,"start_time":"2024-03-11T14:46:46.008762","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# grouping by eeg_id + getting the first value of the spectrogram id and he min value of the label offset\ntrain_gp = df.groupby('eeg_id')\ntrain_df = train_gp[['spectrogram_id','spectrogram_label_offset_seconds']].agg({\n    'spectrogram_id':'first',\n    'spectrogram_label_offset_seconds':'min'\n})\n\ntrain_df.columns = ['spectrogram_id','min']\n\nprint(train_df.shape)\ntrain_df.head()","metadata":{"papermill":{"duration":0.047495,"end_time":"2024-03-11T14:46:46.087496","exception":false,"start_time":"2024-03-11T14:46:46.040001","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# now getting the max value for the offset and inserting in the training dataset\naux_df = train_gp[['spectrogram_id','spectrogram_label_offset_seconds']].agg({\n    'spectrogram_label_offset_seconds':'max'\n})\n\ntrain_df['max'] = aux_df","metadata":{"papermill":{"duration":0.025746,"end_time":"2024-03-11T14:46:46.129088","exception":false,"start_time":"2024-03-11T14:46:46.103342","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# getting the patient id\naux_df = train_gp[['patient_id']].agg('first')\ntrain_df['patient_id'] = aux_df","metadata":{"papermill":{"duration":0.024767,"end_time":"2024-03-11T14:46:46.169650","exception":false,"start_time":"2024-03-11T14:46:46.144883","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# inserting the labels and agregating by summing the votes\naux_df = train_gp[label_columns].agg('sum')\nfor label in label_columns:\n    train_df[label] = aux_df[label].values","metadata":{"papermill":{"duration":0.030865,"end_time":"2024-03-11T14:46:46.216569","exception":false,"start_time":"2024-03-11T14:46:46.185704","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# normalizing the summed values of the spredictions\ny = train_df[label_columns].values\ny = y / y.sum(axis=1,keepdims=True)\ntrain_df[label_columns] = y","metadata":{"papermill":{"duration":0.025661,"end_time":"2024-03-11T14:46:46.296952","exception":false,"start_time":"2024-03-11T14:46:46.271291","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Label name\naux_df = train_gp[['expert_consensus']].agg('first')\ntrain_df['target'] = aux_df","metadata":{"papermill":{"duration":0.042262,"end_time":"2024-03-11T14:46:46.355079","exception":false,"start_time":"2024-03-11T14:46:46.312817","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = train_df.reset_index()\nprint('Train non-overlapp eeg_id shape:', train_df.shape )\ntrain_df.head()","metadata":{"papermill":{"duration":0.036975,"end_time":"2024-03-11T14:46:46.407840","exception":false,"start_time":"2024-03-11T14:46:46.370865","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Read the training spectrograms\n\nIn this section we'll read the spectrograms parquet files. This would normally take 11mintes by using pandas, hence the amount of spectrograms is very high (~11k). Alternatively, we are reading 1 file from [this](https://www.kaggle.com/datasets/cdeotte/brain-spectrograms) Kaggle Dataset, which we can read in less than 1 minute.\n\nThe resulting dictionary has the spectrogram_id as the keys and the spectrogram sequence as values. The shape of the those reads is: (timesteps, 400), a 2-dimensional np.array.\n\nThe second dimension of the array can be separated in 4 (LL, RL, LP, RP), each 100 of those correspond to a region where the electrodes are placed.","metadata":{"papermill":{"duration":0.016022,"end_time":"2024-03-11T14:46:46.440036","exception":false,"start_time":"2024-03-11T14:46:46.424014","status":"completed"},"tags":[]}},{"cell_type":"code","source":"%%time\npaths_spectrograms = glob(paths.TRAIN_SPECTROGRAMS + \"*.parquet\")\nprint(f'There are {len(paths_spectrograms)} spectrogram parquets')\n\nall_spectrograms = np.load(paths.PRE_LOADED_SPECTROGRAMS, allow_pickle=True).item()","metadata":{"papermill":{"duration":65.515711,"end_time":"2024-03-11T14:47:51.971878","exception":false,"start_time":"2024-03-11T14:46:46.456167","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"idx = np.random.randint(0,len(paths_spectrograms))\nspectrogram_path = paths_spectrograms[idx]\nplot_spectrogram(spectrogram_path)","metadata":{"papermill":{"duration":3.602493,"end_time":"2024-03-11T14:47:55.590940","exception":false,"start_time":"2024-03-11T14:47:51.988447","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Read the EEG\n\nThe resulting all_eegs dictionary contains eeg_id as keys (int keys) and the values are the eeg sequences (as 3-dimensional np.array) of shape (128, 256, 4)","metadata":{"papermill":{"duration":0.029976,"end_time":"2024-03-11T14:47:55.651751","exception":false,"start_time":"2024-03-11T14:47:55.621775","status":"completed"},"tags":[]}},{"cell_type":"code","source":"%%time\npaths_eegs = glob(paths.TRAIN_EEGS + \"*.npy\")\nprint(f'There are {len(paths_eegs)} EEG spectrograms')\n\nall_eegs = np.load(paths.PRE_LOADED_EEGS, allow_pickle=True).item()","metadata":{"papermill":{"duration":84.487291,"end_time":"2024-03-11T14:49:20.169378","exception":false,"start_time":"2024-03-11T14:47:55.682087","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Creating validation and training partitions","metadata":{"papermill":{"duration":0.037769,"end_time":"2024-03-11T14:49:20.244889","exception":false,"start_time":"2024-03-11T14:49:20.207120","status":"completed"},"tags":[]}},{"cell_type":"code","source":"train_df, val_df = train_test_split(train_df, test_size=0.2)","metadata":{"papermill":{"duration":0.0464,"end_time":"2024-03-11T14:49:20.322194","exception":false,"start_time":"2024-03-11T14:49:20.275794","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Creating the dataset\n\nCustom dataset to load the data","metadata":{"papermill":{"duration":0.030061,"end_time":"2024-03-11T14:49:20.382418","exception":false,"start_time":"2024-03-11T14:49:20.352357","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class HMSDataset(data.Dataset):\n    def __init__(self, df:pd.DataFrame, spect:Dict[int, np.ndarray], eeg_spect:Dict[int, np.ndarray],\n                 use_eeg_spec:bool=True, use_kaggle_spec:bool=True, channel_cat:bool=False, reshape:bool=True, augment:bool=False, mode:str='train'):\n        self.df = df\n        self.augment = augment\n        self.mode = mode\n        self.use_eeg_spec = use_eeg_spec\n        self.use_kaggle_spec = use_kaggle_spec\n        self.reshape = reshape\n        self.spectrograms = spect\n        self.eeg_spectrograms = eeg_spect\n        self.channel_cat = channel_cat\n        \n    \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, index):\n        X, y = self.__generate_data(index)\n        \n        if self.augment:\n            X = self.__transform(X)\n            \n        X = torch.tensor(X, dtype=torch.float32)\n        y = torch.tensor(y, dtype=torch.float32)\n                \n        if self.reshape:\n            X = self.__reshape_input(X)\n        \n        return X, y\n    \n    def __reshape_input(self, x:np.ndarray):\n        \"\"\"\n        Converts the 8-channel tensor created using the spectrograms and the eeg, \n        into an 3 channel image.\n        \n        returns:\n            X:(512,512,3) np.ndarray\n        \"\"\"\n        # Retrieve spectrograms from input array\n        spectrograms = [x[i:i+1, :, :] for i in range(4)]\n        spectrograms = torch.cat(spectrograms, dim=1)\n        \n        # Retrieve eeg_spectrograms from input array\n        eegs = [x[i:i+1, :, :] for i in range(4,8)]\n        eegs = torch.cat(eegs, dim=1)\n        \n        # Define which spectrograms are going to be used for training\n        if self.use_eeg_spec & self.use_kaggle_spec:\n            if self.channel_cat:\n                x = torch.cat([spectrograms, eegs], dim=0)\n            else:\n                x = torch.cat([spectrograms, eegs], dim=2)\n        elif self.use_eeg_spec:\n            x = eegs\n        else:\n            x = spectrograms\n        \n        if not self.channel_cat:\n            x = torch.cat([x,x,x], dim=0)\n        \n        return x\n        \n    def __generate_data(self, index):\n        \"\"\"\n        Divides the spectrograms into 4 images and concatenates with the 4-channels eeg spectrogram\n        \n        returns:\n            X : (128, 256, 8) np.ndarray\n            y : (6,) np.ndarray\n        \"\"\"\n        X = np.zeros((128, 256, 8), dtype='float32')\n        y = np.zeros(6, dtype='float32')\n        img = np.ones((128,256), dtype='float32')\n        \n        row = self.df.iloc[index]\n        \n        if self.mode=='test': \n            r = 0\n        else: \n            r = int((row['min'] + row['max']) // 4)\n        \n        for region in range(4):\n            img = self.spectrograms[row.spectrogram_id][r:r+300, region*100:(region+1)*100].T\n            \n            # Log transform spectrogram\n            img = np.clip(img, np.exp(-4), np.exp(8))\n            img = np.log(img)\n\n            # Standarize per image\n            ep = 1e-6\n            mu = np.nanmean(img.flatten())\n            std = np.nanstd(img.flatten())\n            img = (img-mu)/(std+ep)\n            img = np.nan_to_num(img, nan=0.0)\n            \n            # Assign the image to the input vector X\n            # X : (128-28=100, 256, 8) and image (100, 300-44=256). The slicing is for the shapes to match.\n            X[14:-14, :, region] = img[:, 22:-22] / 2.0\n            \n            # Assigns the eeg\n            img = self.eeg_spectrograms[row.eeg_id]\n            X[:, :, 4:] = img\n                \n            if self.mode != 'test':\n                y = row[label_columns].values.astype(np.float32)\n            \n        X = np.transpose(X, (2, 0, 1))\n            \n        return X, y\n            \n        \n    \n    def __transform(self, img):\n        transforms = A.Compose([\n            A.HorizontalFlip(p=0.5)\n        ])\n        return transforms(image=img)['image']","metadata":{"papermill":{"duration":0.05412,"end_time":"2024-03-11T14:49:20.466655","exception":false,"start_time":"2024-03-11T14:49:20.412535","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training functions","metadata":{"papermill":{"duration":0.029903,"end_time":"2024-03-11T14:49:20.526708","exception":false,"start_time":"2024-03-11T14:49:20.496805","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def train_step(model, train_dl, loss_fn, optimizer, device, softmax=True):\n    model.train()\n\n    train_loss = AverageMeter()\n    for sample, target in train_dl:\n        sample = sample.to(device)\n        target = target.to(device)\n        \n        output = model(sample)\n        \n        if softmax:\n            output = F.log_softmax(output, dim=1)\n            \n        optimizer.zero_grad()\n        loss = loss_fn(output, target)\n        \n        train_loss.update(loss.item(), len(sample))\n        \n        loss.backward()\n        optimizer.step()\n    return train_loss.avg\n\ndef validation_step(model, val_dl, loss_fn, device,softmax=True):\n    model.eval()\n    val_loss=AverageMeter()\n\n    with torch.no_grad():\n        for sample, target in val_dl:\n            sample = sample.to(device)\n            target = target.to(device)\n\n            output = model(sample)\n            \n            if softmax:\n                output = F.log_softmax(output, dim=1)\n            \n            loss = loss_fn(output, target)\n            val_loss.update(loss.item(), len(sample))\n    return val_loss.avg","metadata":{"papermill":{"duration":0.041266,"end_time":"2024-03-11T14:49:20.598096","exception":false,"start_time":"2024-03-11T14:49:20.556830","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_model(model, train_dl, valid_dl, optimizer, loss_fn, device, mode):\n    history = {'train_losses':[], 'val_losses':[]}\n    best_loss = np.inf\n\n    start_t = time.time()\n    for epoch in range(config.EPOCHS):\n        train_loss = train_step(model, train_dl, loss_fn, optimizer, device)\n        end_t = time.time()\n\n        val_loss = validation_step(model, valid_dl, loss_fn, device)\n\n        history['train_losses'].append(train_loss)\n        history['val_losses'].append(val_loss)\n        \n        # Save the model with the best loss on validation\n        if val_loss < best_loss:\n            best_loss = val_loss\n            torch.save(model.state_dict(), paths.OUTPUT_DIR + f\"sub23_EfficientNet_v2_s_best_{mode}.pth\")\n\n        print(f\"Epoch [{epoch+1}/{config.EPOCHS}] | Time: {(end_t - start_t) // 60}min\")\n        sep()\n        print(f\"Train loss: {round(train_loss, 6):<7}\")\n        print(f\"Valid. loss: {round(val_loss, 6):<7}\")\n        print(f\"Best loss: {round(best_loss, 6):<7}\\n\")\n    \n    torch.cuda.empty_cache()\n    gc.collect()\n    \n    return history","metadata":{"papermill":{"duration":0.041205,"end_time":"2024-03-11T14:49:20.669340","exception":false,"start_time":"2024-03-11T14:49:20.628135","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model Paths\nthis version fine-tunes the previous .47 version (sub23), mimicking the behaviour of a step-scheduler: it had run for 10 epochs on 1e-4 now its running 10 epochs on 1e-6, in order to avoid major changes o the weights.","metadata":{}},{"cell_type":"code","source":"models_sufix = ['both.pth', 'eeg.pth', 'kag.pth']\nmodel_paths = [paths.MODEL_PATH_PREFIX + suf for suf in models_sufix]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training with both espectrograms","metadata":{"papermill":{"duration":0.030506,"end_time":"2024-03-11T14:49:20.729932","exception":false,"start_time":"2024-03-11T14:49:20.699426","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"## Dataloader and Dataset","metadata":{"papermill":{"duration":0.030429,"end_time":"2024-03-11T14:49:20.792002","exception":false,"start_time":"2024-03-11T14:49:20.761573","status":"completed"},"tags":[]}},{"cell_type":"code","source":"train_dataset_both = HMSDataset(train_df, all_spectrograms, all_eegs, mode=\"train\")\n\ntrain_loader_both = data.DataLoader(\n    train_dataset_both,\n    batch_size=config.BATCH_SIZE_TRAIN,\n    shuffle=False,\n    num_workers=config.NUM_WORKERS, pin_memory=True, drop_last=True)\n\nval_dataset_both = HMSDataset(val_df, all_spectrograms, all_eegs, mode='train')\n\nval_loader_both = data.DataLoader(\n    train_dataset_both,\n    batch_size=config.BATCH_SIZE_VALID,\n    shuffle=False,\n    num_workers=config.NUM_WORKERS, pin_memory=True, drop_last=False)","metadata":{"papermill":{"duration":0.039886,"end_time":"2024-03-11T14:49:20.862333","exception":false,"start_time":"2024-03-11T14:49:20.822447","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img, labels = next(iter(train_loader_both))\nprint(img.shape, labels.shape)","metadata":{"papermill":{"duration":1.247516,"end_time":"2024-03-11T14:49:22.140067","exception":false,"start_time":"2024-03-11T14:49:20.892551","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model","metadata":{"papermill":{"duration":0.030334,"end_time":"2024-03-11T14:49:22.202141","exception":false,"start_time":"2024-03-11T14:49:22.171807","status":"completed"},"tags":[]}},{"cell_type":"code","source":"%%capture\n# Listing pytorch pre-trained models\nmodel = models.efficientnet_v2_s()\n\n# Altering the last linear ayer of the net -> match the number of classes from our problem\nmodel.classifier[1] = nn.Linear(\n    in_features=model.classifier[1].in_features,\n    out_features=labels[0].shape[0]\n)\n\nmodel.to(device)\nmodel.load_state_dict(torch.load(model_paths[0]))","metadata":{"papermill":{"duration":2.051386,"end_time":"2024-03-11T14:49:24.284094","exception":false,"start_time":"2024-03-11T14:49:22.232708","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Optimizer and Loss function","metadata":{"papermill":{"duration":0.030047,"end_time":"2024-03-11T14:49:24.345394","exception":false,"start_time":"2024-03-11T14:49:24.315347","status":"completed"},"tags":[]}},{"cell_type":"code","source":"optimizer = optim.Adam(model.parameters(), lr=config.LR, weight_decay=config.WEIGHT_DECAY)\nloss_fn = nn.KLDivLoss(reduction='batchmean')","metadata":{"papermill":{"duration":0.041518,"end_time":"2024-03-11T14:49:24.417169","exception":false,"start_time":"2024-03-11T14:49:24.375651","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training loop","metadata":{"papermill":{"duration":0.030929,"end_time":"2024-03-11T14:49:24.478262","exception":false,"start_time":"2024-03-11T14:49:24.447333","status":"completed"},"tags":[]}},{"cell_type":"code","source":"history = train_model(model, train_loader_both, val_loader_both, optimizer, loss_fn, config.DEVICE, mode=\"both\")","metadata":{"papermill":{"duration":5028.889661,"end_time":"2024-03-11T16:13:13.398844","exception":false,"start_time":"2024-03-11T14:49:24.509183","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig = plt.figure()\nplt.title(\"Loss function\")\nplt.plot(history[\"train_losses\"], c=\"g\", label=\"train\")\nplt.plot(history[\"val_losses\"], c=\"r\", label=\"valid\")\nplt.legend()","metadata":{"papermill":{"duration":0.44107,"end_time":"2024-03-11T16:13:13.871203","exception":false,"start_time":"2024-03-11T16:13:13.430133","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training with eeg-spec only","metadata":{"papermill":{"duration":0.031744,"end_time":"2024-03-11T16:13:13.935360","exception":false,"start_time":"2024-03-11T16:13:13.903616","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"## Dataset and Dataloader","metadata":{"papermill":{"duration":0.031534,"end_time":"2024-03-11T16:13:13.998552","exception":false,"start_time":"2024-03-11T16:13:13.967018","status":"completed"},"tags":[]}},{"cell_type":"code","source":"train_dataset_eeg = HMSDataset(train_df, all_spectrograms, all_eegs, use_kaggle_spec=False, mode=\"train\")\n\ntrain_loader_eeg = data.DataLoader(\n    train_dataset_eeg,\n    batch_size=config.BATCH_SIZE_TRAIN,\n    shuffle=False,\n    num_workers=config.NUM_WORKERS, pin_memory=True, drop_last=True)\n\nval_dataset_eeg = HMSDataset(val_df, all_spectrograms, all_eegs, use_kaggle_spec=False, mode='train')\n\nval_loader_eeg = data.DataLoader(\n    train_dataset_eeg,\n    batch_size=config.BATCH_SIZE_VALID,\n    shuffle=False,\n    num_workers=config.NUM_WORKERS, pin_memory=True, drop_last=False)","metadata":{"papermill":{"duration":0.042576,"end_time":"2024-03-11T16:13:14.073028","exception":false,"start_time":"2024-03-11T16:13:14.030452","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img, labels = next(iter(train_loader_eeg))\nprint(img.shape, labels.shape)","metadata":{"papermill":{"duration":0.994025,"end_time":"2024-03-11T16:13:15.099057","exception":false,"start_time":"2024-03-11T16:13:14.105032","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model","metadata":{"papermill":{"duration":0.032253,"end_time":"2024-03-11T16:13:15.164719","exception":false,"start_time":"2024-03-11T16:13:15.132466","status":"completed"},"tags":[]}},{"cell_type":"code","source":"%%capture\n# Listing pytorch pre-trained models\nmodel = models.efficientnet_v2_s()\n\n# Altering the last linear ayer of the net -> match the number of classes from our problem\nmodel.classifier[1] = nn.Linear(\n    in_features=model.classifier[1].in_features,\n    out_features=labels[0].shape[0]\n)\n\nmodel.to(device)\nmodel.load_state_dict(torch.load(model_paths[1]))","metadata":{"papermill":{"duration":0.821509,"end_time":"2024-03-11T16:13:16.018638","exception":false,"start_time":"2024-03-11T16:13:15.197129","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Optimizer and loss","metadata":{"papermill":{"duration":0.0371,"end_time":"2024-03-11T16:13:16.099325","exception":false,"start_time":"2024-03-11T16:13:16.062225","status":"completed"},"tags":[]}},{"cell_type":"code","source":"optimizer = optim.Adam(model.parameters(), lr=config.LR, weight_decay=config.WEIGHT_DECAY)\nloss_fn = nn.KLDivLoss(reduction='batchmean')","metadata":{"papermill":{"duration":0.053606,"end_time":"2024-03-11T16:13:16.185369","exception":false,"start_time":"2024-03-11T16:13:16.131763","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training loop","metadata":{"papermill":{"duration":0.033428,"end_time":"2024-03-11T16:13:16.251156","exception":false,"start_time":"2024-03-11T16:13:16.217728","status":"completed"},"tags":[]}},{"cell_type":"code","source":"history = train_model(model, train_loader_eeg, val_loader_eeg, optimizer, loss_fn, config.DEVICE, mode=\"eeg\")","metadata":{"papermill":{"duration":2974.536092,"end_time":"2024-03-11T17:02:50.819800","exception":false,"start_time":"2024-03-11T16:13:16.283708","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig = plt.figure()\nplt.title(\"Loss function\")\nplt.plot(history[\"train_losses\"], c=\"g\", label=\"train\")\nplt.plot(history[\"val_losses\"], c=\"r\", label=\"valid\")\nplt.legend()","metadata":{"papermill":{"duration":0.437791,"end_time":"2024-03-11T17:02:51.290449","exception":false,"start_time":"2024-03-11T17:02:50.852658","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training with kaggle-spec only","metadata":{"papermill":{"duration":0.033351,"end_time":"2024-03-11T17:02:51.358110","exception":false,"start_time":"2024-03-11T17:02:51.324759","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"## Dataset and Dataloader","metadata":{"papermill":{"duration":0.034006,"end_time":"2024-03-11T17:02:51.425796","exception":false,"start_time":"2024-03-11T17:02:51.391790","status":"completed"},"tags":[]}},{"cell_type":"code","source":"train_dataset_kag = HMSDataset(train_df, all_spectrograms, all_eegs, use_eeg_spec=False, mode=\"train\")\n\ntrain_loader_kag = data.DataLoader(\n    train_dataset_kag,\n    batch_size=config.BATCH_SIZE_TRAIN,\n    shuffle=False,\n    num_workers=config.NUM_WORKERS, pin_memory=True, drop_last=True)\n\nval_dataset_kag = HMSDataset(val_df, all_spectrograms, all_eegs, use_eeg_spec=False, mode='train')\n\nval_loader_kag = data.DataLoader(\n    train_dataset_kag,\n    batch_size=config.BATCH_SIZE_VALID,\n    shuffle=False,\n    num_workers=config.NUM_WORKERS, pin_memory=True, drop_last=False)","metadata":{"papermill":{"duration":0.043514,"end_time":"2024-03-11T17:02:51.502925","exception":false,"start_time":"2024-03-11T17:02:51.459411","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img, labels = next(iter(train_loader_kag))\nprint(img.shape, labels.shape)","metadata":{"papermill":{"duration":1.015604,"end_time":"2024-03-11T17:02:52.552468","exception":false,"start_time":"2024-03-11T17:02:51.536864","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model","metadata":{"papermill":{"duration":0.033995,"end_time":"2024-03-11T17:02:52.620998","exception":false,"start_time":"2024-03-11T17:02:52.587003","status":"completed"},"tags":[]}},{"cell_type":"code","source":"%%capture\n# Listing pytorch pre-trained models\nmodel = models.efficientnet_v2_s()\n\n# Altering the last linear ayer of the net -> match the number of classes from our problem\nmodel.classifier[1] = nn.Linear(\n    in_features=model.classifier[1].in_features,\n    out_features=labels[0].shape[0]\n)\n\nmodel.to(device)\nmodel.load_state_dict(torch.load(model_paths[2]))","metadata":{"papermill":{"duration":0.743654,"end_time":"2024-03-11T17:02:53.398681","exception":false,"start_time":"2024-03-11T17:02:52.655027","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Optimizer and loss","metadata":{"papermill":{"duration":0.033505,"end_time":"2024-03-11T17:02:53.466479","exception":false,"start_time":"2024-03-11T17:02:53.432974","status":"completed"},"tags":[]}},{"cell_type":"code","source":"optimizer = optim.Adam(model.parameters(), lr=config.LR, weight_decay=config.WEIGHT_DECAY)\nloss_fn = nn.KLDivLoss(reduction='batchmean')","metadata":{"papermill":{"duration":0.05513,"end_time":"2024-03-11T17:02:53.555752","exception":false,"start_time":"2024-03-11T17:02:53.500622","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training loop","metadata":{"papermill":{"duration":0.03944,"end_time":"2024-03-11T17:02:53.633010","exception":false,"start_time":"2024-03-11T17:02:53.593570","status":"completed"},"tags":[]}},{"cell_type":"code","source":"history = train_model(model, train_loader_kag, val_loader_kag, optimizer, loss_fn, config.DEVICE, mode=\"kag\")","metadata":{"papermill":{"duration":2991.629175,"end_time":"2024-03-11T17:52:45.301860","exception":false,"start_time":"2024-03-11T17:02:53.672685","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig = plt.figure()\nplt.title(\"Loss function\")\nplt.plot(history[\"train_losses\"], c=\"g\", label=\"train\")\nplt.plot(history[\"val_losses\"], c=\"r\", label=\"valid\")\nplt.legend()","metadata":{"papermill":{"duration":0.498732,"end_time":"2024-03-11T17:52:45.836742","exception":false,"start_time":"2024-03-11T17:52:45.338010","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]}]}