{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":91844,"databundleVersionId":11361821,"sourceType":"competition"},{"sourceId":11652644,"sourceType":"datasetVersion","datasetId":7312805},{"sourceId":11865265,"sourceType":"datasetVersion","datasetId":7455967}],"dockerImageVersionId":31012,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"### Section 0: Notebook configuration\n\nThis notebook section is used to configure the notebook environment and import base packages.","metadata":{}},{"cell_type":"code","source":"!pip install --upgrade librosa","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Imports\nimport os\nimport shutil\nimport gc\nimport re\nfrom typing import Union\nimport wandb\nfrom kaggle_secrets import UserSecretsClient\nimport itertools\nimport math\nimport random","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Login to WandB\nuser_secrets = UserSecretsClient()\nwandb_api_key = user_secrets.get_secret(\"WANDB_API_KEY\")\nwandb.login(key=wandb_api_key)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Section 1: Audio Data Loading\n\nThe first notebook section is used to load raw audio data. The `load_training_audio` method accepts an optional `classes` parameter which corresponds to audio sub-directories. If the `classes` parameter is specified, audio is only loaded for the specified classes.","metadata":{}},{"cell_type":"code","source":"# Imports\nimport pandas as pd\nimport torch\nimport torchaudio","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Config\nSAVE_PATH = \"/kaggle/working\"\nPROCESSED_DATA_PATH = SAVE_PATH + \"/\" + \"processed_audio\"\nTRAINING_DATA_PATH = \"/kaggle/input/birdclef-2025/train_audio\"\nTRAINING_DATA_METADATA_FILE = \"/kaggle/input/birdclef-2025/train.csv\"\nEXT = \".ogg\"\nMAX_FILES = 1000 # Maximum number of files to load per class; set to arbitrarily large number to load all files\nGBVC_LABELS_FILE = \"/kaggle/input/google-bird-vocalization-labels/google_bird_vocalization_labels.csv\"\nSR = 32000\nLOAD_SLICE = 8 * SR","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_training_metadata(train_index_file: str) -> pd.DataFrame:\n    \"\"\"\n    This function loads the training data index file.\n\n    @param {str} train_index_file: The full path expression for the CSV training data index file.\n    @return {pd.DataFrame} train_index: The training data index file data as a Pandas DataFrame.\n    @raise {Exception} e\n    \"\"\"\n    try:\n        return pd.read_csv(train_index_file)\n    except Exception as e:\n        print(str(e))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def extract_classes_not_in_gbvc(bc_labels: list, gbvc_labels_file: str) -> list:\n    \"\"\"\n    Returns the list of BirdCLEF+ 2025 classes that are not part of the Google Bird Vocalization (GBV) classifier.\n\n    @param {list} bc_labels: List of unique labels in the BirdCLEF+ 2025 dataset.\n    @param {str} gbvc_labels_files: The labels file for the GBV classifier.\n    @return {list} missing_classes: The BirdCLEF+ 2025 classes not covered by the GBV classifier.\n    \"\"\"\n    gbvc_labels = pd.read_csv(GBVC_LABELS_FILE)[\"ebird2021\"].unique()\n    return [bc_label for bc_label in bc_labels if bc_label not in gbvc_labels]","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_training_audio(train_data_path: str, make_copy: bool, processed_data_path: str, classes: list = [], use_slice: bool = False) -> Union[dict, str]:\n    \"\"\"\n    Loads training audio data and generates statistics for loaded data.\n\n    @param {str} train_data_path: The full path to training data sub-directories.\n    @param {bool} make_copy: If `True`, a copy of audio data is saved to `processed_data_path` under the appropriate sub-directory (i.e. class name)\n    @param {str} processed_data_path: The full path to processed audio data sub-directories; used in later notebook sections.\n    @param {list} classes: Optional; if supplied, only classes in list are loaded.\n    @param {bool} use_slice: Optional; if `True`, audio is sliced using LOAD_SLICE constant.\n    @return {Union[dict, str]} _audio, _stats\n    @raise {Exception} e\n    \"\"\"\n    try:\n        _audio = {}\n        _stats = \"class,sampling_rate,num_files_loaded,num_secs_loaded,num_files_loaded\\n\"\n        _dirs = os.listdir(train_data_path)\n\n        if make_copy:\n            if not os.path.exists(processed_data_path):\n                os.mkdir(processed_data_path)\n        \n        if 0 < len(classes):\n            _dirs = classes\n\n        for _dir in _dirs:\n            if make_copy:\n                if not os.path.exists(processed_data_path + \"/\" + _dir):\n                    os.mkdir(processed_data_path + \"/\" + _dir)\n            \n            _files = os.listdir(train_data_path + \"/\" + _dir)\n            _audio[_dir] = []\n            _file_count = 0\n            _secs = 0\n            try:\n                for _file in _files:\n                    # Manual removal of certain files\n                    if \"52884\" == _dir and \"CSA18980\" == _file:\n                        continue\n                    if \"whfant1\" == _dir and \"iNat192767.ogg\" == _file:\n                        continue\n                    if \"gycwor1\" == _dir and \"iNat589108.ogg\" == _file:\n                        continue    \n                    \n                    if _file_count < MAX_FILES:\n                        # Load audio\n                        _audio_tensor, _sampling_rate = read_audio_data(train_data_path + \"/\" + _dir + \"/\" + _file)\n                        \n                        if use_slice:\n                            if len(_audio_tensor[0]) > LOAD_SLICE:\n                                _audio_tensor = _audio_tensor[:,:LOAD_SLICE] # Return all rows of tensor, but only columns up to SLICE length\n                        \n                        if make_copy:\n                            # Save a copy of this audio to `processed_data_path` as `.wav` file\n                            _save_file = re.sub(r\"\\.[a-z]{3}$\", \".wav\", _file)\n                            torchaudio.save(f\"{processed_data_path}/{_dir}/{_save_file}\", _audio_tensor, SR)\n                        \n                        _audio[_dir].append((_file, _audio_tensor))\n                        _secs += math.floor(len(_audio_tensor[0])/SR)\n                        _file_count += 1\n                    else:\n                        break\n            except:\n                # Couldn't process audio file; continue loop\n                print(f\"couldn't process {_dir + '/' + _file}\")\n            _stats += f\"{_dir},{_sampling_rate},{len(_files)},{_secs},{_file_count}\\n\"\n        return _audio, _stats\n    except Exception as e:\n        print(str(e))\n\ndef read_audio_data(file: str) -> torch.Tensor:\n    _audio_tensor, _sampling_rate = torchaudio.load(file, normalize=True)\n    return _audio_tensor, _sampling_rate","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Extract BirdCLEF+ 2025 classes that aren't covered by GBV classifier\nbc_labels = load_training_metadata(TRAINING_DATA_METADATA_FILE)[\"primary_label\"].unique()\nmissing_classes = extract_classes_not_in_gbvc(bc_labels, GBVC_LABELS_FILE)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"### DEBUG/OUTPUT CELL ###\nassert 63 == len(missing_classes)\nprint(missing_classes)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load audio and generate stats\naudio_by_class, stats = load_training_audio(TRAINING_DATA_PATH, True, PROCESSED_DATA_PATH, missing_classes, False)\n# audio_by_class, _ = load_training_audio(TRAINING_DATA_PATH, True, PROCESSED_DATA_PATH, [], False)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"### DEBUG/OUTPUT CELL ###\naudio_by_class_slice = dict(itertools.islice(audio_by_class.items(), 3))\nprint(audio_by_class_slice)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"### DEBUG/OUTPUT CELL ###\nprint(stats)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Section 2: Audio Data Processing\n\nThe second notebook section is used to process the raw audio data.\n\n1. Silent segments are stripped. This eliminates irrelevant parts of the audio signal. Since there tend to be many silent segments at the beginning of audio samples, this processing helps to ensure relevant audio occurs within the first few seconds of each audio sample.\n2. Audio for minority classes is augmented to help address the class imbalance. (The class imbalance is also addressed later through the choice of the loss function.) Audio augmentation consists of adding a randomly generated noise signal to raw audio samples in order to generate new audio samples.","metadata":{}},{"cell_type":"code","source":"# Imports\nimport numpy as np\nimport torchaudio.transforms as T","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Config\nUSE_REMOVE_SILENCE_AND_HUMAN_ANNOT = True\nUSE_AUGMENTATIONS = True\n\n# Silence processing constants\nSIL_FRAME_PCT_OF_SR = 0.25\nSIL_FRAME = int(SR * SIL_FRAME_PCT_OF_SR)\nSIL_HOP = int(1.0 * SIL_FRAME)\nSIL_THRESHOLD = 5e-5\nSIL_REPLACE_VAL = -1000 # Value used to replace audio signal values within silent segments\n\n# The `ANNOT_BREAKPOINT` constant is the num of seconds in each audio sample after which check for silent segments is\n# stopped; only applies to those samples with at least one silent segment *before* the breakpoint. \n# So, audio samples without any silent segments are kept in their entirety.\nANNOT_BREAKPOINT = 10\n\n# The `SLICE_FRAME` constant is used to return a specific number of seconds of\n# audio starting at the beginning of the audio from a sample in the \n# `remove_silence_and_human_annot` function.\nSLICE_FRAME = 8\n\nTEMPO_RNG_LOW = 0.5\nTEMPO_RNG_HIGH = 1.5\n\nNOISE_RNG_LOW = 0.0001\nNOISE_RNG_HIGH = 0.0009","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def detect_silence(audio_by_class: dict, threshold: float = SIL_THRESHOLD, frame_length: int = SIL_FRAME, hop_length: int = SIL_HOP) -> dict:\n    \"\"\"\n    Returns dictionary of detected silence segments for each waveform in each class.\n\n    @param {dict} audio_by_class: Dictionary of waveforms by class.\n    @param {float} threshold: RMS energy threshold below which it's considered silence.\n    @param {int} frame_length: Length of the window for analysis.\n    @param {int} hop_length: Step size between windows.\n    @return {dict} silent_segments: Dictionary of silent segments, using class names as keys. \n    \"\"\"\n    silent_segments = {}\n    for class_name in audio_by_class:\n        silent_segments[class_name] = {}\n        for tup in audio_by_class[class_name]:\n            file = tup[0]\n            wf = tup[1]\n            npwf = wf.numpy()[0]\n            silent_segments[class_name][file] = []\n            for i in range(0, len(npwf) - frame_length, hop_length):\n                frame = npwf[i:i + frame_length]\n                rms = np.sqrt(np.mean(frame**2)) # Calculate RMS\n\n                if rms < threshold:\n                    silent_segments[class_name][file].append(i)\n    return silent_segments","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def remove_silence_and_human_annot(audio_by_class: dict, silent_segments: dict, processed_data_path: str, save_audio: bool = True, only_return_slice: bool = True, test_num_classes: int = None) -> dict:\n    \"\"\"\n    Strips silent segments from audio samples. Audio samples with at least one silent segment before the breakpoint \n    specified by `ANNOT_BREAKPOINT` are sliced at the last silent segment and only the beginning slice is kept. This is a\n    simple approach to remove human annotations which tend to occur after a short silent segment that follows a birdsong\n    recording. Audio samples with no silent segments are kept in their entirety. Processed audio is saved to disk.\n\n    @param {dict} audio_by_class: Dictionary of raw audio samples.\n    @param {dict} silent_segments: Dictionary of silent segments per audio sample per class.\n    @param {str} processed_data_path: Full path to processed audio data directores.\n    @param {bool} save_audio: If `True`, processed audio is saved to disk.\n    @param {bool} only_return_slice: If `True`, only `SLICE_FRAME` seconds starting from the beginning of each audio sample is saved after processing.\n    @param {int} test_num_classes: Optional; if specified, only the first `test_num_classes` are processed.\n    @return {dict} total_secs_audio_per_class: Dictionary to total number of seconds of audio per class after processing.\n    \"\"\"\n    # processed_audio = {}\n    total_secs_audio_per_class = {}\n    total_secs = 0\n    num_samples = 0\n    \n    if None != test_num_classes:\n        silent_segments = dict(itertools.islice(silent_segments.items(), test_num_classes))\n    \n    for class_name in silent_segments:\n        # processed_audio[class_name] = {}        \n        \n        audio_segments = silent_segments[class_name]\n        audio_secs_for_class = 0\n        for i, file in enumerate(audio_segments):\n            segments = audio_segments[file]\n            audio = audio_by_class[class_name][i][1].numpy()[0].tolist()        \n            if 0 == len(segments):\n                # Do not process this audio\n                pass\n            elif 1 == len(segments) and 0 == segments[0]:\n                # Do not process this audio\n                pass\n            else:    \n                last_segment = len(segments) - 1\n                REPL_VAL = SIL_REPLACE_VAL\n                for segment_start in segments:\n                    if SR * ANNOT_BREAKPOINT > segment_start and segments.index(segment_start) != last_segment:\n                        audio[segment_start:segment_start + SIL_FRAME + 1] = [REPL_VAL] * SIL_FRAME\n                    else:\n                        # This segment exists outside the 10s window or is the last segment\n                        if SR > segment_start:\n                            # This is a last segment that occurs less than 1 second into the audio.\n                            # Let's take a guess and just take first `SLICE_FRAME` seconds of audio.\n                            audio = audio[:SLICE_FRAME * SR]\n                        else:\n                            # This segment occurs outside the 10s window.\n                            # Grab everything up to this segment and stop processing.\n                            audio = audio[:segment_start]\n                            break # Don't perform any more processing on this audio            \n                audio = [v for v in audio if v != -1000]\n                \n            if only_return_slice:\n                if  SR * SLICE_FRAME < len(audio):\n                    audio = audio[:SR * SLICE_FRAME]\n            \n            # processed_audio[class_name][file] = audio\n            if save_audio:\n                save_file = re.sub(r\"\\.[a-z]{3}$\", \".wav\", file)\n                torchaudio.save(f\"{processed_data_path}/{class_name}/{save_file}\", torch.from_numpy(np.array([audio])), SR)\n            \n            audio_secs_for_class += len(audio)/SR\n            num_samples += 1\n        \n        total_secs_audio_per_class[class_name] = math.floor(audio_secs_for_class)\n        total_secs += audio_secs_for_class\n\n    avg_secs = total_secs/num_samples\n    return total_secs_audio_per_class, math.floor(avg_secs)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_audio_stats(total_secs_audio_per_class: dict, audio_by_class: dict) -> dict:\n    \"\"\"\n    Utility function to get audio statistics. If `total_secs_audio_per_class` is \n    supplied, that dictionary is used to calculate audio statistics. Otherwise,\n    the dictionary `audio_by_class` is used to calculate audio statistics.\n\n    @param {dict} total_secs_audio_per_class: Dictionary with class names as keys and total seconds of audio per class as values.\n    @param {dict} audio_by_class: Dictionary of raw audio samples.\n    @return {Union[dict, int, int, int]|Union[int, int]}\n    \"\"\"\n    USE_RAW_AUDIO = False\n    total_secs = 0\n    max_secs = 0\n    num_samples = 0\n\n    if None == total_secs_audio_per_class:\n        total_secs_audio_per_class = {}\n        USE_RAW_AUDIO = True\n\n    if USE_RAW_AUDIO and None != audio_by_class:\n        for class_name in audio_by_class:\n            audio_secs_for_class = 0\n            for _, audio in audio_by_class[class_name]:\n                secs = len(audio[0])\n                audio_secs_for_class += secs/SR\n                num_samples += 1\n            total_secs_audio_per_class[class_name] = audio_secs_for_class\n            total_secs += audio_secs_for_class\n            if max_secs < audio_secs_for_class:\n                max_secs = audio_secs_for_class\n        avg_secs = total_secs/num_samples\n        avg_secs_per_class = total_secs/len(audio_by_class)\n        return total_secs_audio_per_class, math.floor(max_secs), math.floor(avg_secs), math.floor(avg_secs_per_class)\n    else:\n        for class_name in total_secs_audio_per_class:\n            audio_secs_for_class = total_secs_audio_per_class[class_name]\n            total_secs += audio_secs_for_class\n            if max_secs < audio_secs_for_class:\n                max_secs = audio_secs_for_class\n            avg_secs_per_class = total_secs/len(total_secs_audio_per_class)\n        return math.floor(max_secs), math.floor(avg_secs_per_class)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if USE_REMOVE_SILENCE_AND_HUMAN_ANNOT:\n    # Generate lists of silent segments for each audio sample in each class\n    silent_segments = detect_silence(audio_by_class)\n    \n    # Remove silence and most human annotations\n    total_secs_audio_per_class, avg_secs = remove_silence_and_human_annot(audio_by_class, silent_segments, PROCESSED_DATA_PATH, True, True, None)\n\n    # Get audio stats\n    max_secs, avg_secs_per_class = get_audio_stats(total_secs_audio_per_class, None)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"### DEBUG/OUTPUT CELL ###\nif USE_REMOVE_SILENCE_AND_HUMAN_ANNOT:\n    print(dict(itertools.islice(silent_segments.items(), 5)))\n    print(total_secs_audio_per_class)\n    print(avg_secs)\n    print(max_secs)\n    print(avg_secs_per_class)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_augmentation_turns_per_class_using_avg(total_secs_audio_per_class: dict, avg_secs_per_class: int, avg_secs: int) -> dict:\n    \"\"\"\n    Calculates number of augmentation turns for classes with total seconds of audio\n    below average.\n\n    @param {dict} total_secs_audio_per_class\n    @param {int} avg_secs_per_class: The average number of seconds per class.\n    @param {int} avg_secs: The average number of seconds per audio sample.\n    @return {dict} turns\n    \"\"\"\n    turns = {}\n    AVG_SECS_FACTOR = 1\n    for class_name in total_secs_audio_per_class:\n        if (avg_secs_per_class * AVG_SECS_FACTOR) > total_secs_audio_per_class[class_name]:\n            # This class has total seconds of audio below the average\n            num_turns = math.floor(((avg_secs_per_class * AVG_SECS_FACTOR) - total_secs_audio_per_class[class_name])/avg_secs)\n            turns[class_name] = num_turns\n        else:\n            # This class has total seconds of audio above the average\n            turns[class_name] = 0\n    return turns","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_augmentation_turns_per_class_using_max(total_secs_audio_per_class: dict, max_secs_audio: int, avg_secs_audio: int) -> dict:\n    \"\"\"\n    Calculates number of augmentation turns for classes with total seconds of audio\n    below max.\n\n    @param {dict} total_secs_audio_per_class\n    @param {int} max_secs_audio\n    @return {dict} turns\n    \"\"\"\n    turns = {}\n    MAX_SECS_FACTOR = 1\n    for class_name in total_secs_audio_per_class:\n        if total_secs_audio_per_class[class_name] < (max_secs_audio * MAX_SECS_FACTOR):\n            num_turns = math.ceil((max_secs_audio - total_secs_audio_per_class[class_name])/avg_secs_audio)\n            turns[class_name] = num_turns\n        else:\n            turns[class_name] = 0\n    return turns","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def add_noise(waveform: torch.Tensor, noise_factor: Union[float|str] = \"rand\") -> torch.Tensor:\n    \"\"\"\n    Adds random noise signal to audio waveform.\n\n    @param {torch.Tensor} waveform: The audio waveform.\n    @param {Union[float|str]} noise_factor: Factor used to generate the noise signal.\n                                            Defaults to `rand` which selects a noise\n                                            factor between `NOISE_RNG_LOW` and\n                                            `NOISE_RNG_HIGH`. Alternatively, pass a \n                                            float value.\n    @return {torch.Tensor} augmented_waveform: The augmented waveform.\n    @raise {Exception} e\n    \"\"\"\n    try:\n        if isinstance(noise_factor, str) and \"rand\" == noise_factor:\n            noise_factor = round(random.uniform(NOISE_RNG_LOW, NOISE_RNG_HIGH), 1)\n        \n        waveform_clone = waveform.clone()\n        noise_signal = torch.randn_like(waveform_clone) * noise_factor\n        augmented_waveform = (waveform_clone + noise_signal)\n        return augmented_waveform\n    except Exception as e:\n        print(str(e))\n\ndef change_tempo(waveform: torch.Tensor, sample_rate: int, tempo_ratio: Union[float|str] = \"rand\"):\n    \"\"\"\n    Changes the tempo of an audio waveform using torchaudio. \n\n    @param {torch.Tensor} waveform: The waveform to perturb.\n    @param {int} sample_rate: The sample rate for the waveform.\n    @param {Union[float|str]} tempo_ratio: Tempo factor for perturbation; defaults to \"rand\" which\n                                           randomly selects a tempo factor between `TEMPO_RNG_LOW`\n                                           and `TEMPO_RNG_HIGH`. Alternatively, pass a float value.\n    @return {torch.Tensor} perturbed_waveform: The speed-adjusted waveform.\n    @raise {Exception} e\n    \"\"\"\n    try:\n        if isinstance(tempo_ratio, str) and \"rand\" == tempo_ratio:\n            tempo_ratio = round(random.uniform(TEMPO_RNG_LOW, TEMPO_RNG_HIGH), 1)\n        \n        # Create the tempo transform\n        tempo_transformer = T.SpeedPerturbation(\n            orig_freq=sample_rate,\n            factors=[tempo_ratio]\n        )\n\n        # Apply the tempo change\n        with torch.inference_mode(): # Reduce memory usage and disable autograd\n            perturbed_waveform = tempo_transformer(waveform) # Returns a tuple\n            perturbed_waveform = perturbed_waveform[0]\n        return perturbed_waveform\n    except Exception as e:\n        print(str(e))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def run_augmentations(audio_by_class: dict, turns: dict, processed_data_path: str, save_audio: bool = True, processed_audio: dict = None) -> None:\n    \"\"\"\n    Utility function to run audio augmentations.\n\n    @param {dict} audio_by_class: Dictionary of raw audio with classes as keys.\n    @param {dict} turns: Number of augmentation turns by class.\n    @param {str} processed_data_path: Full path to processed audio data directores.\n    @param {bool} save_audio: If `True`, saves audio to disk.\n    @param {dict} processed_audio: Optional; if dictionary is supplied, augmented audio is added to dictionary.\n    @return {void|dict}\n    \"\"\"\n    for class_name in turns:\n        num_turns = turns[class_name]\n        if 0 == num_turns:\n            continue\n        else:\n            audios = audio_by_class[class_name]\n            print(f\"before augmentation {class_name} has {len(audios)} samples\")\n            \n            for t in range(0, num_turns):\n                rand_tup = audios[random.randrange(0, len(audios))]\n                rand_file = re.sub(r\"\\.[a-z]{3}$\", \".wav\", rand_tup[0])\n                                \n                # Load the processed audio for this class and file\n                rand_audio, _ = read_audio_data(processed_data_path + \"/\" + class_name + \"/\" + rand_file)\n\n                # Randomly select integer to decide which augmentations to perform for this waveform\n                # 0 = Add noise only\n                # 1 = Change tempo only\n                # 2 = Add noise and change tempo\n                aug_choice = random.randint(0, 2)                \n                \n                # Perform augmentations\n                if 0 == aug_choice:\n                    augmented_audio = add_noise(rand_audio, \"rand\")\n                elif 1 == aug_choice:\n                    augmented_audio = change_tempo(rand_audio, SR, \"rand\")\n                else:\n                    augmented_audio = add_noise(rand_audio, \"rand\")\n                    augmented_audio = change_tempo(augmented_audio, SR, \"rand\")\n                \n                if save_audio:\n                    file_stem = re.sub(r\"\\.[a-z]{3}$\", \"\", rand_file)\n                    save_file = f\"{file_stem}_aug_{t}\"\n                    torchaudio.save(f\"{processed_data_path}/{class_name}/{save_file}.wav\", augmented_audio, SR)\n\n                if None != processed_audio:\n                    processed_audio[class_name][save_filename] = augmented_audio\n\n            print(f\"after augmentation {class_name} has {len(audios) + num_turns} samples\")\n    \n    if None != processed_audio:\n        return processed_audio","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if USE_AUGMENTATIONS:\n    if not USE_REMOVE_SILENCE_AND_HUMAN_ANNOT:\n        total_secs_audio_per_class, max_secs, avg_secs, avg_secs_per_class = get_audio_stats(audio_by_class)\n\n    ### DEBUG ###\n    print(total_secs_audio_per_class)\n    print(max_secs)\n    print(avg_secs)\n    print(avg_secs_per_class)\n    ### DEBUG ###\n    \n    # Get augmentation turns per class\n    turns = get_augmentation_turns_per_class_using_avg(total_secs_audio_per_class, avg_secs_per_class, avg_secs)\n    # turns = get_augmentation_turns_per_class_using_max(total_secs_audio_per_class, max_secs, avg_secs)\n    \n    # Run augmentations\n    run_augmentations(audio_by_class, turns, PROCESSED_DATA_PATH, True, None)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"### DEBUG/OUTPUT CELL ###\nif USE_AUGMENTATIONS:\n    print(turns)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Clean up\ndel audio_by_class\ngc.collect()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"processed_audio_by_class, _ = load_training_audio(PROCESSED_DATA_PATH, False, None, [], False)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Section 3: Mel Spectrogram Generation and Input Preparation\n\nThe third notebook section is used to:\n\n1. Split processed audio data into training and validation lists.\n2. Split audio into 5 second frames.\n3. Generate mel spectrograms for each 5 second audio frame.\n4. Resize mel spectrograms to a target size of `(224, 224)`.\n5. Optionally load pseudo-labeled data samples to augment training data.\n6. One-hot encode training data and validation data labels.\n7. Construct TensorFlow `Dataset` objects from training and validation data lists.\n8. Optionally use MixUp logic to augment training data.","metadata":{}},{"cell_type":"code","source":"# Imports\nimport librosa\nimport librosa.display\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nimport keras\nfrom keras import layers\nimport tensorflow as tf\nfrom tensorflow import data as tf_data\nfrom tensorflow.random import gamma as tf_random_gamma","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Config\nSPLIT = 0.20\nFRAME_LENGTH = 5\nFRAME_STEP = 5\n\n# Mel spectrogram parameters\nN_FFT = 1024  # FFT size\nHOP_SIZE = 256\nN_MELS = 256\nFMIN = 1000  # minimum frequency\nFMAX = 10000 # maximum frequency\n\n# Other\nIMG_SIZE = 224  # target image shape\nBATCH_SIZE = 8\nUSE_MIXUP = True\n\nUSE_PSEUDO_LABELS = False\nPSEUDO_LABELS_DATA_PATH = \"/kaggle/input/birdclef-2025-pseudo-labels\"\nMAX_PSEUDO_LABELS_FILES = 5","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def generate_training_and_validation_data_lists(audio_by_class: dict) -> Union[list[torch.Tensor], list[int], list[torch.Tensor], list[int]]:\n    \"\"\"\n    Splits processed audio data into training and validation data lists.\n\n    @param {dict} audio_by_class: Dictionary of audio data using classes as keys.\n    @return {Union[list, list, list, list} training_audio, training_labels, validation_audio, validation_labels\n    \"\"\"\n    \n    training_audio = []\n    training_labels = []\n    validation_audio = []\n    validation_labels = []\n\n    for class_name in audio_by_class:\n        class_samples = len(audio_by_class[class_name])\n        num_val = math.ceil(class_samples * SPLIT)\n        num_train = class_samples - num_val\n\n        assert(num_train + num_val == class_samples)\n\n        # Shuffle audio for this class before split\n        random.shuffle(audio_by_class[class_name])\n        \n        train_tups = audio_by_class[class_name][:-num_val]\n        val_tups = audio_by_class[class_name][-num_val:]\n\n        class_train_audio = []\n        for tup in train_tups:\n            class_train_audio.append(tup[1])\n        \n        class_train_labels = [class_name] * num_train\n        \n        training_audio = training_audio + class_train_audio\n        training_labels = training_labels + class_train_labels\n        \n        if 0 != num_val:\n            class_validation_audio = []\n            for tup in val_tups:\n                class_validation_audio.append(tup[1])\n            \n            class_validation_labels = [class_name] * num_val\n            validation_audio = validation_audio + class_validation_audio\n            validation_labels = validation_labels + class_validation_labels\n    \n    return training_audio, training_labels, validation_audio, validation_labels\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Generate training data and validation data lists\ntraining_audio, training_labels, validation_audio, validation_labels = generate_training_and_validation_data_lists(processed_audio_by_class)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"### DEBUG/OUTPUT CELL ###\nprint(len(training_audio))\nprint(len(training_labels))\nprint(len(validation_audio))\nprint(len(validation_labels))\n\nprint(training_audio[0])\nprint(training_labels[0])\nprint(validation_audio[0])\nprint(validation_labels[0])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Convert audio to NumPy arrays\ntraining_audio = [audio.numpy() for audio in training_audio]\nvalidation_audio = [audio.numpy() for audio in validation_audio]","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"### DEBUG/OUTPUT CELL ###\n# Check list lengths\nprint(len(training_audio))\nprint(len(validation_audio))\n\n# Print one training data example and one validation data example\nprint(training_audio[0])\nprint(validation_audio[0])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def frame_audio(audio_data: np.ndarray, frame_length: int = FRAME_LENGTH, frame_step: int = FRAME_STEP, sample_rate = SR) -> list[tf.Tensor]:\n    \"\"\"\n    Returns an audio signal as a set of frames based on supplied frame length and frame step.\n\n    @param {np.ndarray} audio_data: The raw audio signal data.\n    @param {float} frame_length: The frame length expressed in seconds.\n    @param {float} hop_size: The frame step expressed in seconds.\n    @param {int} sample_rate: The sampling rate expressed in Hz.\n    @return {tf.Tensor} framed_audio: Set of audio frames derived from original audio signal data.\n    \"\"\"\n    if frame_length is None or frame_length < 0:\n        return audio_data[np.newaxis, :]\n    length = int(frame_length * sample_rate)\n    step = int(frame_step * sample_rate)\n    frames = tf.signal.frame(audio_data, length, step, pad_end=True, pad_value=0, axis=-1)\n    return frames","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Generate audio frames\ntraining_frames = []\ntraining_frames_labels = []\nvalidation_frames = []\nvalidation_frames_labels = []\n\nfor i in range(len(training_audio)):\n    label = training_labels[i]\n    frames = frame_audio(training_audio[i], FRAME_LENGTH, FRAME_STEP, SR)\n    np_frames = [frame.numpy() for frame in frames[0]]\n    training_frames = training_frames + np_frames\n    training_frames_labels = training_frames_labels + ([label] * len(np_frames))\n\nfor j in range(len(validation_audio)):\n    label = validation_labels[j]\n    frames = frame_audio(validation_audio[j], FRAME_LENGTH, FRAME_STEP, SR)\n    np_frames = [frame.numpy() for frame in frames[0]]\n    validation_frames = validation_frames + np_frames\n    validation_frames_labels = validation_frames_labels + ([label] * len(np_frames))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"### DEBUG/OUTPUT CELL ###\nprint(len(training_frames))\nprint(len(training_frames_labels))\nprint(len(validation_frames))\nprint(len(validation_frames_labels))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"### CLEAN UP ###\ndel processed_audio_by_class\ndel training_audio\ndel training_labels\ndel validation_audio\ndel validation_labels\ngc.collect()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def audio2melspec(audio_data):\n    \"\"\"\n    Convert the raw audio data into a normalized mel spectrogram.\n    \n    @param {torch.tensor} audio_data: The audio samples as a Tensor object.\n    @return {np.ndarray} mel_spec_norm: Normalized mel spectrogram.\n    \"\"\"\n    if np.isnan(audio_data).any():\n        mean_signal = np.nanmean(audio_data)\n        audio_data = np.nan_to_num(audio_data, nan=mean_signal)\n    \n    mel_spec = librosa.feature.melspectrogram(\n        y=audio_data,\n        sr=SR,\n        n_fft=N_FFT,\n        hop_length=HOP_SIZE,\n        n_mels=N_MELS,\n        fmin=FMIN,\n        fmax=FMAX,\n        power=2.0\n    )\n    mel_spec_db = librosa.power_to_db(mel_spec, ref=np.max)\n    mel_spec_norm = (mel_spec_db - mel_spec_db.min()) / (mel_spec_db.max() - mel_spec_db.min() + 1e-8)\n    return mel_spec_norm","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Generate mel spectrograms for training audio\nprint(\"Generating mel spectrograms for training audio\")\nall_training_mel_specs = []\ntraining_frame_count = 0\nfor frame in training_frames:\n    mel_spec = audio2melspec(frame)\n    all_training_mel_specs.append(mel_spec)\n    print(f\"added mel spectrogram for training sample {training_frame_count}\", end=\"\\r\")\n    training_frame_count += 1","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Generate mel spectrograms for validation audio\nprint(\"Generating mel spectrograms for validation audio\")\nall_val_mel_specs = []\nval_frame_count = 0\nfor frame in validation_frames:\n    mel_spec = audio2melspec(frame)\n    all_val_mel_specs.append(mel_spec)\n    print(f\"added mel spectrogram for validation sample {val_frame_count}\", end=\"\\r\")\n    val_frame_count += 1","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"### DEBUG/OUTPUT CELL ###\n# Check list lengths\nprint(len(all_training_mel_specs))\nprint(len(all_val_mel_specs))\n\n# Check dimensions of mel spectrogram data\nprint(len(all_training_mel_specs[0]))\nprint(len(all_training_mel_specs[0][0]))\n\n# Print one training mel spectrogram sample\nprint(all_training_mel_specs[0])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"### DEBUG/OUTPUT CELL ###\nplt.figure(figsize=(5, 5), frameon=True)\nlibrosa.display.specshow(all_training_mel_specs[0], sr=SR, hop_length=HOP_SIZE, x_axis=\"time\", y_axis=\"mel\", fmin=FMIN, fmax=FMAX)\nplt.colorbar(format = \"%+2.0f dB\")\nplt.title(\"Mel Spectrogram\")\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Clean up\ndel training_frames\ndel validation_frames\ngc.collect()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def to_melspectrogram_image(mel_spectrogram_norm: np.ndarray, target_size: tuple = (IMG_SIZE, IMG_SIZE)) -> np.ndarray:\n    \"\"\"\n    Converts an audio file to a mel spectrogram image.\n\n    @param {np.ndarray} mel_spectrogram_norm: Normalized mel spectrogram data.\n    @param {tuple(int, int)} target_size: The target size for the resized image.\n    @return {np.ndarray} mel_spectrogram_resized: A 224x224 NumPy array representing the mel spectrogram image (or None if there's an error).\n    @raise {Exception} e: Outputs exception and returns `None`.\n    \"\"\"\n    try:\n        img = Image.fromarray(mel_spectrogram_norm * 255).convert(\"RGB\")  # Scale to 0-255 and convert to RGB\n        \n        img = img.resize(target_size, Image.Resampling.LANCZOS) # Use LANCZOS for better resizing\n        # img = img.resize(target_size, Image.Resampling.BICUBIC) # Use default resampling filter\n        \n        mel_spectrogram_resized = np.array(img) / 1.0 # No normalization\n        # mel_spectrogram_resized = np.array(img) / 255.0  # Normalize back to [0, 1]\n        \n        return mel_spectrogram_resized\n    except Exception as e:\n        print(str(e))\n        return None","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Resize training sample mel spectrograms\nprint(\"Resizing training sample mel spectograms\")\nall_training_mel_specs_resized = []\ntraining_resized_count = 0\nfor mel_spec in all_training_mel_specs:\n    resized = to_melspectrogram_image(mel_spec)\n    all_training_mel_specs_resized.append(resized)\n    print(f\"added resized mel spectrogram for training sample {training_resized_count}\", end=\"\\r\")\n    training_resized_count += 1","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Resize validation sample mel spectrograms\nprint(\"Resizing validation sample mel spectograms\")\nall_val_mel_specs_resized = []\nval_resized_count = 0\nfor mel_spec in all_val_mel_specs:\n    resized = to_melspectrogram_image(mel_spec)\n    all_val_mel_specs_resized.append(resized)\n    print(f\"added resized mel spectrogram for validation sample {val_resized_count}\", end=\"\\r\")\n    val_resized_count += 1","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"### DEBUG/OUTPUT CELL ###\n# Check list lengths\nprint(len(all_training_mel_specs_resized))\nprint(len(all_val_mel_specs_resized))\n\n# Check the type of a resized training mel spectrogram sample\nprint(type(all_training_mel_specs_resized[0]))\n\n# Check the shape of a resized training mel spectrogram sample\nprint(all_training_mel_specs_resized[0].shape)\n\n# Print one resized training mel spectrogram sample\nprint(all_training_mel_specs_resized[0])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_pseudo_labels_data(DATA_PATH: str) -> Union[list, list]:\n    pseudo_labels = []\n    pseudo_mel_specs_resized = []\n    \n    _dirs = os.listdir(DATA_PATH)\n    for _dir in _dirs:\n        _pseudo_label = _dir\n        _files = os.listdir(DATA_PATH + \"/\" + _dir)\n        _count = 1\n        for _file in _files:\n            if MAX_PSEUDO_LABELS_FILES > _count: \n                with Image.open(DATA_PATH + \"/\" + _dir + \"/\" + _file) as _img:\n                    _img_data = np.array(_img)/1.0 # No normalization\n                pseudo_mel_specs_resized.append(_img_data)\n                pseudo_labels.append(_pseudo_label)\n                _count += 1\n            else:\n                break\n    return pseudo_labels, pseudo_mel_specs_resized            ","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load pseudo-labels data\nif USE_PSEUDO_LABELS:\n    pseudo_labels, pseudo_mel_specs_resized = load_pseudo_labels_data(PSEUDO_LABELS_DATA_PATH)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"### DEBUG/OUTPUT CELL ###\n# Check list lengths\n# if USE_PSEUDO_LABELS:\n#     print(len(pseudo_labels))\n#     print(len(pseudo_mel_specs_resized))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Add all pseudo-labels data to training data; no data is added to validation data\nif USE_PSEUDO_LABELS:\n    training_frames_labels = training_frames_labels + pseudo_labels\n    all_training_mel_specs_resized = all_training_mel_specs_resized + pseudo_mel_specs_resized","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"### DEBUG/OUTPUT CELL ###\n# Check list lengths\nprint(len(all_training_mel_specs_resized))\nprint(len(training_frames_labels))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"### CLEAN UP ###\ndel all_training_mel_specs\ndel all_val_mel_specs\n\nif USE_PSEUDO_LABELS:\n    del pseudo_labels\n    del pseudo_mel_specs_resized\ngc.collect()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def process_labels(label: str, labels: list) -> np.array:\n    \"\"\"\n    One-hot encodes labels.\n\n    @param {str} label: The label that will be encoded.\n    @param {list} labels: List of unique labels.\n    @return {np.array} one_hot: The one-hot encoding.\n    \"\"\"\n    one_hot = np.zeros((len(labels),))\n    idx = np.array(labels.index(label))\n    one_hot[idx] = 1\n    return one_hot","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Prepare the input data\ninput_train_images = all_training_mel_specs_resized\ninput_val_images = all_val_mel_specs_resized\n\ninput_train_labels = []\ninput_validation_labels = []\n\nLABELS = {label for label in training_frames_labels}\nLABELS = list(LABELS)\n\n# Check LABELS length\nprint(f\"unique LABELS length is {len(LABELS)}\")\n\nfor label in training_frames_labels:\n    one_hot_label = process_labels(label, LABELS)\n    input_train_labels.append(one_hot_label)\n\nfor label in validation_frames_labels:\n    one_hot_label = process_labels(label, LABELS)\n    input_validation_labels.append(one_hot_label)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"### DEBUG/OUTPUT CELL ###\n# Check list lengths\nprint(len(input_train_images))\nprint(len(input_train_labels))\nprint(len(input_val_images))\nprint(len(input_validation_labels))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Create training and validation datasets\ntrain_ds_mu = None\ntrain_ds_one = None\ntrain_ds_two = None\nval_ds = tf.data.Dataset.from_tensor_slices((input_val_images, input_validation_labels)).batch(batch_size=BATCH_SIZE, drop_remainder=True)\nif USE_MIXUP:\n    train_ds_one = tf.data.Dataset.from_tensor_slices((input_train_images, input_train_labels)).shuffle(BATCH_SIZE * 100).batch(batch_size=BATCH_SIZE, drop_remainder=True)\n    train_ds_two = tf.data.Dataset.from_tensor_slices((input_train_images, input_train_labels)).shuffle(BATCH_SIZE * 100).batch(batch_size=BATCH_SIZE, drop_remainder=True)\nelse:\n    train_ds_mu = tf.data.Dataset.from_tensor_slices((input_train_images, input_train_labels)).shuffle(BATCH_SIZE * 100).batch(batch_size=BATCH_SIZE, drop_remainder=True)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"### DEBUG/OUTPUT CELL ###\n# Check dataset lengths and structures\nif USE_MIXUP:\n    print(f\"training dataset one has length: {len(train_ds_one)}\")\n    print(f\"training dataset two has length: {len(train_ds_two)}\")\n    print(f\"validation dataset has length: {len(val_ds)}\")\n    print(f\"training dataset one:\\n{train_ds_one}\")\n    print(f\"training dataset two:\\n{train_ds_two}\")\n    print(f\"validation dataset:\\n{val_ds}\")\nelse:\n    print(f\"training dataset has length: {len(train_ds_mu)}\")\n    print(f\"validation dataset has length: {len(val_ds)}\")\n    print(f\"training dataset:\\n{train_ds_mu}\")\n    print(f\"validation dataset:\\n{val_ds}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Clean up\ndel all_training_mel_specs_resized\ndel all_val_mel_specs_resized\ngc.collect()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def sample_beta_distribution(size, concentration_0=0.2, concentration_1=0.2):\n    gamma_1_sample = tf_random_gamma(shape=[size], alpha=concentration_1)\n    gamma_2_sample = tf_random_gamma(shape=[size], alpha=concentration_0)\n    return gamma_1_sample / (gamma_1_sample + gamma_2_sample)\n\n\ndef mix_up(ds_one, ds_two, alpha=0.2):\n    images_one, labels_one = ds_one\n    images_two, labels_two = ds_two\n    batch_size = keras.ops.shape(images_one)[0]\n\n    l = sample_beta_distribution(batch_size, alpha, alpha)\n    x_l = tf.cast(keras.ops.reshape(l, (batch_size, 1, 1, 1)), tf.float64)\n    y_l = tf.cast(keras.ops.reshape(l, (batch_size, 1)), tf.float64)\n\n    images = images_one * x_l + images_two * (1 - x_l)\n    labels = labels_one * y_l + labels_two * (1 - y_l)\n    return (images, labels)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# If `USE_MIXUP` is `True`, mixup the images and their corresponding labels.\nif USE_MIXUP:\n    train_ds = tf_data.Dataset.zip((train_ds_one, train_ds_two))\n    \n    # Create the new dataset using the `mix_up` utility\n    train_ds_mu = train_ds.map(\n        lambda ds_one, ds_two: mix_up(ds_one, ds_two, alpha=0.2),\n        num_parallel_calls=tf_data.AUTOTUNE,\n    )\n\n    train_ds_mu = train_ds_mu.concatenate(train_ds_one).concatenate(train_ds_two)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"### DEBUG/OUTPUT CELL ###\nprint(train_ds_mu)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"### CLEAN UP ###\ndel input_train_images\ndel input_train_labels\ndel input_val_images\ndel input_validation_labels\n\nif USE_MIXUP:\n    del train_ds_one\n    del train_ds_two\n    del train_ds\n\ngc.collect()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Section 4: Model Training\n\nThe fourth notebook section is used to:\n\n1. Initialize and configure a Weights and Biases project to capture training run data.\n2. Build and compile the EfficientNet B0 model.\n3. Train the model.\n4. Save the trained model to disk. ","metadata":{}},{"cell_type":"code","source":"# Imports\nfrom tensorflow.keras.applications import EfficientNetB0\nfrom tensorflow.keras.callbacks import CSVLogger, LearningRateScheduler","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Config\nwandb.init(project=\"birdclef-2025-finetuning_enb0-bird_classification\")\nconfig = wandb.config\nconfig.batch_size = BATCH_SIZE\nconfig.epochs = 30\nconfig.image_size = IMG_SIZE\nconfig.num_classes = len(LABELS)\n\n# Regularization constants\nL2_RATE = 1\n\n# CategoricalFocalCrossentropy loss constants\nLOSS_ALPHA = 0.25\nLOSS_GAMMA = 2.0\n\nUNFREEZE_LAYERS = 20\nDROPOUT = 0.5","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def build_model(num_classes: int, img_size: int, unfreeze_layers: int, dropout: float, loss_alpha: float, loss_gamma: float):\n    inputs = layers.Input(shape=(img_size, img_size, 3))\n    model = EfficientNetB0(include_top=False, input_tensor=inputs, weights=\"imagenet\")\n    \n    # Freeze the pretrained weights\n    model.trainable = False\n\n    # Unfreeze last `unfreeze_layers` layers and add regularization\n    for layer in model.layers[-unfreeze_layers:]:\n        if not isinstance(layer, layers.BatchNormalization):\n            layer.trainable = True\n            # layer.kernel_regularizer = tf.keras.regularizers.l2(L2_RATE)\n    \n    # Rebuild top\n    x = layers.GlobalAveragePooling2D(name=\"avg_pool\")(model.output)\n    x = layers.BatchNormalization()(x)\n\n    top_dropout_rate = dropout\n    x = layers.Dropout(top_dropout_rate, name=\"top_dropout\")(x)\n    outputs = layers.Dense(num_classes, activation=\"softmax\", name=\"pred\")(x)\n\n    # Compile\n    model = keras.Model(inputs, outputs, name=\"EfficientNetB0\")\n    optimizer = keras.optimizers.Adam()\n    loss = keras.losses.CategoricalFocalCrossentropy(alpha=loss_alpha, gamma=loss_gamma, from_logits=False, label_smoothing=0.0, axis=-1, reduction=\"sum_over_batch_size\")\n    model.compile(\n        optimizer=optimizer, loss=loss, metrics=[\"accuracy\"]\n    )\n    return model","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = build_model(len(LABELS), config.image_size, UNFREEZE_LAYERS, DROPOUT, LOSS_ALPHA, LOSS_GAMMA)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"### DEBUG/OUTPUT CELL ###\nmodel.summary()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Callback for learning rate scheduler\ndef lr_scheduler(epoch: int) -> float:    \n    if epoch < 30:\n        return 1.5e-4\n    elif epoch < 40:\n        return (1.5e-4) * 0.5\n    else:\n        return ((1.5e-4) * 0.5) * 0.5\n\nlr_callback = LearningRateScheduler(lr_scheduler)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Callback for logging training data\ncsv_logger = CSVLogger(\"training_log.csv\", append=True)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class WandBLogger(tf.keras.callbacks.Callback):\n    def on_epoch_end(self, epoch, logs=None):\n        wandb.log({\"accuracy\": logs[\"accuracy\"], \"loss\": logs[\"loss\"], \"val_accuracy\": logs[\"val_accuracy\"], \"val_loss\": logs[\"val_loss\"]})\nwandb_logger = WandBLogger()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Training the model\nhist = model.fit(train_ds_mu, epochs=config.epochs, validation_data=val_ds, callbacks=[csv_logger, lr_callback, wandb_logger])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.save(\"bird_vocalization_classifier.keras\")","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}