{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.11.11"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":91844,"databundleVersionId":11361821,"sourceType":"competition"},{"sourceId":12052350,"sourceType":"datasetVersion","datasetId":7585054},{"sourceId":189366517,"sourceType":"kernelVersion"}],"dockerImageVersionId":31041,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"This is the train notebook for the inference notebook : https://www.kaggle.com/code/hosseinbahraminekoo/birdclef-2025-inference","metadata":{}},{"cell_type":"markdown","source":"**Score: 0.747**","metadata":{}},{"cell_type":"markdown","source":"\nImport all necessary libraries for data handling, audio processing, modeling, evaluation, and visualization.\n* Includes PyTorch for deep learning, librosa and torchaudio for audio feature extraction, sklearn for metrics and cross-validation,timm for pretrained models, and Kaggle utilities for leaderboard evaluation.\n* Sets the computation device (GPU if available).\n","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport librosa\nimport glob\nimport torch\nimport torchaudio.transforms as T\nimport torch.nn as nn\nimport os\nimport random\nfrom matplotlib import pyplot as plt\nimport seaborn as sns\nfrom ast import literal_eval\nimport timm\nimport pandas.api.types\n\nimport kaggle_metric_utilities\n\nimport sklearn.metrics\nfrom sklearn.model_selection import GroupKFold\nfrom sklearn.utils.class_weight import compute_class_weight\nfrom tqdm import tqdm\nimport gc\nfrom warnings import filterwarnings\nfilterwarnings(\"ignore\")\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","trusted":true,"execution":{"iopub.status.busy":"2025-06-03T10:55:18.710760Z","iopub.execute_input":"2025-06-03T10:55:18.711357Z","iopub.status.idle":"2025-06-03T10:55:31.196574Z","shell.execute_reply.started":"2025-06-03T10:55:18.711331Z","shell.execute_reply":"2025-06-03T10:55:31.195648Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":" Configuration class holding all global settings and constants used throughout the pipeline.\n * Includes dataset paths, audio sampling and spectrogram parameters, and training metadata.\n","metadata":{}},{"cell_type":"code","source":"class Config:\n    train_dir = \"/kaggle/input/birdclef-2025/train_audio\"\n    seed = 42\n    train_csv = \"/kaggle/input/birdclef-2025/train.csv\"\n    sample_csv = \"/kaggle/input/birdclef-2025/sample_submission.csv\"\n    test_soundscapes = \"/kaggle/input/birdclef-2025/test_soundscapes\"\n    sr = int(32e3)\n\n    num_classes = 206\n    n_fft = 2048\n    hop_length = 500\n\n    n_mels = 256\n    fmin = 50\n    fmax = 16000\n    power = 2\n    \n    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-03T10:55:31.198193Z","iopub.execute_input":"2025-06-03T10:55:31.199064Z","iopub.status.idle":"2025-06-03T10:55:31.203770Z","shell.execute_reply.started":"2025-06-03T10:55:31.199042Z","shell.execute_reply":"2025-06-03T10:55:31.202994Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"* Ensures reproducibility by setting seeds for Python, NumPy, and PyTorch (CPU and GPU).\n* Also configures PyTorch for deterministic behavior to reduce result variance.\n","metadata":{}},{"cell_type":"code","source":"def set_seed(seed: int = Config.seed) -> None:\n    random.seed(seed)\n    np.random.seed(seed)\n    # reproducible weight initialization\n    torch.manual_seed(seed)\n\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed(seed)\n        torch.cuda.manual_seed_all(seed)\n        \n    torch.backends.cudnn.determinstic = True\n    torch.backends.cudnn.benchmark = False\n\n    print(f\"[info] set seed: {seed}\")\n\nset_seed()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-03T10:55:31.204783Z","iopub.execute_input":"2025-06-03T10:55:31.205025Z","iopub.status.idle":"2025-06-03T10:55:31.229514Z","shell.execute_reply.started":"2025-06-03T10:55:31.205008Z","shell.execute_reply":"2025-06-03T10:55:31.228621Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"* Loads training metadata CSV and processes label columns.\n* Converts list-like strings to '###'-joined strings and prepends full file paths to filenames.","metadata":{}},{"cell_type":"code","source":"data_df = pd.read_csv(Config.train_csv)\nfor col in ('secondary_labels', 'type'):\n    data_df[col] = data_df[col].apply(lambda x: '###'.join(literal_eval(x)))\n\ndata_df['filename'] = data_df['filename'].apply(lambda x: Config.train_dir + '/' + x)\ndata_df.sample(10)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-03T10:55:31.231482Z","iopub.execute_input":"2025-06-03T10:55:31.231814Z","iopub.status.idle":"2025-06-03T10:55:31.857531Z","shell.execute_reply.started":"2025-06-03T10:55:31.231796Z","shell.execute_reply":"2025-06-03T10:55:31.856726Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Null Check**","metadata":{}},{"cell_type":"code","source":"data_df.isnull().sum()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-03T10:55:31.858572Z","iopub.execute_input":"2025-06-03T10:55:31.859407Z","iopub.status.idle":"2025-06-03T10:55:31.883125Z","shell.execute_reply.started":"2025-06-03T10:55:31.859377Z","shell.execute_reply":"2025-06-03T10:55:31.882439Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Distribution of the ratings**","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize = (15, 3))\nsns.histplot(data_df, x = 'rating')\nplt.xticks(np.arange(0, 5.5, 0.5))\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-03T10:55:31.884072Z","iopub.execute_input":"2025-06-03T10:55:31.884390Z","iopub.status.idle":"2025-06-03T10:55:32.144999Z","shell.execute_reply.started":"2025-06-03T10:55:31.884364Z","shell.execute_reply":"2025-06-03T10:55:32.144218Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Distribution of the primary label**","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize = (15, 3))\nsns.countplot(data_df, x = 'primary_label')\nplt.xticks(rotation = 90)\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-03T10:55:32.145856Z","iopub.execute_input":"2025-06-03T10:55:32.146206Z","iopub.status.idle":"2025-06-03T10:55:33.416845Z","shell.execute_reply.started":"2025-06-03T10:55:32.146157Z","shell.execute_reply":"2025-06-03T10:55:33.415850Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Distribution of the primary label with different ratings**","metadata":{}},{"cell_type":"code","source":"for r in range(0, 6):\n    plt.figure(figsize = (20, 3))\n    sns.countplot(data_df[data_df['rating'] == float(r)], x = 'primary_label')\n    plt.title(f\"Rating {r}\")\n    plt.xticks(rotation = 90)\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-03T10:55:33.417618Z","iopub.execute_input":"2025-06-03T10:55:33.417831Z","iopub.status.idle":"2025-06-03T10:55:39.460877Z","shell.execute_reply.started":"2025-06-03T10:55:33.417813Z","shell.execute_reply":"2025-06-03T10:55:39.459838Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Stats of audio duration**","metadata":{}},{"cell_type":"code","source":"durations = []\nfor idx, row in data_df.sample(100).iterrows():\n    data, _ = librosa.load(row['filename'], sr = Config.sr)\n    durations.append(librosa.get_duration(y = data, sr = Config.sr))\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-03T10:55:39.462061Z","iopub.execute_input":"2025-06-03T10:55:39.462444Z","iopub.status.idle":"2025-06-03T10:55:59.813844Z","shell.execute_reply.started":"2025-06-03T10:55:39.462424Z","shell.execute_reply":"2025-06-03T10:55:59.813194Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"d_df = pd.DataFrame(columns = ['durations'], data = durations)\n\nplt.figure(figsize = (10, 5))\nplt.title(\"Distribution of audio legnths\")\nsns.histplot(d_df, x = 'durations', bins = 100)\nplt.show()\n\nd_df.describe()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-03T10:55:59.816270Z","iopub.execute_input":"2025-06-03T10:55:59.816753Z","iopub.status.idle":"2025-06-03T10:56:00.120551Z","shell.execute_reply.started":"2025-06-03T10:55:59.816735Z","shell.execute_reply":"2025-06-03T10:56:00.119672Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Check out spectrograms**","metadata":{}},{"cell_type":"markdown","source":"* Visualizes audio signal in multiple representations.\n* Displays raw waveform, linear and log-scaled spectrograms, and mel spectrogram for a given file.","metadata":{}},{"cell_type":"code","source":"def show_signal(file_path):\n    class_, collector = file_path.split(\"/\")[-2:]\n    \n    y, sr = librosa.load(file_path, sr = Config.sr)\n    \n    fig, axes = plt.subplots(2, 2, figsize = (20, 10))\n    fig.suptitle(f\"Class: {class_} | Collector: {collector}\", fontsize = 16)\n\n    # Plotting raw signal\n    librosa.display.waveshow(y, sr = sr, ax = axes[0, 0])\n    axes[0, 0].set_title('Raw Signal')\n\n    # Plotting Fourier Transformed Signal\n    ft = np.abs(librosa.stft(\n        y,\n        n_fft = Config.n_fft,\n        hop_length = Config.hop_length\n    ))\n    im1 = librosa.display.specshow(\n        ft,\n        sr = sr,\n        x_axis = 'time',\n        y_axis = 'linear',\n        ax = axes[0, 1]\n    )\n    fig.colorbar(im1, ax = axes[0, 1])\n    axes[0, 1].set_title(\"Spectrogram\")\n\n    # Plotting log scaled fourier transformed signal\n    ft_db = librosa.amplitude_to_db(ft, ref = np.max)\n    im2 = librosa.display.specshow(\n        ft_db,\n        sr = sr,\n        x_axis = 'time',\n        y_axis = 'log',\n        ax = axes[1, 0]\n    )\n    fig.colorbar(im2, ax = axes[1, 0])\n    axes[1, 0].set_title(\"Log scaled spectrogram\")\n\n    # Plotting mel spectrogram\n    mel_sp = librosa.feature.melspectrogram(\n        y = y,\n        sr = Config.sr,\n        fmin = Config.fmin,\n        fmax = Config.fmax,\n        power = Config.power,\n        n_mels = Config.n_mels  \n    )\n    mel_sp = librosa.power_to_db(mel_sp, ref = np.max)\n    im3 = librosa.display.specshow(\n        mel_sp,\n        y_axis = 'mel',\n        sr = Config.sr,\n        fmin = Config.fmin,\n        x_axis = 'time',\n        fmax = Config.fmax,\n        ax = axes[1, 1]\n    )\n    fig.colorbar(im3, ax = axes[1, 1])\n    axes[1, 1].set_title(\"Mel spectrogram\")\n\n    plt.show()    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-03T10:56:00.121470Z","iopub.execute_input":"2025-06-03T10:56:00.122379Z","iopub.status.idle":"2025-06-03T10:56:00.131439Z","shell.execute_reply.started":"2025-06-03T10:56:00.122348Z","shell.execute_reply":"2025-06-03T10:56:00.130580Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for idx, row in data_df.sample(10).iterrows():\n    show_signal(row['filename'])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-03T10:56:00.132228Z","iopub.execute_input":"2025-06-03T10:56:00.132537Z","iopub.status.idle":"2025-06-03T10:56:55.120745Z","shell.execute_reply.started":"2025-06-03T10:56:00.132510Z","shell.execute_reply":"2025-06-03T10:56:55.119798Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**BirdClefDataset class and label preparation for BirdCLEF 2025.**\n\n* This section creates label-index mappings and prepares the main PyTorch Dataset used for training and validation.\n* Each audio file is loaded, optionally augmented with Gaussian noise, and converted into a normalized mel spectrogram.\n* SpecAugment is applied to the spectrogram for additional robustness during training.\n* The dataset returns either a multi-hot encoded label vector for classification or the spectrogram alone for inference.\n* This flexible and efficient data pipeline ensures proper preprocessing and augmentation for high-performance model training.\n","metadata":{}},{"cell_type":"code","source":"label_mapper = {\n    label: idx\n    for idx, label in enumerate(sorted(data_df['primary_label'].unique()))\n}\n# Map labels to indices for class_weight\ndata_df['label_idx'] = data_df['primary_label'].map(label_mapper)\n\nrev_mapper = {\n    idx: label\n    for label, idx in label_mapper.items()\n}\n\nConfig.num_classes = len(label_mapper)  # Update the config\n\nclass BirdClefDataset(torch.utils.data.Dataset):\n    def __init__(self, df, mode = 'train', aug = True):\n        self.df = df\n        self.mode = mode\n        # Initialize augmentations\n        self.aug = aug\n        self.spec_aug = torch.nn.Sequential(\n            T.TimeMasking(time_mask_param=16),  \n            T.FrequencyMasking(freq_mask_param=4)  \n        ) if aug else None\n\n        # Gaussian noise parameters\n        self.noise_prob = 0.5  # 50% chance to apply noise\n        self.noise_scale = 0.01  # Adjust based on your audio volume\n\n    def __len__(self): return len(self.df)\n\n    def _add_gaussian_noise(self, audio):\n        \"\"\"Add Gaussian noise to raw audio waveform\"\"\"\n        if random.random() < self.noise_prob:\n            noise = np.random.normal(0, self.noise_scale, len(audio))\n            return audio + noise\n        return audio\n\n    def process(self, audio_path):\n        data, _ = librosa.load(audio_path, sr = Config.sr)\n\n        data = data * 1024\n\n        # Apply Gaussian noise to RAW AUDIO (before spectrogram)\n        if self.aug:\n            data = self._add_gaussian_noise(data)\n        \n        chunk_duration = 10\n        min_len = chunk_duration * Config.sr\n\n        # If the audio signal is less than min_len\n        if len(data) < min_len: \n            cnt = int(np.ceil(min_len / len(data)))\n            data = np.tile(data, cnt)\n\n        # Making the data length divisible by min_len\n        leftover = len(data) % min_len\n        if leftover > 0:\n            front_crop = leftover // 2\n            back_crop = leftover - front_crop\n            data = data[front_crop : len(data) - back_crop]\n\n        # Truncating the signal to min_len\n        data = data[:min_len]\n        data = data.reshape(-1, min_len)\n\n        # Creating Mel  spectrogram\n        mel_sp = librosa.feature.melspectrogram(\n            y = data,\n            sr = Config.sr,\n            fmin = Config.fmin,\n            fmax = Config.fmax,\n            power = Config.power,\n            n_mels = Config.n_mels,\n            n_fft = Config.n_fft,\n            hop_length = Config.hop_length\n        )\n\n        mel_sp = librosa.power_to_db(mel_sp, ref = 1)\n\n        # Normalizing the features\n        eps = 1e-12\n        mel_sp = (mel_sp - mel_sp.min())/(mel_sp.max() - mel_sp.min() + eps)\n\n        mel_sp = mel_sp[:, :, :640]\n        return mel_sp\n            \n    \n    def __getitem__(self, idx):\n        row = self.df.loc[idx, :]\n        filename = row['filename']\n\n        # TODO: Spectrogram conversion\n        x = self.process(filename) # Returns numpy array\n\n        # Convert to tensor FIRST before augmentation\n        x = torch.from_numpy(x).float()\n\n        if self.mode == 'train':\n            if self.aug:\n                x = self.spec_aug(x)  # Apply augmentations    \n            try:\n                # Get label index\n                label_idx = label_mapper[row['primary_label']]\n                # Convert to multi-hot vector\n                label_vec = torch.zeros(Config.num_classes, dtype=torch.float)\n                label_vec[label_idx] = 1.0\n                # y = label_mapper[row['primary_label']]\n            except KeyError as e:\n                print(f\"Error: Label '{row['primary_label']}' not in label_mapper!\")\n                raise\n            return x, label_vec\n\n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-03T10:56:55.121758Z","iopub.execute_input":"2025-06-03T10:56:55.122043Z","iopub.status.idle":"2025-06-03T10:56:55.148666Z","shell.execute_reply.started":"2025-06-03T10:56:55.122021Z","shell.execute_reply":"2025-06-03T10:56:55.148005Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"* Define a wrapper for a pretrained image classification model (from timm library)\n* tailored for single-channel mel spectrogram input.\n* The model is initialized with \npretrained weights, dropout regularization, and a final layer adapted to the number \nof bird species classes in BirdCLEF 2025. This class enables flexible use of \narchitectures like EfficientNet, ConvNeXt, etc., as the backbone.","metadata":{}},{"cell_type":"code","source":"class Model(nn.Module):\n    def __init__(self, model_name: str):\n        super().__init__()\n        self.base_model = timm.create_model(\n            model_name = model_name,\n            pretrained = True,        # use pretrained weights\n            in_chans = 1,             # for single-channel input\n            num_classes = Config.num_classes,\n            drop_rate = 0.3           # Dropout before classifier\n        )\n\n\n    def forward(self, x):\n        return self.base_model(x)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-03T10:56:55.170924Z","iopub.execute_input":"2025-06-03T10:56:55.171518Z","iopub.status.idle":"2025-06-03T10:56:55.187525Z","shell.execute_reply.started":"2025-06-03T10:56:55.171492Z","shell.execute_reply":"2025-06-03T10:56:55.186667Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"* Custom scoring utilities for the BirdCLEF 2025 competition.\n* Implements a macro-averaged ROC-AUC metric that only considers classes \nwith at least one positive label (to avoid penalizing rare or absent classes).\n* The `score` function ensures the input is valid and computes the score using \nsklearn.metrics.roc_auc_score, while the `cal_score` function formats predictions \nand labels as DataFrames and calls the custom scoring logic. These functions \nprovide a reliable offline evaluation consistent with Kaggle's official metric.","metadata":{}},{"cell_type":"code","source":"class ParticipantVisibleError(Exception):\n    pass\n\n\ndef score(solution: pd.DataFrame, submission: pd.DataFrame, row_id_column_name: str) -> float:\n    '''\n    Version of macro-averaged ROC-AUC score that ignores all classes that have no true positive labels.\n    '''\n    del solution[row_id_column_name]\n    del submission[row_id_column_name]\n\n    if not pandas.api.types.is_numeric_dtype(submission.values):\n        bad_dtypes = {x: submission[x].dtype  for x in submission.columns if not pandas.api.types.is_numeric_dtype(submission[x])}\n        raise ParticipantVisibleError(f'Invalid submission data types found: {bad_dtypes}')\n\n    solution_sums = solution.sum(axis=0)\n    scored_columns = list(solution_sums[solution_sums > 0].index.values)\n    assert len(scored_columns) > 0\n\n    return kaggle_metric_utilities.safe_call_score(sklearn.metrics.roc_auc_score, solution[scored_columns].values, submission[scored_columns].values, average='macro')\n\n\ndef cal_score(labels, preds):\n    labels = np.concatenate(labels)\n    preds = np.concatenate(preds)\n\n    labels_df = pd.DataFrame(labels > 0.5, columns = list(label_mapper.keys()))\n    pred_df = pd.DataFrame(preds, columns = list(label_mapper.keys()))\n\n    labels_df['id'] = np.arange(len(labels_df))\n    pred_df['id'] = np.arange(len(pred_df))\n\n    return score(labels_df, pred_df, row_id_column_name = 'id')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-03T10:56:55.226126Z","iopub.execute_input":"2025-06-03T10:56:55.226370Z","iopub.status.idle":"2025-06-03T10:56:55.240937Z","shell.execute_reply.started":"2025-06-03T10:56:55.226353Z","shell.execute_reply":"2025-06-03T10:56:55.240037Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"##### Create directory to save model checkpoints, if it doesn't already exist","metadata":{}},{"cell_type":"code","source":"# Create save directory\nsave_dir = \"model_checkpoints\"\nos.makedirs(save_dir, exist_ok=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-03T10:56:55.241675Z","iopub.execute_input":"2025-06-03T10:56:55.241937Z","iopub.status.idle":"2025-06-03T10:56:55.260711Z","shell.execute_reply.started":"2025-06-03T10:56:55.241915Z","shell.execute_reply":"2025-06-03T10:56:55.259834Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"* Define training configuration parameters including epochs, learning rate, and number of folds for cross-validation.\n* Prepare the dataset and initialize GroupKFold to split data into folds while ensuring no data leakage between groups defined by 'filename'.\n* Assign fold numbers to the dataframe for later use in training and validation splits.","metadata":{}},{"cell_type":"code","source":"# Training configs\n# epochs = 20\nepochs = 15\nnum_folds = 3\nsave_interval = 3\nlr = 1e-3\n\ntarget_col = 'primary_label'\n#df = data_df.sample(1000).reset_index()\ndf = data_df\n\n# Setup GroupKFold\ngkf = GroupKFold(\n    n_splits = num_folds,\n)\n\n# 1. Assign fold numbers once\ndf['kfold'] = -1\nfor fold, (train_idx, val_idx) in enumerate(gkf.split(df, groups = df['filename'])):\n    df.loc[val_idx, 'kfold'] = fold","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-03T10:56:55.261663Z","iopub.execute_input":"2025-06-03T10:56:55.261934Z","iopub.status.idle":"2025-06-03T10:56:55.372410Z","shell.execute_reply.started":"2025-06-03T10:56:55.261907Z","shell.execute_reply":"2025-06-03T10:56:55.371335Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-03T10:56:55.373250Z","iopub.execute_input":"2025-06-03T10:56:55.373486Z","iopub.status.idle":"2025-06-03T10:56:55.390301Z","shell.execute_reply.started":"2025-06-03T10:56:55.373468Z","shell.execute_reply":"2025-06-03T10:56:55.389576Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"##### Perform Mixup data augmentation by blending pairs of inputs and their labels with a weighted factor sampled from a Beta distribution.","metadata":{}},{"cell_type":"code","source":"def mixup_data(x, y, alpha=0.4):\n    '''Apply Mixup to inputs and labels'''\n    if alpha > 0:\n        lam = np.random.beta(alpha, alpha)\n    else:\n        lam = 1\n\n    batch_size = x.size()[0]\n    index = torch.randperm(batch_size).to(x.device)\n\n    mixed_x = lam * x + (1 - lam) * x[index, :]\n    y_a, y_b = y, y[index]\n    return mixed_x, y_a, y_b, lam","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-03T10:56:55.391196Z","iopub.execute_input":"2025-06-03T10:56:55.391536Z","iopub.status.idle":"2025-06-03T10:56:55.404752Z","shell.execute_reply.started":"2025-06-03T10:56:55.391501Z","shell.execute_reply":"2025-06-03T10:56:55.403811Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### Define the main training function that:\n- Initializes the model, optimizer, and learning rate scheduler (with optional checkpoint resume)\n- Performs stratified GroupKFold cross-validation on the dataset\n- Computes per-fold class weights to handle class imbalance dynamically\n- Prepares training and validation datasets with augmentations and data loaders\n- Implements the training loop with Mixup data augmentation, gradient clipping, and loss calculation using weighted BCEWithLogitsLoss\n- Evaluates model performance on validation data after each epoch and tracks metrics such as AUC and loss\n- Applies early stopping based on validation AUC to prevent overfitting\n- Saves the best model per fold and periodically saves training checkpoints for fault tolerance\n- Manages GPU memory efficiently and resets epoch counter for each fold training cycle\n","metadata":{}},{"cell_type":"code","source":"def train_model(resume_checkpoint=None):\n    # Initialize model and optimizer ONCE (outside folds)\n    model = Model(model_name='tf_efficientnet_b0').to(device)\n\n    optimizer = torch.optim.AdamW(\n        model.parameters(), \n        lr = lr, \n        weight_decay=0\n    )\n    scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='max', patience=2)\n    \n    start_fold, start_epoch = 0, 0\n    if resume_checkpoint:\n        checkpoint = torch.load(resume_checkpoint)\n        model.load_state_dict(checkpoint['model_state_dict'])\n        optimizer.load_state_dict(checkpoint['optimizer_state_dict'])\n        start_fold = checkpoint['fold']\n        scheduler.load_state_dict(checkpoint[\"scheduler\"])\n        # start_epoch = checkpoint['epoch'] + 1\n        start_epoch = checkpoint['epoch'] \n\n    for fold in range(start_fold, num_folds):\n        train_df = df[df['kfold'] != fold].reset_index(drop=True)\n        val_df = df[df['kfold'] == fold].reset_index(drop=True)\n\n        # Compute class weights for this fold\n        y_train = train_df['label_idx'].values\n        present_classes = np.unique(y_train)\n        computed_weights = compute_class_weight(\n            class_weight='balanced',\n            classes = present_classes,\n            y = y_train\n        )\n        # Ensure consistent shape (num_classes) with model output\n        class_weights_tensor = torch.zeros(Config.num_classes, dtype=torch.float32).to(device)\n        for cls, w in zip(present_classes, computed_weights):\n            class_weights_tensor[cls] = w\n            \n        criterion = torch.nn.BCEWithLogitsLoss(pos_weight = class_weights_tensor)\n        \n        train_ds = BirdClefDataset(train_df, mode='train', aug = True)\n        val_ds = BirdClefDataset(val_df, mode='train', aug = False)\n        \n        train_loader = torch.utils.data.DataLoader(\n            train_ds, batch_size=16, shuffle=True, num_workers=2, drop_last=True\n        )\n        val_loader = torch.utils.data.DataLoader(\n            val_ds, batch_size=16, shuffle=False, num_workers=2, drop_last=False\n        )\n\n        best_auc = 0\n        patience = 3  # Early stopping patience\n        patience_counter = 0\n\n        for epoch in range(start_epoch, epochs):\n            model.train()\n            pred_train, label_train = [], []\n            running_loss = 0.0\n\n            for x, y in tqdm(train_loader, desc=\"Training\"):\n                x, y = x.to(device), y.to(device)\n            \n                # Apply Mixup\n                x, y_a, y_b, lam = mixup_data(x, y, alpha=0.4)\n            \n                optimizer.zero_grad()\n                outputs = model(x)\n                loss = lam * criterion(outputs, y_a) + (1 - lam) * criterion(outputs, y_b)\n                loss.backward()\n\n                torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)  # Gradient clipping\n\n                optimizer.step()\n                   \n                running_loss += loss.item()\n                probs = torch.softmax(outputs, dim=1)\n                pred_train.append(probs.detach().cpu().numpy())\n                #label_train.append(y_one_hot.detach().cpu().numpy())  # <-- Store one-hot for AUC\n                label_train.append(y.detach().cpu().numpy())\n\n            # Validation\n            model.eval()\n            pred_val, label_val = [], []\n            running_val_loss = 0.0\n\n            with torch.no_grad():\n                for x, y in tqdm(val_loader, desc=\"Validation\"):\n                    x, y = x.to(device), y.to(device)\n                    \n                    outputs = model(x)\n                    loss = criterion(outputs, y)\n                    running_val_loss += loss.item()\n                    \n                    probs = torch.softmax(outputs, dim = 1)\n                    pred_val.append(probs.detach().cpu().numpy())\n                    #label_val.append(y_one_hot.detach().cpu().numpy())\n                    label_val.append(y.detach().cpu().numpy())\n\n            # Compute metrics \n            auc_train = cal_score(label_train, pred_train)\n            auc_val = cal_score(label_val, pred_val)\n            avg_train_loss = running_loss / len(train_loader)\n            avg_val_loss = running_val_loss / len(val_loader)\n\n            print(f\"Fold {fold} | Epoch {epoch} | Train Loss: {avg_train_loss:.4f} | Val Loss: {avg_val_loss:.4f}\")\n            print(f\"Fold {fold} | Epoch {epoch} | Train AUC: {auc_train:.4f} | Val AUC: {auc_val:.4f}\")\n\n            # Update LR scheduler based on validation AUC\n            scheduler.step(auc_val)\n            for i, param_group in enumerate(optimizer.param_groups):\n                current_lr = param_group['lr']\n                print(f\"Epoch {epoch} | Fold {fold} | Learning Rate (group {i}): {current_lr}\")\n\n            # Save best model \n            if auc_val > best_auc:\n                best_auc = auc_val\n                torch.save(\n                    model.state_dict(), \n                    f\"fold_{fold}_epoch_{epoch}_best_effnetB0_val_auc_{auc_val:.4f}_val_loss_{avg_val_loss}.pth\"\n                )\n                patience_counter = 0\n            else:\n                patience_counter += 1\n                if patience_counter >= patience:\n                    print(f\"Early stopping at epoch {epoch} for fold {fold}\")\n                    break\n\n            # Save checkpoint every 3 epochs\n            if (epoch + 1) % save_interval == 0:\n                checkpoint_path = os.path.join(\n                    save_dir, \n                    f\"fold{fold}_epoch{epoch+1}.pth.tar\"\n                )\n                \n                torch.save({\n                    'epoch': epoch + 1,\n                    'fold': fold,\n                    'model_state_dict': model.state_dict(),\n                    'optimizer_state_dict': optimizer.state_dict(),\n                    'scheduler': scheduler.state_dict()\n                }, checkpoint_path)\n                \n                print(f\"Saved checkpoint to {checkpoint_path}\")\n        \n        start_epoch = 0  # Reset for next fold\n        torch.cuda.empty_cache()\n\ntrain_model()  # Start fresh\n#checkpoint_path = os.path.join(\n#    save_dir, \n#    \"fold2_epoch3.pth.tar\"\n#)\n#train_model(\"/kaggle/input/birdclef-fold1-epoch-9-mixup/fold1_epoch9.pth.tar\")  # Resume from a checkpoint","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-03T10:56:55.405720Z","iopub.execute_input":"2025-06-03T10:56:55.406055Z","execution_failed":"2025-06-03T11:20:24.165Z"}},"outputs":[],"execution_count":null}]}