{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"}],"dockerImageVersionId":30699,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"\n!pip install keras","metadata":{"execution":{"iopub.status.busy":"2024-05-11T06:55:59.626276Z","iopub.execute_input":"2024-05-11T06:55:59.626612Z","iopub.status.idle":"2024-05-11T06:56:12.703439Z","shell.execute_reply.started":"2024-05-11T06:55:59.626586Z","shell.execute_reply":"2024-05-11T06:56:12.702500Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nos.environ[\"KERAS_BACKEND\"] = \"jax\" # you can also use tensorflow or torch\n\nimport keras_cv\nimport keras\n# from keras import ops\nimport tensorflow as tf\nimport cv2\nimport pandas as pd\nimport numpy as np\nfrom glob import glob\nfrom tqdm.notebook import tqdm\nimport joblib\nimport matplotlib.pyplot as plt \nimport seaborn as sns\nfrom sklearn.preprocessing import LabelEncoder\nfrom sklearn.model_selection import train_test_split\nfrom tensorflow.keras.models import Sequential\nfrom tensorflow.keras.layers import Conv2D, MaxPooling2D,BatchNormalization, Flatten, Dense, Dropout,GlobalAveragePooling2D\nfrom sklearn.tree import DecisionTreeClassifier\nfrom sklearn.ensemble import RandomForestClassifier\nfrom sklearn.svm import SVC\nfrom sklearn.metrics import accuracy_score","metadata":{"execution":{"iopub.status.busy":"2024-05-11T06:57:34.787163Z","iopub.execute_input":"2024-05-11T06:57:34.787543Z","iopub.status.idle":"2024-05-11T06:57:53.109788Z","shell.execute_reply.started":"2024-05-11T06:57:34.787507Z","shell.execute_reply":"2024-05-11T06:57:53.108876Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"TensorFlow:\", tf.__version__)\nprint(\"Keras:\", keras.__version__)\nprint(\"KerasCV:\", keras_cv.__version__)","metadata":{"execution":{"iopub.status.busy":"2024-05-11T06:57:53.113501Z","iopub.execute_input":"2024-05-11T06:57:53.114075Z","iopub.status.idle":"2024-05-11T06:57:53.119387Z","shell.execute_reply.started":"2024-05-11T06:57:53.114048Z","shell.execute_reply":"2024-05-11T06:57:53.118614Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    verbose = 1  # Verbosity\n    seed = 42  # Random seed\n    preset = \"efficientnet\"  # Name of pretrained classifier\n    image_size = [400, 300]  # Input image size\n    epochs = 30 # Training epochs\n    batch_size = 64  # Batch size\n    lr_mode = \"cos\" # LR scheduler mode from one of \"cos\", \"step\", \"exp\"\n    drop_remainder = True  # Drop incomplete batches\n    num_classes = 6 # Number of classes in the dataset\n    fold = 0 # Which fold to set as validation data\n    class_names = ['Seizure', 'LPD', 'GPD', 'LRDA','GRDA', 'Other']\n    label2name = dict(enumerate(class_names))\n    name2label = {v:k for k, v in label2name.items()}","metadata":{"execution":{"iopub.status.busy":"2024-05-11T06:57:53.120877Z","iopub.execute_input":"2024-05-11T06:57:53.121475Z","iopub.status.idle":"2024-05-11T06:57:53.141221Z","shell.execute_reply.started":"2024-05-11T06:57:53.121433Z","shell.execute_reply":"2024-05-11T06:57:53.140143Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"keras.utils.set_random_seed(CFG.seed)","metadata":{"execution":{"iopub.status.busy":"2024-05-11T06:58:01.141899Z","iopub.execute_input":"2024-05-11T06:58:01.142488Z","iopub.status.idle":"2024-05-11T06:58:01.146910Z","shell.execute_reply.started":"2024-05-11T06:58:01.142458Z","shell.execute_reply":"2024-05-11T06:58:01.145971Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BASE_PATH = \"/kaggle/input/hms-harmful-brain-activity-classification\"\n\nSPEC_DIR = \"/tmp/dataset/hms-hbac\"\nos.makedirs(SPEC_DIR+'/train_spectrograms', exist_ok=True)\nos.makedirs(SPEC_DIR+'/test_spectrograms', exist_ok=True)","metadata":{"execution":{"iopub.status.busy":"2024-05-11T06:58:04.528950Z","iopub.execute_input":"2024-05-11T06:58:04.529338Z","iopub.status.idle":"2024-05-11T06:58:04.534694Z","shell.execute_reply.started":"2024-05-11T06:58:04.529312Z","shell.execute_reply":"2024-05-11T06:58:04.533770Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Train + Valid\ndf = pd.read_csv(f'{BASE_PATH}/train.csv')\ndf['eeg_path'] = f'{BASE_PATH}/train_eegs/'+df['eeg_id'].astype(str)+'.parquet'\ndf['spec_path'] = f'{BASE_PATH}/train_spectrograms/'+df['spectrogram_id'].astype(str)+'.parquet'\ndf['spec2_path'] = f'{SPEC_DIR}/train_spectrograms/'+df['spectrogram_id'].astype(str)+'.npy'\ndf['class_name'] = df.expert_consensus.copy()\ndf['class_label'] = df.expert_consensus.map(CFG.name2label)\ndisplay(df.head(2))\n\n# Test\ntest_df = pd.read_csv(f'{BASE_PATH}/test.csv')\ntest_df['eeg_path'] = f'{BASE_PATH}/test_eegs/'+test_df['eeg_id'].astype(str)+'.parquet'\ntest_df['spec_path'] = f'{BASE_PATH}/test_spectrograms/'+test_df['spectrogram_id'].astype(str)+'.parquet'\ntest_df['spec2_path'] = f'{SPEC_DIR}/test_spectrograms/'+test_df['spectrogram_id'].astype(str)+'.npy'\ndisplay(test_df.head(2))","metadata":{"execution":{"iopub.status.busy":"2024-05-11T06:58:08.420447Z","iopub.execute_input":"2024-05-11T06:58:08.421323Z","iopub.status.idle":"2024-05-11T06:58:08.958781Z","shell.execute_reply.started":"2024-05-11T06:58:08.421287Z","shell.execute_reply":"2024-05-11T06:58:08.957841Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.isnull().sum()","metadata":{"execution":{"iopub.status.busy":"2024-05-11T06:58:14.140879Z","iopub.execute_input":"2024-05-11T06:58:14.141237Z","iopub.status.idle":"2024-05-11T06:58:14.208592Z","shell.execute_reply.started":"2024-05-11T06:58:14.141208Z","shell.execute_reply":"2024-05-11T06:58:14.207744Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df","metadata":{"execution":{"iopub.status.busy":"2024-05-11T06:58:17.346362Z","iopub.execute_input":"2024-05-11T06:58:17.347020Z","iopub.status.idle":"2024-05-11T06:58:17.372337Z","shell.execute_reply.started":"2024-05-11T06:58:17.346985Z","shell.execute_reply":"2024-05-11T06:58:17.371414Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df","metadata":{"execution":{"iopub.status.busy":"2024-05-11T06:58:21.878681Z","iopub.execute_input":"2024-05-11T06:58:21.879919Z","iopub.status.idle":"2024-05-11T06:58:21.894758Z","shell.execute_reply.started":"2024-05-11T06:58:21.879716Z","shell.execute_reply":"2024-05-11T06:58:21.893593Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define a function to process a single eeg_id\n\ndef process_spec(spec_id, split=\"train\"):\n    spec_path = f\"{BASE_PATH}/{split}_spectrograms/{spec_id}.parquet\"\n    spec = pd.read_parquet(spec_path)\n    spec = spec.fillna(0).values[:, 1:].T # fill NaN values with 0, transpose for (Time, Freq) -> (Freq, Time)\n    spec = spec.astype(\"float32\")\n    np.save(f\"{SPEC_DIR}/{split}_spectrograms/{spec_id}.npy\", spec)\n\n# Get unique spec_ids of train and valid data\nspec_ids = df[\"spectrogram_id\"].unique()\n\n# Parallelize the processing using joblib for training data\n_ = joblib.Parallel(n_jobs=-1, backend=\"loky\")(\n    joblib.delayed(process_spec)(spec_id, \"train\")\n    for spec_id in tqdm(spec_ids,total=len(spec_ids))\n)\n\n# Get unique spec_ids of test data\ntest_spec_ids = test_df[\"spectrogram_id\"].unique()\n\n# Parallelize the processing using joblib for test data\n_ = joblib.Parallel(n_jobs=-1, backend=\"loky\")(\n    joblib.delayed(process_spec)(spec_id, \"test\")\n    for spec_id in tqdm(test_spec_ids, total=len(test_spec_ids))\n)","metadata":{"execution":{"iopub.status.busy":"2024-05-11T06:58:25.102852Z","iopub.execute_input":"2024-05-11T06:58:25.103205Z","iopub.status.idle":"2024-05-11T07:01:45.576351Z","shell.execute_reply.started":"2024-05-11T06:58:25.103177Z","shell.execute_reply":"2024-05-11T07:01:45.575463Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Remove all the warnings\nimport warnings\nwarnings.filterwarnings('ignore')","metadata":{"execution":{"iopub.status.busy":"2024-05-11T07:01:50.228374Z","iopub.execute_input":"2024-05-11T07:01:50.229163Z","iopub.status.idle":"2024-05-11T07:01:50.233978Z","shell.execute_reply.started":"2024-05-11T07:01:50.229126Z","shell.execute_reply":"2024-05-11T07:01:50.232624Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Shape of the dataset\ndf.shape","metadata":{"execution":{"iopub.status.busy":"2024-05-11T07:01:53.040426Z","iopub.execute_input":"2024-05-11T07:01:53.041027Z","iopub.status.idle":"2024-05-11T07:01:53.048222Z","shell.execute_reply.started":"2024-05-11T07:01:53.040994Z","shell.execute_reply":"2024-05-11T07:01:53.045965Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.describe()","metadata":{"execution":{"iopub.status.busy":"2024-05-11T07:01:55.953473Z","iopub.execute_input":"2024-05-11T07:01:55.954307Z","iopub.status.idle":"2024-05-11T07:01:56.045783Z","shell.execute_reply.started":"2024-05-11T07:01:55.954274Z","shell.execute_reply":"2024-05-11T07:01:56.044925Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Check unique values of `expert_consensus`\ndf['expert_consensus'].unique()","metadata":{"execution":{"iopub.status.busy":"2024-05-11T07:01:58.983069Z","iopub.execute_input":"2024-05-11T07:01:58.983869Z","iopub.status.idle":"2024-05-11T07:01:58.999252Z","shell.execute_reply.started":"2024-05-11T07:01:58.983834Z","shell.execute_reply":"2024-05-11T07:01:58.998395Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Plot the distribution of `expert_consensus`\nsns.countplot(x='expert_consensus', data=df)\nplt.title('Distribution of expert consensus')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-05-11T07:02:02.337672Z","iopub.execute_input":"2024-05-11T07:02:02.338450Z","iopub.status.idle":"2024-05-11T07:02:02.760111Z","shell.execute_reply.started":"2024-05-11T07:02:02.338419Z","shell.execute_reply":"2024-05-11T07:02:02.759154Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Encode the `expert_consensus` column\nlabel_encoder = LabelEncoder()\ndf['expert_consensus'] = label_encoder.fit_transform(df['expert_consensus'])","metadata":{"execution":{"iopub.status.busy":"2024-05-11T07:02:06.696434Z","iopub.execute_input":"2024-05-11T07:02:06.696818Z","iopub.status.idle":"2024-05-11T07:02:06.732291Z","shell.execute_reply.started":"2024-05-11T07:02:06.696789Z","shell.execute_reply":"2024-05-11T07:02:06.731337Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Number of unique patients\nnum_patients = df['patient_id'].nunique()\nprint(f\"Number of unique patients in train dataset: {num_patients}\")\n\n# Number of unique EEG IDs\nnum_eeg_ids = df['eeg_id'].nunique()\nprint(f\"Number of unique EEG IDs in train dataset: {num_eeg_ids}\")","metadata":{"execution":{"iopub.status.busy":"2024-05-11T07:02:09.864025Z","iopub.execute_input":"2024-05-11T07:02:09.864808Z","iopub.status.idle":"2024-05-11T07:02:09.873824Z","shell.execute_reply.started":"2024-05-11T07:02:09.864772Z","shell.execute_reply":"2024-05-11T07:02:09.872883Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Histograms for numerical columns\ndf.hist(figsize=(15, 10))\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-05-11T07:02:12.859526Z","iopub.execute_input":"2024-05-11T07:02:12.859918Z","iopub.status.idle":"2024-05-11T07:02:15.798323Z","shell.execute_reply.started":"2024-05-11T07:02:12.859887Z","shell.execute_reply":"2024-05-11T07:02:15.797424Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Class Distribution\ndf['expert_consensus'].value_counts().plot(kind='bar')\nplt.title('Class Distribution')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-05-11T07:02:19.812641Z","iopub.execute_input":"2024-05-11T07:02:19.813451Z","iopub.status.idle":"2024-05-11T07:02:20.074417Z","shell.execute_reply.started":"2024-05-11T07:02:19.813422Z","shell.execute_reply":"2024-05-11T07:02:20.073626Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Load an EEG file\neeg = pd.read_parquet('/kaggle/input/hms-harmful-brain-activity-classification/train_eegs/1484166292.parquet')\neeg\n# List of columns to plot\ncolumns_to_plot = ['Fp1', 'F3', 'C3', 'P3', 'F7', 'T3', 'T5', 'O1', 'Fz', 'Cz', 'Pz', 'Fp2', 'F4', 'C4', 'P4', 'F8', 'T4', 'T6', 'O2', 'EKG']\n# Determine the number of rows/columns needed for subplots\nnum_plots = len(columns_to_plot)\nnum_columns = 2  # Set to 2 as per the previous code\nnum_rows = num_plots // num_columns + (num_plots % num_columns > 0)\n# Create subplots\nfig, axes = plt.subplots(num_rows, num_columns, figsize=(20, num_rows * 4))\n# Flatten the axes array for easy iteration\naxes = axes.flatten()\n# Plot each column in a subplot\nfor i, col in enumerate(columns_to_plot):\n    axes[i].plot(eeg[col])\n    axes[i].set_title(f'Electrode: {col}', fontsize=14)\n# Hide any unused subplots\nfor ax in axes[len(columns_to_plot):]:\n    ax.set_visible(False)\n# Set the overall figure title\nfig.suptitle('EEG Data Visualization Based on the International 10-20 System', fontsize=30, y=1.02)\n# Adjust layout to prevent overlap\nplt.tight_layout()\n# Show the plot\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-05-11T07:02:24.331906Z","iopub.execute_input":"2024-05-11T07:02:24.332280Z","iopub.status.idle":"2024-05-11T07:02:29.979746Z","shell.execute_reply.started":"2024-05-11T07:02:24.332250Z","shell.execute_reply":"2024-05-11T07:02:29.978778Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_egg_id = 3911565283\ntest_spec_id = 853520\n\neeg_test = pd.read_parquet(f'/kaggle/input/hms-harmful-brain-activity-classification/test_eegs/{test_egg_id}.parquet')    \nplt.figure(figsize=(20,5))\nplt.plot(eeg_test['Fp1'])\n\nspec_test = pd.read_parquet(f'/kaggle/input/hms-harmful-brain-activity-classification/test_spectrograms/{test_spec_id}.parquet')    \nplt.figure(figsize=(20,5))\nplt.plot(spec_test['time'], spec_test['LL_0.59'])","metadata":{"execution":{"iopub.status.busy":"2024-05-11T07:02:57.241865Z","iopub.execute_input":"2024-05-11T07:02:57.242646Z","iopub.status.idle":"2024-05-11T07:02:57.812473Z","shell.execute_reply.started":"2024-05-11T07:02:57.242611Z","shell.execute_reply":"2024-05-11T07:02:57.811603Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## load data","metadata":{}},{"cell_type":"code","source":"def build_augmenter(dim=CFG.image_size):\n    augmenters = [\n        keras_cv.layers.MixUp(alpha=2.0),\n        keras_cv.layers.RandomCutout(height_factor=(1.0, 1.0),\n                                     width_factor=(0.06, 0.1)), # freq-masking\n        keras_cv.layers.RandomCutout(height_factor=(0.06, 0.1),\n                                     width_factor=(1.0, 1.0)), # time-masking\n    ]\n    \n    def augment(img, label):\n        data = {\"images\":img, \"labels\":label}\n        for augmenter in augmenters:\n            if tf.random.uniform([]) < 0.5:\n                data = augmenter(data, training=True)\n        return data[\"images\"], data[\"labels\"]\n    \n    return augment","metadata":{"execution":{"iopub.status.busy":"2024-05-11T07:03:01.899671Z","iopub.execute_input":"2024-05-11T07:03:01.900330Z","iopub.status.idle":"2024-05-11T07:03:01.907390Z","shell.execute_reply.started":"2024-05-11T07:03:01.900293Z","shell.execute_reply":"2024-05-11T07:03:01.906450Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def build_decoder(with_labels=True, target_size=CFG.image_size, dtype=32):\n    def decode_signal(path, offset=None):\n        # Read .npy files and process the signal\n        file_bytes = tf.io.read_file(path)\n        sig = tf.io.decode_raw(file_bytes, tf.float32)\n        sig = sig[1024//dtype:]  # Remove header tag\n        sig = tf.reshape(sig, [400, -1])\n        \n        # Extract labeled subsample from full spectrogram using \"offset\"\n        if offset is not None: \n            offset = offset // 2  # Only odd values are given\n            sig = sig[:, offset:offset+300]\n            \n            # Pad spectrogram to ensure the same input shape of [400, 300]\n            pad_size = tf.math.maximum(0, 300 - tf.shape(sig)[1])\n            sig = tf.pad(sig, [[0, 0], [0, pad_size]])\n            sig = tf.reshape(sig, [400, 300])\n        \n        # Log spectrogram \n        sig = tf.clip_by_value(sig, tf.math.exp(-4.0), tf.math.exp(8.0)) # avoid 0 in log\n        sig = tf.math.log(sig)\n        # Normalize spectrogram\n        sig -= tf.math.reduce_mean(sig)\n        sig /= tf.math.reduce_std(sig) + 1e-6\n        \n        # Mono channel to 3 channels to use \"ImageNet\" weights\n        sig = tf.tile(sig[..., None], [1, 1, 3])\n        return sig\n    \n    def decode_label(label):\n        label = tf.one_hot(label, CFG.num_classes)\n        label = tf.cast(label, tf.float32)\n        label = tf.reshape(label, [CFG.num_classes])\n        return label\n    \n    def decode_with_labels(path, offset=None, label=None):\n        sig = decode_signal(path, offset)\n        label = decode_label(label)\n        return (sig, label)\n    \n    return decode_with_labels if with_labels else decode_signal","metadata":{"execution":{"iopub.status.busy":"2024-05-11T07:03:11.429850Z","iopub.execute_input":"2024-05-11T07:03:11.430647Z","iopub.status.idle":"2024-05-11T07:03:11.445938Z","shell.execute_reply.started":"2024-05-11T07:03:11.430607Z","shell.execute_reply":"2024-05-11T07:03:11.444908Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def build_dataset(paths, offsets=None, labels=None, batch_size=32, cache=True,\n                  decode_fn=None, augment_fn=None,\n                  augment=False, repeat=True, shuffle=1024, \n                  cache_dir=\"\", drop_remainder=False):\n    if cache_dir != \"\" and cache is True:\n        os.makedirs(cache_dir, exist_ok=True)\n    \n    if decode_fn is None:\n        decode_fn = build_decoder(labels is not None)\n    \n    if augment_fn is None:\n        augment_fn = build_augmenter()\n    \n    AUTO = tf.data.experimental.AUTOTUNE\n    slices = (paths, offsets) if labels is None else (paths, offsets, labels)\n    \n    ds = tf.data.Dataset.from_tensor_slices(slices)\n    ds = ds.map(decode_fn, num_parallel_calls=AUTO)\n    ds = ds.cache(cache_dir) if cache else ds\n    ds = ds.repeat() if repeat else ds\n    if shuffle: \n        ds = ds.shuffle(shuffle, seed=CFG.seed)\n        opt = tf.data.Options()\n        opt.experimental_deterministic = False\n        ds = ds.with_options(opt)\n    ds = ds.batch(batch_size, drop_remainder=drop_remainder)\n    ds = ds.map(augment_fn, num_parallel_calls=AUTO) if augment else ds\n    ds = ds.prefetch(AUTO)\n    return ds","metadata":{"execution":{"iopub.status.busy":"2024-05-11T07:03:19.450653Z","iopub.execute_input":"2024-05-11T07:03:19.451422Z","iopub.status.idle":"2024-05-11T07:03:19.461249Z","shell.execute_reply.started":"2024-05-11T07:03:19.451388Z","shell.execute_reply":"2024-05-11T07:03:19.460342Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## splitting data","metadata":{}},{"cell_type":"code","source":"from sklearn.model_selection import StratifiedGroupKFold\n\nsgkf = StratifiedGroupKFold(n_splits=5, shuffle=True, random_state=CFG.seed)\n\ndf[\"fold\"] = -1\ndf.reset_index(drop=True, inplace=True)\nfor fold, (train_idx, valid_idx) in enumerate(\n    sgkf.split(df, y=df[\"class_label\"], groups=df[\"patient_id\"])\n):\n    df.loc[valid_idx, \"fold\"] = fold\ndf.groupby([\"fold\", \"class_name\"])[[\"eeg_id\"]].count().T","metadata":{"execution":{"iopub.status.busy":"2024-05-11T07:03:24.295307Z","iopub.execute_input":"2024-05-11T07:03:24.295984Z","iopub.status.idle":"2024-05-11T07:03:25.533300Z","shell.execute_reply.started":"2024-05-11T07:03:24.295952Z","shell.execute_reply":"2024-05-11T07:03:25.532489Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train and Valid Dataset","metadata":{}},{"cell_type":"code","source":"# Sample from full data\nsample_df = df.groupby(\"spectrogram_id\").head(1).reset_index(drop=True)\ntrain_df = sample_df[sample_df.fold != CFG.fold]\nvalid_df = sample_df[sample_df.fold == CFG.fold]\nprint(f\"# Num Train: {len(train_df)} | Num Valid: {len(valid_df)}\")\n\n# Train\ntrain_paths = train_df.spec2_path.values\ntrain_offsets = train_df.spectrogram_label_offset_seconds.values.astype(int)\ntrain_labels = train_df.class_label.values\ntrain_ds = build_dataset(train_paths, train_offsets, train_labels, batch_size=CFG.batch_size,\n                         repeat=True, shuffle=True, augment=True, cache=True)\n\n# Valid\nvalid_paths = valid_df.spec2_path.values\nvalid_offsets = valid_df.spectrogram_label_offset_seconds.values.astype(int)\nvalid_labels = valid_df.class_label.values\nvalid_ds = build_dataset(valid_paths, valid_offsets, valid_labels, batch_size=CFG.batch_size,\n                         repeat=False, shuffle=False, augment=False, cache=True)","metadata":{"execution":{"iopub.status.busy":"2024-05-11T07:03:30.380229Z","iopub.execute_input":"2024-05-11T07:03:30.380798Z","iopub.status.idle":"2024-05-11T07:03:33.605373Z","shell.execute_reply.started":"2024-05-11T07:03:30.380767Z","shell.execute_reply":"2024-05-11T07:03:33.604280Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Sample Images","metadata":{}},{"cell_type":"code","source":"imgs, tars = next(iter(train_ds))\n\nnum_imgs = 8\nplt.figure(figsize=(4*4, num_imgs//4*5))\nfor i in range(num_imgs):\n    plt.subplot(num_imgs//4, 4, i + 1)\n    img = imgs[i].numpy()[...,0]  # Adjust as per your image data format\n    img -= img.min()\n    img /= img.max() + 1e-4\n    tar = CFG.label2name[np.argmax(tars[i].numpy())]\n    plt.imshow(img)\n    plt.title(f\"Target: {tar}\")\n    plt.axis('off')\n    \nplt.tight_layout()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-05-11T07:03:36.910474Z","iopub.execute_input":"2024-05-11T07:03:36.910854Z","iopub.status.idle":"2024-05-11T07:03:39.738726Z","shell.execute_reply.started":"2024-05-11T07:03:36.910825Z","shell.execute_reply":"2024-05-11T07:03:39.737523Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## EfficientNet Model","metadata":{}},{"cell_type":"code","source":"LOSS = keras.losses.KLDivergence()","metadata":{"execution":{"iopub.status.busy":"2024-05-11T07:03:48.271718Z","iopub.execute_input":"2024-05-11T07:03:48.272584Z","iopub.status.idle":"2024-05-11T07:03:48.276720Z","shell.execute_reply.started":"2024-05-11T07:03:48.272531Z","shell.execute_reply":"2024-05-11T07:03:48.275768Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from keras.applications import EfficientNetB0\nfrom keras.layers import GlobalAveragePooling2D, Dense\nfrom keras.models import Model","metadata":{"execution":{"iopub.status.busy":"2024-05-11T07:03:51.780549Z","iopub.execute_input":"2024-05-11T07:03:51.781419Z","iopub.status.idle":"2024-05-11T07:03:51.785648Z","shell.execute_reply.started":"2024-05-11T07:03:51.781386Z","shell.execute_reply":"2024-05-11T07:03:51.784621Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"base_model = EfficientNetB0(weights='imagenet', include_top=False)","metadata":{"execution":{"iopub.status.busy":"2024-05-11T07:04:33.364373Z","iopub.execute_input":"2024-05-11T07:04:33.364746Z","iopub.status.idle":"2024-05-11T07:04:55.964677Z","shell.execute_reply.started":"2024-05-11T07:04:33.364716Z","shell.execute_reply":"2024-05-11T07:04:55.963771Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Add custom classification head\nx = GlobalAveragePooling2D()(base_model.output)\nx = Dense(CFG.num_classes, activation='softmax')(x)\n","metadata":{"execution":{"iopub.status.busy":"2024-05-11T07:04:55.966385Z","iopub.execute_input":"2024-05-11T07:04:55.966685Z","iopub.status.idle":"2024-05-11T07:04:56.115160Z","shell.execute_reply.started":"2024-05-11T07:04:55.966660Z","shell.execute_reply":"2024-05-11T07:04:56.114214Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create the final model\nmodel = Model(inputs=base_model.input, outputs=x)\n","metadata":{"execution":{"iopub.status.busy":"2024-05-11T07:05:01.766982Z","iopub.execute_input":"2024-05-11T07:05:01.767705Z","iopub.status.idle":"2024-05-11T07:05:01.784398Z","shell.execute_reply.started":"2024-05-11T07:05:01.767670Z","shell.execute_reply":"2024-05-11T07:05:01.783511Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Compile the model\nmodel.compile(optimizer=keras.optimizers.Adam(learning_rate=1e-4),\n              loss=LOSS,\n              metrics=['accuracy'])\n\n# Model Summary\nmodel.summary()","metadata":{"execution":{"iopub.status.busy":"2024-05-11T07:05:05.849440Z","iopub.execute_input":"2024-05-11T07:05:05.849799Z","iopub.status.idle":"2024-05-11T07:05:06.192033Z","shell.execute_reply.started":"2024-05-11T07:05:05.849773Z","shell.execute_reply":"2024-05-11T07:05:06.191147Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## lR Scheduling","metadata":{}},{"cell_type":"code","source":"import math\ndef get_lr_callback(batch_size=8, mode='cos', epochs=10, plot=False):\n    lr_start, lr_max, lr_min = 5e-5, 6e-6 * batch_size, 1e-5\n    lr_ramp_ep, lr_sus_ep, lr_decay = 3, 0, 0.75\n\n    def lrfn(epoch):  # Learning rate update function\n        if epoch < lr_ramp_ep: lr = (lr_max - lr_start) / lr_ramp_ep * epoch + lr_start\n        elif epoch < lr_ramp_ep + lr_sus_ep: lr = lr_max\n        elif mode == 'exp': lr = (lr_max - lr_min) * lr_decay**(epoch - lr_ramp_ep - lr_sus_ep) + lr_min\n        elif mode == 'step': lr = lr_max * lr_decay**((epoch - lr_ramp_ep - lr_sus_ep) // 2)\n        elif mode == 'cos':\n            decay_total_epochs, decay_epoch_index = epochs - lr_ramp_ep - lr_sus_ep + 3, epoch - lr_ramp_ep - lr_sus_ep\n            phase = math.pi * decay_epoch_index / decay_total_epochs\n            lr = (lr_max - lr_min) * 0.5 * (1 + math.cos(phase)) + lr_min\n        return lr\n\n    if plot:  # Plot lr curve if plot is True\n        plt.figure(figsize=(10, 5))\n        plt.plot(np.arange(epochs), [lrfn(epoch) for epoch in np.arange(epochs)], marker='o')\n        plt.xlabel('epoch'); plt.ylabel('lr')\n        plt.title('LR Scheduler')\n        plt.show()\n    return keras.callbacks.LearningRateScheduler(lrfn, verbose=False)  # Create lr callback","metadata":{"execution":{"iopub.status.busy":"2024-05-11T07:05:20.257491Z","iopub.execute_input":"2024-05-11T07:05:20.257867Z","iopub.status.idle":"2024-05-11T07:05:20.267994Z","shell.execute_reply.started":"2024-05-11T07:05:20.257837Z","shell.execute_reply":"2024-05-11T07:05:20.267068Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"lr_cb = get_lr_callback(CFG.batch_size, mode=CFG.lr_mode, plot=True)","metadata":{"execution":{"iopub.status.busy":"2024-05-11T07:05:25.450524Z","iopub.execute_input":"2024-05-11T07:05:25.450904Z","iopub.status.idle":"2024-05-11T07:05:25.673176Z","shell.execute_reply.started":"2024-05-11T07:05:25.450876Z","shell.execute_reply":"2024-05-11T07:05:25.672256Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Model Checkpoint Callback\nckpt_cb = keras.callbacks.ModelCheckpoint(f\"best_model_efficientnet_custom_fold{CFG.fold}.keras\",\n                                           monitor='val_loss',\n                                           save_best_only=True,\n                                           save_weights_only=False,\n                                           mode='min')","metadata":{"execution":{"iopub.status.busy":"2024-05-11T07:05:30.085597Z","iopub.execute_input":"2024-05-11T07:05:30.085953Z","iopub.status.idle":"2024-05-11T07:05:30.090536Z","shell.execute_reply.started":"2024-05-11T07:05:30.085926Z","shell.execute_reply":"2024-05-11T07:05:30.089639Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Train the model\nhistory = model.fit(train_ds,\n                    epochs=CFG.epochs,\n                    callbacks=[lr_cb, ckpt_cb],\n                    steps_per_epoch=len(train_df) // CFG.batch_size,\n                    validation_data=valid_ds,\n                    verbose=CFG.verbose)","metadata":{"execution":{"iopub.status.busy":"2024-05-11T07:05:36.390287Z","iopub.execute_input":"2024-05-11T07:05:36.390654Z","iopub.status.idle":"2024-05-11T07:48:22.540788Z","shell.execute_reply.started":"2024-05-11T07:05:36.390626Z","shell.execute_reply":"2024-05-11T07:48:22.539608Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import plotly.graph_objects as go","metadata":{"execution":{"iopub.status.busy":"2024-05-11T07:48:22.577854Z","iopub.execute_input":"2024-05-11T07:48:22.578163Z","iopub.status.idle":"2024-05-11T07:48:22.586309Z","shell.execute_reply.started":"2024-05-11T07:48:22.578140Z","shell.execute_reply":"2024-05-11T07:48:22.585430Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def display_training_curves(training, validation, yaxis):\n    if yaxis == \"loss\":\n        ylabel = \"Loss\"\n        title = \"Loss vs. Epochs\"\n    else:\n        ylabel = \"Accuracy\"\n        title = \"Accuracy vs. Epochs\"\n        \n    fig = go.Figure()\n        \n    fig.add_trace(\n        go.Scatter(x=np.arange(1, CFG.epochs+1), mode='lines+markers', y=training, marker=dict(color=\"dodgerblue\"),\n               name=\"Train\"))\n    \n    fig.add_trace(\n        go.Scatter(x=np.arange(1, CFG.epochs+1), mode='lines+markers', y=validation, marker=dict(color=\"darkorange\"),\n               name=\"Val\"))\n    \n    fig.update_layout(title_text=title, yaxis_title=ylabel, xaxis_title=\"Epochs\", template=\"plotly_white\")\n    fig.show()","metadata":{"execution":{"iopub.status.busy":"2024-05-11T07:48:22.587466Z","iopub.execute_input":"2024-05-11T07:48:22.587802Z","iopub.status.idle":"2024-05-11T07:48:22.596342Z","shell.execute_reply.started":"2024-05-11T07:48:22.587773Z","shell.execute_reply":"2024-05-11T07:48:22.595474Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"display_training_curves(\n    history.history['accuracy'], \n    history.history['val_accuracy'], \n    'accuracy')","metadata":{"execution":{"iopub.status.busy":"2024-05-11T07:48:22.597495Z","iopub.execute_input":"2024-05-11T07:48:22.597841Z","iopub.status.idle":"2024-05-11T07:48:24.089302Z","shell.execute_reply.started":"2024-05-11T07:48:22.597811Z","shell.execute_reply":"2024-05-11T07:48:24.088369Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Resnet","metadata":{}},{"cell_type":"code","source":"from keras.applications import ResNet50\nfrom keras.layers import GlobalAveragePooling2D, Dense\nfrom keras.models import Model","metadata":{"execution":{"iopub.status.busy":"2024-05-11T07:52:40.163375Z","iopub.execute_input":"2024-05-11T07:52:40.163814Z","iopub.status.idle":"2024-05-11T07:52:40.168556Z","shell.execute_reply.started":"2024-05-11T07:52:40.163781Z","shell.execute_reply":"2024-05-11T07:52:40.167644Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Load pre-trained ResNet50 model without the top layers\nbase_model = ResNet50(weights='imagenet', include_top=False)","metadata":{"execution":{"iopub.status.busy":"2024-05-11T07:52:57.479684Z","iopub.execute_input":"2024-05-11T07:52:57.480466Z","iopub.status.idle":"2024-05-11T07:53:04.314033Z","shell.execute_reply.started":"2024-05-11T07:52:57.480433Z","shell.execute_reply":"2024-05-11T07:53:04.313204Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}