{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":70203,"databundleVersionId":8068726},{"sourceType":"modelInstanceVersion","sourceId":6127,"databundleVersionId":7429415,"modelInstanceId":4598},{"sourceType":"modelInstanceVersion","sourceId":6091,"databundleVersionId":7429345,"modelInstanceId":4623}],"dockerImageVersionId":30886,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport os\nos.environ[\"KERAS_BACKEND\"] = \"tensorflow\"  # \"jax\" or \"tensorflow\" or \"torch\"\n\nimport keras_cv\nimport keras\nimport keras.backend as K\nimport tensorflow as tf\nimport numpy as np\nimport pandas as pd\n\nfrom glob import glob\nfrom tqdm import tqdm\n\nimport librosa\nimport IPython.display as ipd\nimport librosa.display as lid\n\nimport matplotlib.pyplot as plt\nimport matplotlib as mpl\nimport matplotlib.pylab as ply\nimport ipywidgets as widgets\nimport seaborn as sns\n\nfrom itertools import cycle\n# Set interactive backend\n%matplotlib inline\ncmap = mpl.cm.get_cmap('coolwarm')\nsns.set_theme(style=\"white\", palette=None)\ncolor_pal = ply.rcParams[\"axes.prop_cycle\"].by_key()[\"color\"]\ncolor_cycle = cycle(ply.rcParams[\"axes.prop_cycle\"].by_key()[\"color\"])\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n# for dirname, _, filenames in os.walk('/kaggle/input'):\n#     for filename in filenames:\n#         print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-05-01T06:08:47.964930Z","iopub.execute_input":"2025-05-01T06:08:47.965265Z","iopub.status.idle":"2025-05-01T06:09:01.354870Z","shell.execute_reply.started":"2025-05-01T06:08:47.965239Z","shell.execute_reply":"2025-05-01T06:09:01.353903Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DATASET_PATH = '/kaggle/input/birdclef-2024'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T06:09:01.356100Z","iopub.execute_input":"2025-05-01T06:09:01.356719Z","iopub.status.idle":"2025-05-01T06:09:01.360181Z","shell.execute_reply.started":"2025-05-01T06:09:01.356692Z","shell.execute_reply":"2025-05-01T06:09:01.359412Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class_names = sorted(os.listdir(f\"{DATASET_PATH}/train_audio/\"))\nnum_classes = len(class_names)\nclass_labels = list(range(num_classes))\nlabel2name = dict(zip(class_labels, class_names))\nname2label = {v:k for k,v in label2name.items()}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T06:09:06.452968Z","iopub.execute_input":"2025-05-01T06:09:06.453420Z","iopub.status.idle":"2025-05-01T06:09:06.458772Z","shell.execute_reply.started":"2025-05-01T06:09:06.453382Z","shell.execute_reply":"2025-05-01T06:09:06.458017Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## Print out the first 5 items in the label2name and name2label dictionaries\nprint(f\"Number of classes: {num_classes}\")\nprint({k: label2name[k] for k in list(label2name)[:5]})\nprint({k: name2label[k] for k in list(name2label)[:5]})","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T06:09:08.207580Z","iopub.execute_input":"2025-05-01T06:09:08.207897Z","iopub.status.idle":"2025-05-01T06:09:08.214835Z","shell.execute_reply.started":"2025-05-01T06:09:08.207874Z","shell.execute_reply":"2025-05-01T06:09:08.213887Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = pd.read_csv(f'{DATASET_PATH}/train_metadata.csv')\ndf['filepath'] = DATASET_PATH + '/train_audio/' + df.filename\ndf['target'] = df.primary_label.map(name2label)\ndf['filename'] = df.filepath.map(lambda x: x.split('/')[-1])\ndf['xc_id'] = df.filepath.map(lambda x: x.split('/')[-1].split('.')[0])\n\n## display a few rows of the dataframe from columns ['scientific_name', 'scientific_name',  'filepath']\ndf = df.sample(frac=1, random_state=42)\ndf.head(5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T06:09:10.531133Z","iopub.execute_input":"2025-05-01T06:09:10.531475Z","iopub.status.idle":"2025-05-01T06:09:10.750488Z","shell.execute_reply.started":"2025-05-01T06:09:10.531448Z","shell.execute_reply":"2025-05-01T06:09:10.749551Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## Display the number of samples per class and save the result in a dictionary\nclass_counts = df.primary_label.value_counts()\nclass_counts = class_counts.sort_index()\nclass_counts\n# ## Save to a csv file\npd.DataFrame(class_counts.items(), columns=['class', 'count']).to_csv('class_counts.csv', index=False)\n\n## Show the largest and smallest classes\nclass_counts_csv = pd.read_csv('class_counts.csv')\n## Show the largest and smallest classes with the corresponding counts\n# Find the minimum and maximum counts\nmin_count = class_counts.min()\nmax_count = class_counts.max()\n \nprint(f\"Smallest class: {class_counts_csv['class'][class_counts_csv['count'].idxmin()]} {min_count}\")\nprint(f\"Largest class: {class_counts_csv['class'][class_counts_csv['count'].idxmax()]} {max_count}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T06:09:13.826051Z","iopub.execute_input":"2025-05-01T06:09:13.826425Z","iopub.status.idle":"2025-05-01T06:09:14.101896Z","shell.execute_reply.started":"2025-05-01T06:09:13.826391Z","shell.execute_reply":"2025-05-01T06:09:14.100810Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## Explore more statistics of the dataset\nstatistics = class_counts_csv['count'].describe(percentiles=[.25, .5, .75])\n\n# Show statistics\nprint(f\"Mean: {statistics['mean']}\")\nprint(f\"Median (50%): {statistics['50%']}\")\nprint(f\"Standard Deviation: {statistics['std']}\")\nprint(f\"Minimum: {statistics['min']}\")\nprint(f\"Maximum: {statistics['max']}\")\nprint(f\"25th Percentile: {statistics['25%']}\")\nprint(f\"50th Percentile (Median): {statistics['50%']}\")\nprint(f\"75th Percentile: {statistics['75%']}\")\n\n# You can also calculate specific quantiles, for example:\nquantile_10 = class_counts_csv['count'].quantile(0.10)\nquantile_90 = class_counts_csv['count'].quantile(0.90)\n\nprint(f\"10th Percentile: {quantile_10}\")\nprint(f\"90th Percentile: {quantile_90}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T06:09:17.058110Z","iopub.execute_input":"2025-05-01T06:09:17.058458Z","iopub.status.idle":"2025-05-01T06:09:17.073114Z","shell.execute_reply.started":"2025-05-01T06:09:17.058431Z","shell.execute_reply":"2025-05-01T06:09:17.072162Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"percentiles = [0.10, 0.25, 0.50, 0.75, 0.90]\npercentile_counts = {p: (class_counts_csv['count'] < class_counts_csv['count'].quantile(p)).sum() for p in percentiles}\nprint(\"Number of classes below each percentile:\")\nfor p, count in percentile_counts.items():\n    print(f\"{int(p*100)}%: {count} classes\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T06:09:19.537107Z","iopub.execute_input":"2025-05-01T06:09:19.537502Z","iopub.status.idle":"2025-05-01T06:09:19.548929Z","shell.execute_reply.started":"2025-05-01T06:09:19.537470Z","shell.execute_reply":"2025-05-01T06:09:19.548072Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## Number of classes with less than 14 samples\nless_than_20 = (class_counts_csv['count'] < 20).sum()\nprint(f\"Number of classes with less than 20 samples: {less_than_20}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T06:09:22.452518Z","iopub.execute_input":"2025-05-01T06:09:22.452869Z","iopub.status.idle":"2025-05-01T06:09:22.457987Z","shell.execute_reply.started":"2025-05-01T06:09:22.452839Z","shell.execute_reply":"2025-05-01T06:09:22.456955Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## Plot the distribution of class sizes\nplt.figure(figsize=(12, 6))\nplt.bar(class_counts.index, class_counts.values, color=color_pal)\nplt.xticks(rotation=90)\nplt.xlabel('Class')\nplt.ylabel('Number of samples')\nplt.title('Distribution of class sizes')\nplt.show()\n\n\n\n### Group the classes by the class counts and plot the distribution of class sizes\n# Group the classes by the class counts\nclass_counts = pd.DataFrame(class_counts)\nclass_counts['class'] = class_counts.index\nclass_counts['count'] = class_counts['count'].astype(int)\nclass_counts['group'] = pd.cut(class_counts['count'], bins=[0, 10, 20, 50, 100, 200, 500, 1000, 2000, 5000], labels=['0-10', '10-20', '20-50', '50-100', '100-200', '200-500', '500-1000', '1000-2000', '2000-5000'])\n# Plot the distribution of class sizes\nplt.figure(figsize=(12, 6))\nplt.bar(class_counts['group'].cat.codes, class_counts['count'], color=color_pal)\nplt.xticks(rotation=90)\nplt.xlabel('Class')\nplt.ylabel('Number of samples')\nplt.title('Distribution of class sizes')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T06:09:30.756378Z","iopub.execute_input":"2025-05-01T06:09:30.756710Z","iopub.status.idle":"2025-05-01T06:09:32.957434Z","shell.execute_reply.started":"2025-05-01T06:09:30.756685Z","shell.execute_reply":"2025-05-01T06:09:32.956357Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## Check the size of the dataframe\n## The shape of the dataframe should be (24459, 15) which means there are 24459 rows and 15 columns\nprint(f\"Dataframe shape: {df.shape}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T06:09:36.141779Z","iopub.execute_input":"2025-05-01T06:09:36.142098Z","iopub.status.idle":"2025-05-01T06:09:36.146801Z","shell.execute_reply.started":"2025-05-01T06:09:36.142067Z","shell.execute_reply":"2025-05-01T06:09:36.145873Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## Drop classes with less than 20 samples\n# 1. Load class names and metadata\nclass_names = sorted(os.listdir(f\"{DATASET_PATH}/train_audio/\"))\ndf = pd.read_csv(f'{DATASET_PATH}/train_metadata.csv')\ndf['filepath'] = DATASET_PATH + '/train_audio/' + df.filename\ndf['filename'] = df.filepath.map(lambda x: x.split('/')[-1])\ndf['xc_id'] = df.filepath.map(lambda x: x.split('/')[-1].split('.')[0])\ndf['primary_label'] = df['primary_label'].astype(str)  # Ensure consistent type\n\n# 2. Filter out classes with fewer than 20 samples\nlabel_counts = df['primary_label'].value_counts()\nvalid_labels = label_counts[label_counts >= 20].index.tolist()\ndf = df[df['primary_label'].isin(valid_labels)]\n\n# 3. Recompute class info\nclass_names = sorted(valid_labels)\nnum_classes = len(class_names)\nclass_labels = list(range(num_classes))\nlabel2name = dict(zip(class_labels, class_names))\nname2label = {v: k for k, v in label2name.items()}\n\n# 4. Reset the `target` column using updated name2label\ndf['target'] = df['primary_label'].map(name2label)\n\n# 5. Shuffle and show\ndf = df.sample(frac=1, random_state=42)\nprint(f\"Number of classes: {num_classes}\")\nprint({k: label2name[k] for k in list(label2name)[:5]})\nprint({k: name2label[k] for k in list(name2label)[:5]})\ndf.head(5)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T06:09:39.111748Z","iopub.execute_input":"2025-05-01T06:09:39.112157Z","iopub.status.idle":"2025-05-01T06:09:39.272935Z","shell.execute_reply.started":"2025-05-01T06:09:39.112127Z","shell.execute_reply":"2025-05-01T06:09:39.271966Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Function to retreive an audio file 🎵\n**librosa is a python package for music and audio analysis. It provides the building blocks necessary to create music information retrieval systems**\n[Documentation here](https://librosa.org/doc/latest/index.html)","metadata":{}},{"cell_type":"code","source":"## Load the audio as a waveform `y`\n# Store the sampling rate as `sr`\ndef load_audio(filepath):\n    audio, sr = librosa.load(filepath)\n    return audio, sr","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T05:13:39.303144Z","iopub.execute_input":"2025-05-01T05:13:39.303456Z","iopub.status.idle":"2025-05-01T05:13:39.307417Z","shell.execute_reply.started":"2025-05-01T05:13:39.303434Z","shell.execute_reply":"2025-05-01T05:13:39.306428Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import random\nfor i in range(2):\n    # random_index = random.randint(0, df.shape[0])\n    ipd.Audio(df['filepath'].iloc[i])\n    audio, sr = load_audio(df['filepath'].iloc[i])\n    plt.figure(figsize=(10, 3))\n    pd.Series(audio).plot(figsize=(10, 5),\n                    lw=1,\n                    # title=f\"{df['scientific_name'].iloc[i]}\",\n                    color=color_pal[0])\n    ## Zoomed in sample to view waves better:\n    plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T05:13:39.308371Z","iopub.execute_input":"2025-05-01T05:13:39.308644Z","iopub.status.idle":"2025-05-01T05:13:40.073521Z","shell.execute_reply.started":"2025-05-01T05:13:39.308616Z","shell.execute_reply":"2025-05-01T05:13:40.072594Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### From the plotted files, we can see that the bird recordings occur where the amplitude (height) of the wave is high","metadata":{}},{"cell_type":"code","source":"#### Understanding the audio data\nfor i in range(2):\n    # random_index = random.randint(0, df.shape[0])\n    ipd.Audio(df['filepath'].iloc[i])\n    audio, sr = load_audio(df['filepath'].iloc[i])\n    print(f\"Audio: {audio}\")\n    print(f\"Shape of the audio: {audio.shape}\")\n## The audio file is a numpy array. However, the size of the arrays are different hence we need to pad/trim the arrays to make them the same size\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T05:13:40.074454Z","iopub.execute_input":"2025-05-01T05:13:40.074759Z","iopub.status.idle":"2025-05-01T05:13:40.144698Z","shell.execute_reply.started":"2025-05-01T05:13:40.074730Z","shell.execute_reply":"2025-05-01T05:13:40.143903Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### Preview a sample of audio spectrograms\nThe STFT represents a signal in the time-frequency domain by computing discrete Fourier transforms (DFT) over short overlapping windows.","metadata":{}},{"cell_type":"code","source":"\"\"\"\n function to preview simple spectrograms in decibels. A decibel is a logarithmic unit that expresses \n the ratio of two values of a physical quantity, often power or intensity.\n \"\"\"\ndef audio_to_spectrogram(audio):\n    D = librosa.stft(audio)\n    S_db = librosa.amplitude_to_db(np.abs(D), ref=np.max)\n    print(S_db.shape)\n\n    fig, ax = plt.subplots(figsize=(10, 5))\n    img = librosa.display.specshow(S_db,\n                                x_axis='time',\n                                y_axis='log',\n                                ax=ax)\n    # ax.set_title(f\"{df['scientific_name'].iloc[i]} Audio Spectogram\", fontsize=20)\n    fig.colorbar(img, ax=ax, format=f'%0.2f')\n    plt.show()\n\nfor i in range(2):\n    audio, sr = load_audio(df['filepath'].iloc[i])\n    audio_to_spectrogram(audio)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T05:13:40.145599Z","iopub.execute_input":"2025-05-01T05:13:40.145953Z","iopub.status.idle":"2025-05-01T05:13:41.698956Z","shell.execute_reply.started":"2025-05-01T05:13:40.145919Z","shell.execute_reply":"2025-05-01T05:13:41.698001Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nWhile a regular spectrogram uses a linear frequency scale, \na Mel spectrogram uses the Mel scale, which is designed to better reflect how humans perceive sound.\n\"\"\"\ndef audio_to_melspectrogram(audio, sr):\n    S = librosa.feature.melspectrogram(y=audio,\n                                   sr=sr,\n                                   n_mels=128 * 2,)\n    S_db_mel = librosa.amplitude_to_db(S, ref=np.max)\n    # print(S_db_mel.shape)\n    fig, ax = plt.subplots(figsize=(10, 5))\n    # Plot the mel spectogram\n    img = librosa.display.specshow(S_db_mel,\n                                x_axis='time',\n                                y_axis='log',\n                                ax=ax)\n    # ax.set_title('Mel Spectogram Example', fontsize=20)\n    fig.colorbar(img, ax=ax, format=f'%0.2f')\n    plt.show()\n\nfor i in range(2):\n    random_index = random.randint(0, df.shape[0])\n    audio, sr = load_audio(df['filepath'].iloc[i])\n    audio_to_melspectrogram(audio, sr)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T05:13:41.700000Z","iopub.execute_input":"2025-05-01T05:13:41.700294Z","iopub.status.idle":"2025-05-01T05:13:42.526153Z","shell.execute_reply.started":"2025-05-01T05:13:41.700267Z","shell.execute_reply":"2025-05-01T05:13:42.525295Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Get the audio spectrogram 🌊. \n**A spectrogram is a visual representation of the spectrum of frequencies of a signal as it varies with time. When applied to an audio signal, spectrograms are sometimes called sonographs, voiceprints, or voicegrams**","metadata":{}},{"cell_type":"code","source":"# Define the sampling rate of the audio signal (32 kHz)\nsample_rate = 32000\n\n# Define the maximum frequency to include in the spectrogram (16 kHz)\nfmax = 16000\n\n# Define the minimum frequency to include in the spectrogram (20 Hz)\nfmin = 20\n\n# Function to compute the Mel-spectrogram of an audio signal\ndef get_spectrogram(audio):\n    # Compute the Mel-spectrogram\n    spec = librosa.feature.melspectrogram(\n        y=audio,  # Input audio signal\n        sr=sample_rate,  # Sampling rate of the audio\n        n_mels=256,  # Number of Mel bands (frequency bins)\n        n_fft=2048,  # Size of the FFT window (determines frequency resolution)\n        hop_length=512,  # Number of samples between successive frames (determines time resolution)\n        fmax=fmax,  # Maximum frequency to include in the spectrogram\n        fmin=fmin,  # Minimum frequency to include in the spectrogram\n    )\n\n    # Convert the power spectrogram to decibel (dB) scale\n    # This makes the values more perceptually meaningful\n    spec = librosa.power_to_db(spec, ref=1.0)  # ref=1.0 is the reference value for dB calculation\n\n    # Normalize the spectrogram to the range [0, 1]\n    min_ = spec.min()  # Minimum value in the spectrogram\n    max_ = spec.max()  # Maximum value in the spectrogram\n    if max_ != min_:  # Avoid division by zero if the spectrogram is constant\n        spec = (spec - min_) / (max_ - min_)  # Normalize using min-max scaling\n\n    # print(spec.shape)\n    # Return the normalized Mel-spectrogram\n    return spec","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T06:09:55.487954Z","iopub.execute_input":"2025-05-01T06:09:55.488311Z","iopub.status.idle":"2025-05-01T06:09:55.494026Z","shell.execute_reply.started":"2025-05-01T06:09:55.488286Z","shell.execute_reply":"2025-05-01T06:09:55.493015Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"duration = 15\naudio_len = duration * sample_rate\ndef display_audio(row):\n    caption = f'Id: {row.filename} | Name: {row.common_name} | Sci.Name: {row.scientific_name}'\n    \n    audio, sr = load_audio(row.filepath)\n    audio = audio[:audio_len]\n    spec = get_spectrogram(audio)\n    \n    # Audio output widget\n    audio_output = widgets.Output()\n    with audio_output:\n        display(ipd.Audio(audio, rate=sample_rate))\n    \n    # Plot output widget\n    plot_output = widgets.Output()\n    with plot_output:\n        fig, ax = plt.subplots(2, 1, figsize=(12, 6), sharex=True, tight_layout=True)\n        # fig.suptitle(caption)\n        \n        # Plot waveform\n        lid.waveshow(audio, sr=sample_rate, ax=ax[0], color='b')\n        \n        # Plot spectrogram\n        lid.specshow(spec, sr=sample_rate, hop_length=512, n_fft=2048,\n                     fmin=fmin, fmax=fmax, x_axis='time', y_axis='mel', \n                     cmap='coolwarm', ax=ax[1])\n        \n        ax[0].set_xlabel('')\n        plt.show()\n\n    # Display side-by-side\n    display(widgets.HBox([audio_output, plot_output]))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T06:09:58.314706Z","iopub.execute_input":"2025-05-01T06:09:58.315027Z","iopub.status.idle":"2025-05-01T06:09:58.321648Z","shell.execute_reply.started":"2025-05-01T06:09:58.315004Z","shell.execute_reply":"2025-05-01T06:09:58.320650Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## Display a few audio samples\nfor i in range(2):\n    display_audio(df.sample(1).iloc[0])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T05:13:42.554483Z","iopub.execute_input":"2025-05-01T05:13:42.554665Z","iopub.status.idle":"2025-05-01T05:13:44.954627Z","shell.execute_reply.started":"2025-05-01T05:13:42.554649Z","shell.execute_reply":"2025-05-01T05:13:44.953865Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Build a decoder parse files into spectrograms🚀 \n\n**The build_decoder() function constructs a decoder that can process audio files into spectrograms.\nIt loads, normalizes, and converts the audio into a Mel-spectrogram.\nIf with_labels=True, it also converts labels into one-hot vectors.\nThe output is an RGB-like spectrogram image that can be used as input to CNNs.**\n[Tensorflow Documentation here](https://www.tensorflow.org/io/api_docs/python/tfio/audio/spectrogram)","metadata":{}},{"cell_type":"code","source":"# Image and audio parameters\nimg_size = [128, 384]  # Spectrogram image size (height, width)\nbatch_size = 64  # Batch size for training\n## The hop length is the number of samples between successive frames. \n# The audio length is divided by the width of the image to determine the hop length\nhop_length = audio_len // (img_size[1] - 1)  # What does this do? \nnfft = 2028  # FFT window size for computing the spectrogram\n\n## Explain the parameters\n\n\ndef build_decoder(with_labels=True, dim=1024):\n    \"\"\"\n    Builds a function to decode and preprocess audio files into spectrograms.\n    \n    Parameters:\n    - with_labels (bool): Whether to return labels along with spectrograms.\n    - dim (int): Target audio length (number of samples).\n    \n    Returns:\n    - Function to decode audio files (with or without labels).\n    \"\"\"\n    def get_audio(filepath):\n        \"\"\"Loads and decodes an audio file from a given filepath using librosa.\"\"\"\n        def _load_audio(filepath):\n            # Load the audio file using librosa\n            audio, _ = librosa.load(filepath.numpy().decode('utf-8'), sr=sample_rate, mono=True)\n            return audio.astype(np.float32)  # Ensure the audio is in float32 format\n\n        # Use tf.py_function to wrap the librosa call\n        audio = tf.py_function(_load_audio, [filepath], tf.float32)\n        audio.set_shape([None])  # Set shape to [None] since the length may vary\n        return audio\n\n    def crop_or_pad(audio, target_len, pad_mode=\"constant\"):\n        \"\"\"Ensures the audio is of fixed length by either cropping or padding.\"\"\"\n        audio_len = tf.shape(audio)[0]  # Get current length of audio\n        diff_len = abs(target_len - audio_len)  # Difference from target length\n\n        if audio_len < target_len:\n            # If audio is shorter, pad it randomly on both sides\n            pad1 = tf.random.uniform([], maxval=diff_len, dtype=tf.int32)\n            pad2 = diff_len - pad1\n            audio = tf.pad(audio, paddings=[[pad1, pad2]], mode=pad_mode)\n\n        elif audio_len > target_len:\n            # If audio is longer, randomly crop a section\n            idx = tf.random.uniform([], maxval=diff_len, dtype=tf.int32)\n            audio = audio[idx : (idx + target_len)]\n\n        return tf.reshape(audio, [target_len])  # Ensure fixed shape\n\n    def apply_preproc(spec):\n        \"\"\"Applies standardization and normalization to the spectrogram.\"\"\"\n        # Standardization: Zero mean and unit variance\n        mean = tf.math.reduce_mean(spec)\n        std = tf.math.reduce_std(spec)\n        spec = tf.where(tf.math.equal(std, 0), spec - mean, (spec - mean) / std)\n\n        # Min-Max Normalization: Scale values between 0 and 1\n        min_val = tf.math.reduce_min(spec)\n        max_val = tf.math.reduce_max(spec)\n        spec = tf.where(\n            tf.math.equal(max_val - min_val, 0), \n            spec - min_val, \n            (spec - min_val) / (max_val - min_val)\n        )\n\n        return spec\n\n    def get_target(target):\n        \"\"\"Converts a label into a one-hot encoded vector.\"\"\"\n        target = tf.reshape(target, [1])  # Reshape to single element tensor\n        target = tf.cast(tf.one_hot(target, num_classes), tf.float32)  # One-hot encoding\n        return tf.reshape(target, [num_classes])  # Reshape to match the output format\n\n    def decode(path):\n        \"\"\"Processes an audio file into a spectrogram image.\"\"\"\n        # Load and preprocess the audio\n        audio = get_audio(path)\n        audio = crop_or_pad(audio, dim)  # Ensure fixed length\n        \n        # Convert audio to a Mel-spectrogram\n        spec = keras.layers.MelSpectrogram(\n            num_mel_bins=img_size[0],  # Number of Mel frequency bins (height of image)\n            fft_length=nfft,  # FFT window size\n            sequence_stride=hop_length,  # Step size between spectrogram columns\n            sampling_rate=sample_rate,  # Sample rate of audio\n        )(audio)\n\n        spec = apply_preproc(spec)  # Apply normalization and standardization\n        \n        # Convert spectrogram into a 3-channel image (for compatibility with CNNs)\n        spec = tf.tile(spec[..., None], [1, 1, 3])  # Repeat values along the last axis\n        return tf.reshape(spec, [*img_size, 3])  # Reshape to (height, width, 3)\n\n    def decode_with_labels(path, label):\n        \"\"\"Processes an audio file into a spectrogram and returns it with its label.\"\"\"\n        return decode(path), get_target(label)\n\n    return decode_with_labels if with_labels else decode","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T06:10:04.972062Z","iopub.execute_input":"2025-05-01T06:10:04.972433Z","iopub.status.idle":"2025-05-01T06:10:04.983709Z","shell.execute_reply.started":"2025-05-01T06:10:04.972404Z","shell.execute_reply":"2025-05-01T06:10:04.982778Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Image Augmentation ♻\n##### augmentation involves applying a variety of transformations to the original dataset, generating new samples that are similar but not identical to the original data. Common augmentations include rotation, flipping, scaling, changes in brightness and contrast, color space adjustments, and geometric transformations","metadata":{}},{"cell_type":"code","source":"def build_augmenter():\n    \"\"\"\n    Creates an augmentation pipeline for spectrogram images.\n    Uses MixUp, time masking, and frequency masking to improve model generalization.\n    \n    Returns:\n        A function that applies random augmentations to images and labels.\n    \"\"\"\n\n    # Define a list of augmentation techniques to apply\n    augmenters = [\n        keras_cv.layers.MixUp(alpha=0.4),  # MixUp augmentation for blending two images\n        keras_cv.layers.RandomCutout(\n            height_factor=(1.0, 1.0), width_factor=(0.06, 0.12)\n        ),  # Time-masking: Randomly removes sections along the time axis\n        keras_cv.layers.RandomCutout(\n            height_factor=(0.06, 0.1), width_factor=(1.0, 1.0)\n        ),  # Frequency-masking: Randomly removes sections along the frequency axis\n    ]\n\n    def augment(img, label):\n        \"\"\"\n        Applies the augmentation pipeline to an image-label pair.\n\n        Args:\n            img (tf.Tensor): Input spectrogram image.\n            label (tf.Tensor): Corresponding label for the image.\n\n        Returns:\n            Augmented image and label.\n        \"\"\"\n\n        # Wrap image and label in a dictionary for compatibility with keras_cv augmenters\n        data = {\"images\": img, \"labels\": label}\n\n        # Apply augmentations with a 35% probability for each augmenter\n        for augmenter in augmenters:\n            if tf.random.uniform([]) < 0.35:\n                data = augmenter(data, training=True)\n\n        # Extract and return augmented image and label\n        return data[\"images\"], data[\"labels\"]\n\n    return augment","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T06:10:09.781633Z","iopub.execute_input":"2025-05-01T06:10:09.781960Z","iopub.status.idle":"2025-05-01T06:10:09.787694Z","shell.execute_reply.started":"2025-05-01T06:10:09.781936Z","shell.execute_reply":"2025-05-01T06:10:09.786650Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Build the data pipeline for training 💰\n","metadata":{}},{"cell_type":"code","source":"seed = 42\ndef build_dataset(\n    paths, \n    labels=None, \n    batch_size=32,\n    decode_fn=None, \n    augment_fn=None, \n    cache=True,\n    augment=False, \n    shuffle=2048\n):\n    \"\"\"\n    Builds a TensorFlow dataset pipeline for audio processing.\n\n    Args:\n        paths (list or tf.Tensor): List of file paths to audio files.\n        labels (list or tf.Tensor, optional): Corresponding labels for classification. Defaults to None.\n        batch_size (int, optional): Number of samples per batch. Defaults to 32.\n        decode_fn (function, optional): Function to decode audio files. Defaults to None.\n        augment_fn (function, optional): Function to apply augmentations. Defaults to None.\n        cache (bool, optional): Whether to cache the dataset in memory. Defaults to True.\n        augment (bool, optional): Whether to apply data augmentation. Defaults to False.\n        shuffle (int or bool, optional): Buffer size for shuffling. Set to False to disable shuffling. Defaults to 2048.\n\n    Returns:\n        tf.data.Dataset: Preprocessed dataset ready for training.\n    \"\"\"\n\n    # Use default decoder if none is provided\n    if decode_fn is None:\n        decode_fn = build_decoder(with_labels=(labels is not None), dim=audio_len)\n\n    # Use default augmentation function if none is provided\n    if augment_fn is None:\n        augment_fn = build_augmenter()\n\n    # Set automatic tuning for dataset performance optimization\n    AUTO = tf.data.experimental.AUTOTUNE\n\n    # Create dataset from file paths (with or without labels)\n    slices = (paths,) if labels is None else (paths, labels)\n    print(f\"Labels: {labels}\")\n    ds = tf.data.Dataset.from_tensor_slices(slices)\n\n    # Apply decoding function to process audio files\n    ds = ds.map(decode_fn, num_parallel_calls=AUTO)\n\n    # Cache dataset in memory to speed up subsequent iterations\n    if cache:\n        ds = ds.cache()\n\n    # Shuffle dataset if required\n    if shuffle:\n        opt = tf.data.Options()\n        ds = ds.shuffle(shuffle, seed=seed)  # Shuffle with seed for reproducibility\n        opt.experimental_deterministic = False  # Improve performance by allowing non-deterministic order\n        ds = ds.with_options(opt)\n\n    # Batch dataset with a fixed size, ensuring even batch sizes\n    ds = ds.batch(batch_size, drop_remainder=True)\n\n    # Apply augmentation if enabled\n    if augment:\n        ds = ds.map(augment_fn, num_parallel_calls=AUTO)\n\n    # Prefetch data to improve training performance\n    ds = ds.prefetch(AUTO)\n\n    return ds\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T06:10:13.956811Z","iopub.execute_input":"2025-05-01T06:10:13.957131Z","iopub.status.idle":"2025-05-01T06:10:13.963798Z","shell.execute_reply.started":"2025-05-01T06:10:13.957108Z","shell.execute_reply":"2025-05-01T06:10:13.962862Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Split the dataset to a test and train set 🚂\n***We used a test size of 0.2***","metadata":{}},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\n\n# First split: 80% train, 20% temp (stratified by original targets)\ntrain_df, temp_df = train_test_split(\n    df,\n    test_size=0.2,\n    stratify=df['target'],  # Use original targets\n    random_state=42\n)\n\n# Second split: 50% validation, 50% test (stratified by temp_df's targets)\nvalid_df, test_df = train_test_split(\n    temp_df,\n    test_size=0.5,\n    stratify=temp_df['target'],  # Critical: Use temp_df's targets!\n    random_state=42\n)\n\nprint(f\"Train: {len(train_df)} | Valid: {len(valid_df)} | Test: {len(test_df)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T06:10:18.068464Z","iopub.execute_input":"2025-05-01T06:10:18.068775Z","iopub.status.idle":"2025-05-01T06:10:18.308237Z","shell.execute_reply.started":"2025-05-01T06:10:18.068753Z","shell.execute_reply":"2025-05-01T06:10:18.307348Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Prepare training dataset\ntrain_paths = train_df.filepath.values  # Extract file paths from training DataFrame\ntrain_labels = train_df.target.values   # Extract corresponding labels\n\ntrain_ds = build_dataset(\n    paths=train_paths, \n    labels=train_labels, \n    batch_size=batch_size,\n    shuffle=True,  # Enable shuffling for training dataset\n    augment=True  # Apply augmentation for training dataset\n)\n\n# Prepare validation dataset\nvalid_paths = valid_df.filepath.values  # Extract file paths from validation DataFrame\nvalid_labels = valid_df.target.values   # Extract corresponding labels\n\nvalid_ds = build_dataset(\n    paths=valid_paths, \n    labels=valid_labels, \n    batch_size=batch_size,\n    shuffle=False,  # No shuffling for validation to ensure consistency\n    augment=False  # No augmentation for validation dataset\n)\n\n# Prepare test dataset\ntest_paths = test_df.filepath.values  # Extract file paths from test DataFrame\ntest_labels = test_df.target.values   # Extract corresponding labels\n\ntest_ds = build_dataset(\n    paths=test_paths, \n    labels=test_labels, \n    batch_size=1,\n    shuffle=False,  # No shuffling for test to ensure consistency\n    augment=False  # No augmentation for test dataset\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T06:10:23.197527Z","iopub.execute_input":"2025-05-01T06:10:23.197925Z","iopub.status.idle":"2025-05-01T06:10:29.049100Z","shell.execute_reply.started":"2025-05-01T06:10:23.197894Z","shell.execute_reply":"2025-05-01T06:10:29.048401Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## Show the shape of the spectrogram from train_ds\nx_train = next(iter(train_ds))[0]\nprint(x_train.shape)\n## The shape of the spectrogram is (64, 128, 384, 3) which means that there are 64 images in the batch,\n## each image has a height of 128 pixels, a width of 384 pixels and 3 channels (RGB)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T05:13:46.433255Z","iopub.execute_input":"2025-05-01T05:13:46.433633Z","iopub.status.idle":"2025-05-01T05:13:49.514885Z","shell.execute_reply.started":"2025-05-01T05:13:46.433611Z","shell.execute_reply":"2025-05-01T05:13:49.514108Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\ndef plot_batch(batch, row=3, col=3, label2name=None,):\n    \"\"\"Plot one batch data\"\"\"\n    if isinstance(batch, tuple) or isinstance(batch, list):\n        specs, tars = batch\n    else:\n        specs = batch\n        tars = None\n    plt.figure(figsize=(col*5, row*3))\n    for idx in range(row*col):\n        ax = plt.subplot(row, col, idx+1)\n        lid.specshow(np.array(specs[idx, ..., 0]), \n                     n_fft=nfft, \n                     hop_length=hop_length, \n                     sr=sample_rate,\n                     x_axis='time',\n                     y_axis='mel',\n                     cmap='coolwarm')\n        if tars is not None:\n            label = tars[idx].numpy().argmax()\n            name = label2name[label]\n            plt.title(name)\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T06:10:31.669250Z","iopub.execute_input":"2025-05-01T06:10:31.669965Z","iopub.status.idle":"2025-05-01T06:10:31.676129Z","shell.execute_reply.started":"2025-05-01T06:10:31.669933Z","shell.execute_reply":"2025-05-01T06:10:31.675130Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sample_ds = train_ds.take(10)\nbatch = next(iter(sample_ds))\nplot_batch(batch, label2name=label2name)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T05:13:49.522498Z","iopub.execute_input":"2025-05-01T05:13:49.522795Z","iopub.status.idle":"2025-05-01T05:13:54.034394Z","shell.execute_reply.started":"2025-05-01T05:13:49.522766Z","shell.execute_reply":"2025-05-01T05:13:54.033392Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## EfficientNetV2 Classifier","metadata":{}},{"cell_type":"code","source":"# Create an input layer for the model\ninp = keras.layers.Input(shape=(None, None, 3))\npreset = 'efficientnetv2_b2_imagenet'\n# Pretrained backbone\nbackbone = keras_cv.models.EfficientNetV2Backbone.from_preset(\n    preset,\n)\nout = keras_cv.models.ImageClassifier(\n    backbone=backbone,\n    num_classes=num_classes,\n    name=\"classifier\"\n)(inp)\n# Build model\nmodel = keras.models.Model(inputs=inp, outputs=out)\n# Compile model with optimizer, loss and metrics\nmodel.compile(optimizer=\"adam\",\n              loss=keras.losses.CategoricalCrossentropy(label_smoothing=0.02),\n                 metrics=[keras.metrics.AUC(name='auc'), \n                       # keras.metrics.CategoricalAccuracy(name='accuracy'), \n                       # keras.metrics.Precision(name='precision'), \n                       # keras.metrics.Recall(name='recall'), \n                       # keras.metrics.F1Score(name='f1_score')\n                        ]\n             )\nmodel.summary()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T05:13:54.035711Z","iopub.execute_input":"2025-05-01T05:13:54.036015Z","iopub.status.idle":"2025-05-01T05:13:59.455449Z","shell.execute_reply.started":"2025-05-01T05:13:54.035988Z","shell.execute_reply":"2025-05-01T05:13:59.454615Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Custom metrics extension","metadata":{}},{"cell_type":"code","source":"\nimport math\n\ndef get_lr_callback(batch_size=8, mode='cos', epochs=10, plot=False):\n    lr_start, lr_max, lr_min = 5e-5, 8e-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\n    return keras.callbacks.LearningRateScheduler(lrfn, verbose=False)  # Create lr callback","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T06:11:10.546985Z","iopub.execute_input":"2025-05-01T06:11:10.547297Z","iopub.status.idle":"2025-05-01T06:11:10.554391Z","shell.execute_reply.started":"2025-05-01T06:11:10.547274Z","shell.execute_reply":"2025-05-01T06:11:10.553530Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"lr_cb = get_lr_callback(batch_size, plot=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T05:13:59.464113Z","iopub.execute_input":"2025-05-01T05:13:59.464367Z","iopub.status.idle":"2025-05-01T05:13:59.704212Z","shell.execute_reply.started":"2025-05-01T05:13:59.464336Z","shell.execute_reply":"2025-05-01T05:13:59.703308Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ckpt_cb = keras.callbacks.ModelCheckpoint(\"efficientNet-method.weights.h5\",\n                                         monitor='val_auc',\n                                         save_best_only=True,\n                                         save_weights_only=True,\n                                         mode='max')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T05:13:59.705626Z","iopub.execute_input":"2025-05-01T05:13:59.705938Z","iopub.status.idle":"2025-05-01T05:13:59.709744Z","shell.execute_reply.started":"2025-05-01T05:13:59.705909Z","shell.execute_reply":"2025-05-01T05:13:59.708856Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Train the classifier for 10 epochs","metadata":{}},{"cell_type":"code","source":"epochs = 10\n# history = model.fit(\n#      train_ds, \n#     validation_data=valid_ds, \n#     epochs=epochs,\n#      callbacks=[\n#         lr_cb,\n#         ckpt_cb\n#     ],\n#     verbose=1\n# )\n\n# 3. Train with GPU\nwith tf.device('/GPU:0'):\n    history = model.fit(\n        train_ds,\n        validation_data=valid_ds,\n        epochs=10,\n        callbacks=[lr_cb, ckpt_cb],\n        verbose=1\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T05:13:59.710567Z","iopub.execute_input":"2025-05-01T05:13:59.710842Z","iopub.status.idle":"2025-05-01T05:50:31.315092Z","shell.execute_reply.started":"2025-05-01T05:13:59.710813Z","shell.execute_reply":"2025-05-01T05:50:31.313228Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Test on test split set","metadata":{}},{"cell_type":"code","source":"## Load the saved weights\nmodel.load_weights(\"efficientNet-method.weights.h5\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T05:50:31.329037Z","iopub.execute_input":"2025-05-01T05:50:31.329303Z","iopub.status.idle":"2025-05-01T05:50:32.576824Z","shell.execute_reply.started":"2025-05-01T05:50:31.329280Z","shell.execute_reply":"2025-05-01T05:50:32.576089Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Get prediction probabilities\ny_pred_proba = model.predict(test_ds, verbose=1)\n# Get predicted class labels\ny_pred = y_pred_proba.argmax(axis=1)\n# True labels\ny_true = test_df.target.values","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T05:50:32.577938Z","iopub.execute_input":"2025-05-01T05:50:32.578170Z","iopub.status.idle":"2025-05-01T05:52:04.744029Z","shell.execute_reply.started":"2025-05-01T05:50:32.578150Z","shell.execute_reply":"2025-05-01T05:52:04.743069Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(y_true.shape, y_pred.shape, y_pred_proba.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T05:52:04.745062Z","iopub.execute_input":"2025-05-01T05:52:04.745427Z","iopub.status.idle":"2025-05-01T05:52:04.750345Z","shell.execute_reply.started":"2025-05-01T05:52:04.745394Z","shell.execute_reply":"2025-05-01T05:52:04.749681Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\n\nunique_classes = np.unique(y_true)\nprint(f\"Unique classes in y_true: {len(unique_classes)} out of {y_pred_proba.shape[1]} total classes\")\n\nmissing_classes = set(range(y_pred_proba.shape[1])) - set(unique_classes)\nprint(f\"Missing classes: {missing_classes}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T05:52:04.751019Z","iopub.execute_input":"2025-05-01T05:52:04.751214Z","iopub.status.idle":"2025-05-01T05:52:04.788642Z","shell.execute_reply.started":"2025-05-01T05:52:04.751196Z","shell.execute_reply":"2025-05-01T05:52:04.787978Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# calculate the accuracy score\nfrom sklearn.metrics import accuracy_score, roc_auc_score, f1_score, precision_score,recall_score\nacc = accuracy_score(y_true, y_pred)\nprint(f\"Accuracy: {acc:.4f}\")\n\n## Calculate auc_roc score\nauc_roc = roc_auc_score(y_true, y_pred_proba, multi_class='ovr')\nprint(f\"AUC ROC: {auc_roc:.4f}\")\n\n## calculate the f1 score\nf1 = f1_score(y_true, y_pred, average='macro')\nprint(f\"F1 Score: {f1:.4f}\")\n## calculate the precision score\nprecision = precision_score(y_true, y_pred, average='macro')\nprint(f\"Precision: {precision:.4f}\")\n## calculate the recall score\nrecall = recall_score(y_true, y_pred, average='macro')\nprint(f\"Recall: {recall:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T05:52:04.789359Z","iopub.execute_input":"2025-05-01T05:52:04.789617Z","iopub.status.idle":"2025-05-01T05:52:04.973098Z","shell.execute_reply.started":"2025-05-01T05:52:04.789587Z","shell.execute_reply":"2025-05-01T05:52:04.972151Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Create another model with a different backbone","metadata":{}},{"cell_type":"code","source":"## Create another model with a different backbone\ninp = keras.layers.Input(shape=(None, None, 3))\npreset = 'resnet50_imagenet'\n# Pretrained backbone\nbackbone = keras_cv.models.ResNet50Backbone.from_preset(\n    preset,\n)\nout = keras_cv.models.ImageClassifier(\n    backbone=backbone,\n    num_classes=num_classes,\n    name=\"classifier\"\n)(inp)\n# Build model\nmodel = keras.models.Model(inputs=inp, outputs=out)\n# Compile model with optimizer, loss and metrics\nmodel.compile(optimizer=\"adam\",\n              loss=keras.losses.CategoricalCrossentropy(label_smoothing=0.02),\n                 metrics=[keras.metrics.AUC(name='auc')]\n             )\nmodel.summary()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T06:10:54.882639Z","iopub.execute_input":"2025-05-01T06:10:54.882967Z","iopub.status.idle":"2025-05-01T06:10:59.937113Z","shell.execute_reply.started":"2025-05-01T06:10:54.882944Z","shell.execute_reply":"2025-05-01T06:10:59.936394Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"lr_cb = get_lr_callback(batch_size, plot=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T06:11:16.070264Z","iopub.execute_input":"2025-05-01T06:11:16.070646Z","iopub.status.idle":"2025-05-01T06:11:16.276420Z","shell.execute_reply.started":"2025-05-01T06:11:16.070615Z","shell.execute_reply":"2025-05-01T06:11:16.275544Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Overwrite previous preset to Resnet50","metadata":{}},{"cell_type":"code","source":"ckpt_cb = keras.callbacks.ModelCheckpoint(\"resnet50-method.weights.h5\",\n                                         monitor='val_auc',\n                                         save_best_only=True,\n                                         save_weights_only=True,\n                                         mode='max')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T06:11:19.367664Z","iopub.execute_input":"2025-05-01T06:11:19.367972Z","iopub.status.idle":"2025-05-01T06:11:19.372193Z","shell.execute_reply.started":"2025-05-01T06:11:19.367948Z","shell.execute_reply":"2025-05-01T06:11:19.371163Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## Train the model for 10 epochs\nepochs = 10\n# history = model.fit(\n#      train_ds, \n#     validation_data=valid_ds, \n#     epochs=epochs,\n#     callbacks=[lr_cb, ckpt_cb], \n#     verbose=1\n# )\nwith tf.device('/GPU:0'):\n    history = model.fit(\n        train_ds,\n        validation_data=valid_ds,\n        epochs=10,\n        callbacks=[lr_cb, ckpt_cb],\n        verbose=1\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T06:11:22.094700Z","iopub.execute_input":"2025-05-01T06:11:22.095058Z","iopub.status.idle":"2025-05-01T06:49:30.117449Z","shell.execute_reply.started":"2025-05-01T06:11:22.095027Z","shell.execute_reply":"2025-05-01T06:49:30.116412Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Load saved model","metadata":{}},{"cell_type":"code","source":"## Load the resnet model weights\nmodel.load_weights(\"resnet50-method.weights.h5\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T06:52:31.245171Z","iopub.execute_input":"2025-05-01T06:52:31.245615Z","iopub.status.idle":"2025-05-01T06:52:32.106495Z","shell.execute_reply.started":"2025-05-01T06:52:31.245580Z","shell.execute_reply":"2025-05-01T06:52:32.105503Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Test on test data","metadata":{}},{"cell_type":"code","source":"# Get prediction probabilities\ny_pred_proba = model.predict(test_ds, verbose=1)\n# Get predicted class labels\ny_pred = y_pred_proba.argmax(axis=1)\n# True labels\ny_true = test_df.target.values","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T06:52:35.485022Z","iopub.execute_input":"2025-05-01T06:52:35.485386Z","iopub.status.idle":"2025-05-01T06:54:02.682420Z","shell.execute_reply.started":"2025-05-01T06:52:35.485351Z","shell.execute_reply":"2025-05-01T06:54:02.681512Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# calculate the accuracy score\nfrom sklearn.metrics import accuracy_score, roc_auc_score, f1_score, precision_score,recall_score\nacc = accuracy_score(y_true, y_pred)\nprint(f\"Accuracy: {acc:.4f}\")\n\n## Calculate auc_roc score\nauc_roc = roc_auc_score(y_true, y_pred_proba, multi_class='ovr')\nprint(f\"AUC ROC: {auc_roc:.4f}\")\n\n## calculate the f1 score\nf1 = f1_score(y_true, y_pred, average='macro')\nprint(f\"F1 Score: {f1:.4f}\")\n## calculate the precision score\nprecision = precision_score(y_true, y_pred, average='macro')\nprint(f\"Precision: {precision:.4f}\")\n## calculate the recall score\nrecall = recall_score(y_true, y_pred, average='macro')\nprint(f\"Recall: {recall:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-01T06:54:02.683608Z","iopub.execute_input":"2025-05-01T06:54:02.683866Z","iopub.status.idle":"2025-05-01T06:54:02.890215Z","shell.execute_reply.started":"2025-05-01T06:54:02.683844Z","shell.execute_reply":"2025-05-01T06:54:02.889211Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}