{"metadata":{"kaggle":{"accelerator":"none","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"}],"dockerImageVersionId":30648,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false},"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"papermill":{"default_parameters":{},"duration":22321.864322,"end_time":"2024-01-28T12:56:49.320132","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2024-01-28T06:44:47.45581","version":"2.4.0"},"widgets":{"application/vnd.jupyter.widget-state+json":{"state":{"020b844a5f1c4a998793589345d58187":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"387ef6c55cf34ee8b06cef960ab5709e":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"DescriptionStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"DescriptionStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"StyleView","description_width":""}},"4551e56c177e48c7b32c50244181acce":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"HTMLModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"HTMLModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"HTMLView","description":"","description_tooltip":null,"layout":"IPY_MODEL_020b844a5f1c4a998793589345d58187","placeholder":"​","style":"IPY_MODEL_a76c83a1114641e9a0278861a509523d","value":" 21.4M/21.4M [00:00&lt;00:00, 39.7MB/s]"}},"54c978d8e366477fa1f2e4de7d812eaa":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"585b9ad521bc42f88613a33772d14b09":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"HTMLModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"HTMLModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"HTMLView","description":"","description_tooltip":null,"layout":"IPY_MODEL_54c978d8e366477fa1f2e4de7d812eaa","placeholder":"​","style":"IPY_MODEL_387ef6c55cf34ee8b06cef960ab5709e","value":"model.safetensors: 100%"}},"5c5e6c736e884230a714ede1aac584e8":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"5d247d0364d444f4bda51aea5d8e9103":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"HBoxModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"HBoxModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"HBoxView","box_style":"","children":["IPY_MODEL_585b9ad521bc42f88613a33772d14b09","IPY_MODEL_f84cec30df65438192b9c608c0d1d8b4","IPY_MODEL_4551e56c177e48c7b32c50244181acce"],"layout":"IPY_MODEL_5c5e6c736e884230a714ede1aac584e8"}},"a76c83a1114641e9a0278861a509523d":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"DescriptionStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"DescriptionStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"StyleView","description_width":""}},"f34a838e5b6b4dcc82ea12ad3977b520":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"ProgressStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"ProgressStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"StyleView","bar_color":null,"description_width":""}},"f84cec30df65438192b9c608c0d1d8b4":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"FloatProgressModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"FloatProgressModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"ProgressView","bar_style":"success","description":"","description_tooltip":null,"layout":"IPY_MODEL_fe2e5aa594d24e889081a556aebdec4f","max":21355344,"min":0,"orientation":"horizontal","style":"IPY_MODEL_f34a838e5b6b4dcc82ea12ad3977b520","value":21355344}},"fe2e5aa594d24e889081a556aebdec4f":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}}},"version_major":2,"version_minor":0}}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Introduction\n## Acknowledgements\nThe original base of this notebook was copied from @andreasbis. We thank them for supplying a useful baseline to expand upon. Please take a look at their work: https://www.kaggle.com/code/andreasbis/hms-train-efficientnetb1.","metadata":{}},{"cell_type":"markdown","source":"# Imports","metadata":{}},{"cell_type":"code","source":"import gc\nimport os\nimport random\nimport warnings\nfrom IPython.display import display\n\nimport numpy as np\nimport pandas as pd\n\nimport timm\nimport torch\nimport optuna\nimport torch.nn as nn  \nimport torch.optim as optim\nimport torch.nn.functional as F\nimport torchvision.transforms as transforms\nfrom torch.optim.lr_scheduler import CosineAnnealingLR, ReduceLROnPlateau\n\nwarnings.filterwarnings('ignore', category=Warning)\ngc.collect()","metadata":{"papermill":{"duration":5.810106,"end_time":"2024-01-28T06:44:56.777717","exception":false,"start_time":"2024-01-28T06:44:50.967611","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-07T13:31:00.068035Z","iopub.execute_input":"2024-03-07T13:31:00.068544Z","iopub.status.idle":"2024-03-07T13:31:14.848431Z","shell.execute_reply.started":"2024-03-07T13:31:00.068499Z","shell.execute_reply":"2024-03-07T13:31:14.846907Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Setup","metadata":{"papermill":{"duration":0.003167,"end_time":"2024-01-28T06:44:56.784246","exception":false,"start_time":"2024-01-28T06:44:56.781079","status":"completed"},"tags":[]}},{"cell_type":"code","source":"labels = ['seizure', 'lpd', 'gpd', 'lrda', 'grda', 'other']\n\nclass Config:\n    seed = 3131\n    image_transform = transforms.Resize((512,512))\n    batch_size = 16\n    num_epochs = 9\n    num_folds = 5\n    num_trials = 20\n    dataset_wide_mean = 0\n    dataset_wide_std = 0\n    manual_pruning_threshold = 0.57\n    optimize_hyperparameters = False\n    \nclass HyperparameterSpaces:\n    lowpass = {\n        \"min\": np.exp(10),\n        \"max\": np.exp(10)\n    }\n    highpass = {\n        \"min\": np.exp(-6),\n        \"max\": np.exp(-6)\n    }\n    learning_rate = {\n        \"min\": 0.0005,\n        \"max\": 0.0015\n    }\n    dropout = {\n        \"min\": 0.17,\n        \"max\": 0.23\n    }\n\n    schedulers = [\"CosineAnnealingLR\", \"ReduceLROnPlateau\"]\n    normalize_dataset_wide = [True, False]\n\nclass HyperparameterPreset:\n    lowpass = np.exp(10)\n    highpass = np.exp(-6)\n    learning_rate = 0.00137263241151172\n    dropout = 0.184235721122803\n    scheduler = \"CosineAnnealingLR\"\n    normalize_dataset_wide = True\n    \ndef set_seed(seed):\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = True\n    \n    torch.manual_seed(seed)\n    np.random.seed(seed)\n    random.seed(seed)\n\ndef kl_loss(p, q):\n    epsilon = 10 ** (-15)\n    \n    p = torch.clamp(p, epsilon, 1 - epsilon)\n    log_p = torch.log(p)\n    log_q = nn.functional.log_softmax(q, dim=1)\n    \n    kl_divergence_per_point = p * (log_p - log_q)\n    kl_divergence_per_label = torch.sum(kl_divergence_per_point, dim=1)\n    \n    return torch.mean(kl_divergence_per_label)\n\nset_seed(Config.seed)\ngc.collect()","metadata":{"papermill":{"duration":0.139048,"end_time":"2024-01-28T06:44:56.926359","exception":false,"start_time":"2024-01-28T06:44:56.787311","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-07T13:31:19.417456Z","iopub.execute_input":"2024-03-07T13:31:19.417913Z","iopub.status.idle":"2024-03-07T13:31:19.793620Z","shell.execute_reply.started":"2024-03-07T13:31:19.417879Z","shell.execute_reply":"2024-03-07T13:31:19.792263Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Loading","metadata":{"papermill":{"duration":0.003288,"end_time":"2024-01-28T06:44:56.933246","exception":false,"start_time":"2024-01-28T06:44:56.929958","status":"completed"},"tags":[]}},{"cell_type":"code","source":"train_df = pd.read_csv(\"/kaggle/input/hms-harmful-brain-activity-classification/train.csv\")\n\ndef extract_vote_count_features(input_data: pd.DataFrame) -> pd.DataFrame:\n    label_votes = pd.DataFrame()\n    \n    for label in labels:\n        input_grouped_by_spectrogram_id = input_data[f'{label}_vote'].groupby(input_data['spectrogram_id']).sum()\n\n        label_vote_sum = pd.DataFrame()\n        label_vote_sum[\"spectrogram_id\"] = input_grouped_by_spectrogram_id.index\n        label_vote_sum[f\"{label}_vote_sum\"] = input_grouped_by_spectrogram_id.values\n\n        if label == labels[0]:\n            label_votes = label_vote_sum\n        else:\n            label_votes = label_votes.merge(label_vote_sum, on='spectrogram_id', how='left')\n            \n    return label_votes\n\ndef extract_features(input_data: pd.DataFrame) -> pd.DataFrame:\n    choose_cols = ['spectrogram_id']\n    feature_df = extract_vote_count_features(input_data)\n    \n    feature_df['total_vote'] = 0\n    for label in labels:\n        choose_cols += [f'{label}_vote']\n        feature_df['total_vote'] += feature_df[f'{label}_vote_sum']\n        \n    for label in labels:\n        feature_df[f'{label}_vote'] = feature_df[f'{label}_vote_sum'] / feature_df['total_vote']\n        \n    feature_df = feature_df[choose_cols]\n    feature_df['path'] = feature_df['spectrogram_id'].apply(lambda x: \"/kaggle/input/hms-harmful-brain-activity-classification/train_spectrograms/\" + str(x) + \".parquet\")\n    \n    return feature_df\n\ntrain_features = extract_features(train_df)\ndisplay(train_features)\n    \ngc.collect()","metadata":{"papermill":{"duration":0.468753,"end_time":"2024-01-28T06:44:57.40548","exception":false,"start_time":"2024-01-28T06:44:56.936727","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-07T13:31:23.555839Z","iopub.execute_input":"2024-03-07T13:31:23.556679Z","iopub.status.idle":"2024-03-07T13:31:24.396275Z","shell.execute_reply.started":"2024-03-07T13:31:23.556639Z","shell.execute_reply":"2024-03-07T13:31:24.394597Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Preprocessing","metadata":{"papermill":{"duration":0.003848,"end_time":"2024-01-28T06:44:57.41335","exception":false,"start_time":"2024-01-28T06:44:57.409502","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def preprocess(path_to_parquet, lowpass, highpass):\n    data = pd.read_parquet(path_to_parquet)\n    data = data.fillna(-1).values[:, 1:].T\n    data = np.clip(data, highpass, lowpass)\n    data = np.log(data)\n    \n    return data\n\ndef get_dataset_wide_mean(paths, lowpass, highpass):\n    data_sum = 0\n    num_values = 0\n\n    for path in paths:      \n        data_point = preprocess(path[0], lowpass, highpass)\n        data_sum += data_point.sum(axis=(0, 1))\n        rows, columns = data_point.shape\n        num_values += rows * columns\n    \n    return data_sum / num_values\n\ndef get_dataset_wide_std(paths, lowpass, highpass):\n    sum_of_stds = 0\n    num_values = 0\n    \n    for path in paths:\n        data_point = preprocess(path[0], lowpass, highpass)\n        sum_of_stds += np.sum((data_point - Config.dataset_wide_mean) ** 2)\n        rows, columns = data_point.shape\n        num_values += rows * columns\n    \n    return np.sqrt(sum_of_stds / (num_values - 1))\n\ndef normalize_dataset_wide(data_point):\n    eps = 1e-6\n\n    data_point = (data_point - Config.dataset_wide_mean) / (Config.dataset_wide_std + eps)\n\n    data_tensor = torch.unsqueeze(torch.Tensor(data_point), dim=0)\n    data_point = Config.image_transform(data_tensor)\n\n    return data_point\n\ndef normalize_instance_wise(data_point):\n    eps = 1e-6\n    \n    data_mean = data_point.mean(axis=(0, 1))\n    data_std = data_point.std(axis=(0, 1))\n    data_point = (data_point - data_mean) / (data_std + eps)\n    \n    data_tensor = torch.unsqueeze(torch.Tensor(data_point), dim=0)\n    data_point = Config.image_transform(data_tensor)\n    \n    return data_point\n\ndef get_batch(paths, lowpass, highpass, normalization_dataset_wide):        \n    batch_data = []\n    \n    for path in paths:\n        data_point = preprocess(path[0], lowpass, highpass)\n        \n        if normalization_dataset_wide:\n            data_point = normalize_dataset_wide(data_point)\n        else:\n            data_point = normalize_instance_wise(data_point)\n        \n        batch_data.append(data_point)\n    batch_data = torch.stack(batch_data)\n\n    return batch_data\n","metadata":{"papermill":{"duration":0.015617,"end_time":"2024-01-28T06:44:57.432829","exception":false,"start_time":"2024-01-28T06:44:57.417212","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-07T13:31:24.797136Z","iopub.execute_input":"2024-03-07T13:31:24.797734Z","iopub.status.idle":"2024-03-07T13:31:24.815718Z","shell.execute_reply.started":"2024-03-07T13:31:24.797690Z","shell.execute_reply":"2024-03-07T13:31:24.814320Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"Calculating dataset-wide mean...\")\nConfig.dataset_wide_mean = get_dataset_wide_mean(train_features[[\"path\"]].values, HyperparameterPreset.lowpass, HyperparameterPreset.highpass)\nprint(\"Finished mean calculation!\\nCalculating dataset-wide standard deviation...\")\nConfig.dataset_wide_std = get_dataset_wide_std(train_features[[\"path\"]].values, HyperparameterPreset.lowpass, HyperparameterPreset.highpass)\nprint(\"Finished standard deviation calculation!\")\n\ngc.collect()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Objective Function","metadata":{"papermill":{"duration":0.003763,"end_time":"2024-01-28T06:44:57.440662","exception":false,"start_time":"2024-01-28T06:44:57.436899","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def get_fold_train_val_indexes(indexes: np.ndarray, fold: int) -> tuple[np.ndarray, np.ndarray]:\n    lower_bound = fold * len(indexes) // Config.num_folds\n    upper_bound = (fold + 1) * len(indexes) // Config.num_folds\n    \n    val_idx = indexes[lower_bound:upper_bound]\n    train_idx = []\n    \n    for index in indexes:\n        if index not in val_idx:\n            train_idx.append(index)\n            \n    train_idx = np.array(train_idx)\n    \n    return (train_idx, val_idx) \n\ndef objective(trial) -> float:    \n    if trial is None:\n        lowpass = HyperparameterPreset.lowpass\n        highpass = HyperparameterPreset.highpass\n        learning_rate = HyperparameterPreset.learning_rate\n        dropout = HyperparameterPreset.dropout\n        scheduler_name = HyperparameterPreset.scheduler\n        normalize_dataset_wide = HyperparameterPreset.normalize_dataset_wide\n    else:\n        lowpass = trial.suggest_float(\"lowpass\", HyperparameterSpaces.lowpass[\"min\"], HyperparameterSpaces.lowpass[\"max\"])\n        highpass = trial.suggest_float(\"highpass\", HyperparameterSpaces.highpass[\"min\"], HyperparameterSpaces.highpass[\"max\"])\n        learning_rate = trial.suggest_float(\"learning_rate\", HyperparameterSpaces.learning_rate[\"min\"], HyperparameterSpaces.learning_rate[\"max\"])\n        dropout = trial.suggest_float(\"dropout\", HyperparameterSpaces.dropout[\"min\"], HyperparameterSpaces.dropout[\"max\"])\n        scheduler_name = trial.suggest_categorical(\"scheduler\", HyperparameterSpaces.schedulers)\n        normalize_dataset_wide = trial.suggest_categorical(\"normalize_dataset_wide\", HyperparameterSpaces.normalize_dataset_wide)\n    \n    for fold in range(Config.num_folds):        \n        train_idx, val_idx = get_fold_train_val_indexes(train_spectrogram_indexes, fold)\n\n        model = timm.create_model(\n            'efficientnet_b1', \n            pretrained=True, \n            num_classes=6, \n            in_chans=1, \n            drop_rate=dropout\n        ).to(device)\n\n        optimizer = optim.AdamW(\n            model.parameters(), \n            lr=learning_rate, \n            betas=(0.5, 0.999),\n            weight_decay=0.01\n        )\n            \n        if scheduler_name == \"CosineAnnealingLR\":\n            scheduler = CosineAnnealingLR(optimizer, T_max=Config.num_epochs)\n        elif scheduler_name == \"ReduceLROnPlateau\":\n            scheduler = ReduceLROnPlateau(optimizer)\n        else:\n            raise ValueError()\n\n        best_val_loss = float('inf')\n        train_losses = []\n        val_losses = []\n\n        print(f\"Starting training for fold {fold + 1}\")\n\n        for epoch in range(Config.num_epochs):\n            print(f\" Epoch: {epoch + 1}\")\n            model.train()\n            train_loss = []\n\n            random_num = np.arange(len(train_idx))\n            np.random.shuffle(random_num)\n            train_idx = train_idx[random_num]\n\n            print(f\"  Train - {len(train_idx)} indexes\")\n            for idx in range(0, len(train_idx), Config.batch_size):\n                optimizer.zero_grad()\n\n                train_batch_idx = train_idx[idx:idx + Config.batch_size]\n                train_batch_idx_paths = train_features[['path']].iloc[train_batch_idx].values\n                train_batch = get_batch(train_batch_idx_paths, lowpass, highpass, normalize_dataset_wide)\n                train_batch = train_batch.to(device)\n\n                train_target = train_features[['seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']].iloc[train_batch_idx].values\n                train_target = torch.Tensor(train_target).to(device)\n\n                train_pred = model(train_batch)\n\n                loss = kl_loss(train_target, train_pred)\n                loss.backward()\n                optimizer.step()\n\n                train_loss.append(loss.item())\n\n            epoch_train_loss = np.mean(train_loss)\n            train_losses.append(epoch_train_loss)\n            print(f\" Epoch {epoch + 1}: Train Loss = {epoch_train_loss:.2f}\")\n\n            if scheduler_name == \"CosineAnnealingLR\":\n                scheduler.step()\n\n            model.eval()\n            val_loss = []\n\n            with torch.no_grad():\n                print(f\"  Validation - {len(val_idx)} indexes\")\n                for idx in range(0, len(val_idx), Config.batch_size):\n                    val_batch_idx = val_idx[idx:idx + Config.batch_size]\n                    val_batch_idx_paths = train_features[['path']].iloc[val_batch_idx].values\n                    val_batch = get_batch(val_batch_idx_paths, lowpass, highpass, normalize_dataset_wide)\n                    val_batch = val_batch.to(device)\n\n                    val_target = train_features[['seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']].iloc[val_batch_idx].values\n                    val_target = torch.Tensor(val_target).to(device)\n\n                    val_pred = model(val_batch)\n\n                    loss = kl_loss(val_target, val_pred)\n                    val_loss.append(loss.item())\n\n            epoch_val_loss = np.mean(val_loss)\n            val_losses.append(epoch_val_loss)\n            print(f\" Epoch {epoch + 1}: Test Loss = {epoch_val_loss:.2f}\")\n            \n            if scheduler_name == \"ReduceLROnPlateau\":\n                    scheduler.step(epoch_val_loss)\n\n            if epoch_val_loss < best_val_loss:\n                best_val_loss = epoch_val_loss\n                torch.save(model.state_dict(), f\"efficientnet_b1_fold{fold}.pth\")\n\n            gc.collect()\n            \n            if trial is not None:\n                trial.report(epoch_val_loss, epoch)\n\n                if trial.should_prune():\n                    raise optuna.TrialPruned()\n        \n    print(f\"Fold {fold + 1} Best Test Loss: {best_val_loss:.2f}\")\n    \n    return best_val_loss\n\ntrain_spectrogram_indexes = np.arange(len(train_features))\nnp.random.shuffle(train_spectrogram_indexes)\n    \ngc.collect()","metadata":{"papermill":{"duration":22310.196541,"end_time":"2024-01-28T12:56:47.64122","exception":false,"start_time":"2024-01-28T06:44:57.444679","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-03-07T13:31:28.590283Z","iopub.execute_input":"2024-03-07T13:31:28.590793Z","iopub.status.idle":"2024-03-07T13:31:28.958877Z","shell.execute_reply.started":"2024-03-07T13:31:28.590754Z","shell.execute_reply":"2024-03-07T13:31:28.957498Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training/Optimization","metadata":{}},{"cell_type":"code","source":"device = 'cuda' if torch.cuda.is_available() else 'cpu'\nprint(f\"Using device: {device}\")\n\ngc.collect()\n\nif Config.optimize_hyperparameters:\n    Config.num_folds = 2\n    Config.num_epochs = 2\n    \n    print(\"### STARTING HYPERPARAMETER OPTIMIZATION ###\")\n    study = optuna.create_study(pruner=optuna.pruners.ThresholdPruner(upper=Config.manual_pruning_threshold))\n    study.optimize(objective, n_trials=Config.num_trials)\n\n    print(\"Best test loss:\", study.best_value)\n    print(\"Best trial run:\", study.best_trial)\n    print(\"Hyperparameter values for best test loss:\")\n    print(study.best_params)\n    print(\"### FINISHED HYPERPARAMETER OPTIMIZATION ###\")\nelse:\n    print(\"### STARTING MODEL TRAINING ###\")\n    objective(None)\n    print(\"### FINISHED MODEL TRAINING ###\")","metadata":{"execution":{"iopub.status.busy":"2024-03-07T13:31:33.764097Z","iopub.execute_input":"2024-03-07T13:31:33.764614Z"},"trusted":true},"execution_count":null,"outputs":[]}]}