{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":91844,"databundleVersionId":11361821,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":12688641,"sourceType":"datasetVersion","datasetId":8018520},{"sourceId":167220511,"sourceType":"kernelVersion"},{"sourceId":189366517,"sourceType":"kernelVersion"},{"sourceId":509081,"sourceType":"modelInstanceVersion","isSourceIdPinned":false,"modelInstanceId":403798,"modelId":421720}],"dockerImageVersionId":30919,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# **BirdCLEF 2025 Training Notebook**\n\nThis is a baseline training pipeline for BirdCLEF 2025 using Transfer learning BEATs (PyTorch implementation). You can check inference and preprocessing notebooks in the following links: \n\n- [BirdCLEF'25 | Transfer learning BEATs [INFERENCE]](https://www.kaggle.com/code/hubfor/birdclef-25-transfer-learning-beats-inference)\n\n  \n- [BEATs models code](https://www.kaggle.com/datasets/hubfor/microsoft-beats-model)  \n\nNote that by default this notebook is in Debug Mode, so it will only train the model with 2 epochs, but the [weight](https://www.kaggle.com/models/hubfor/finetuned_beats_birdclef/PyTorch/default/1) I used in the inference notebook was obtained after 10 epochs of training.\n\n**Features**\n* Implement with Pytorch\n* Stratified 5-fold cross-validation\n* AdamW optimizer with Cosine Annealing LR scheduling\n* Debug mode for quick experimentation with smaller datasets","metadata":{}},{"cell_type":"markdown","source":"## Libraries","metadata":{}},{"cell_type":"code","source":"import os\nimport logging\nimport random\nimport gc\nimport time\nimport cv2\nimport math\nimport warnings\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\nfrom sklearn.model_selection import StratifiedKFold\nfrom sklearn.metrics import roc_auc_score\nimport librosa\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.optim import lr_scheduler\nfrom torch.utils.data import Dataset, DataLoader\n\nimport torchaudio\n\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom tqdm.auto import tqdm\n\nimport scipy.signal\nfrom scipy.interpolate import interp1d\n\n\nwarnings.filterwarnings(\"ignore\")\nlogging.basicConfig(level=logging.ERROR)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-11T14:02:44.844973Z","iopub.execute_input":"2025-08-11T14:02:44.845489Z","iopub.status.idle":"2025-08-11T14:02:51.148618Z","shell.execute_reply.started":"2025-08-11T14:02:44.845444Z","shell.execute_reply":"2025-08-11T14:02:51.147941Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Configuration","metadata":{}},{"cell_type":"code","source":"class CFG:\n    \n    seed = 42\n    debug = True  \n    apex = False\n    print_freq = 100\n    num_workers = 2\n    \n    OUTPUT_DIR = '/kaggle/working/'\n\n    train_datadir = '/kaggle/input/birdclef-2025/train_audio'\n    train_csv = '/kaggle/input/birdclef-2025/train.csv'\n    test_soundscapes = '/kaggle/input/birdclef-2025/test_soundscapes'\n    submission_csv = '/kaggle/input/birdclef-2025/sample_submission.csv'\n    taxonomy_csv = '/kaggle/input/birdclef-2025/taxonomy.csv'\n\n\n    models_code_path = '/kaggle/input/microsoft-beats-model'\n    models_weights_path = '/kaggle/input/microsoft-pretrained-beats-iter3-plus-as2m/pytorch/default/1/BEATs_iter3_plus_AS2M.pt'\n    pretrained = True\n    in_channels = 1\n\n    LOAD_DATA = True  \n    FS = 32000\n    TARGET_DURATION = 5.0\n    TARGET_SHAPE = (256, 256)\n    \n    N_FFT = 1024\n    HOP_LENGTH = 512\n    N_MELS = 128\n    FMIN = 50\n    FMAX = 14000\n    \n    device = 'cuda' if torch.cuda.is_available() else 'cpu'\n    epochs = 20  \n    batch_size = 16  \n    criterion = 'BCEWithLogitsLoss'\n\n    n_fold = 5\n    selected_folds = [0, 1, 2, 3, 4]   \n\n    optimizer = 'AdamW'\n    lr = 1e-4 \n    weight_decay = 1e-5\n  \n    scheduler = 'CosineAnnealingLR'\n    min_lr = 1e-5\n    T_max = epochs\n\n    aug_prob = 0.5  \n    mixup_alpha = 0.5  \n    \n    def update_debug_settings(self):\n        if self.debug:\n            self.epochs = 15\n            self.selected_folds = [0, 1, 2, 3, 4]\n\ncfg = CFG()\ntaxonomy_df = pd.read_csv(cfg.taxonomy_csv)\ncfg.num_classes = len(taxonomy_df)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-11T14:02:51.149737Z","iopub.execute_input":"2025-08-11T14:02:51.150279Z","iopub.status.idle":"2025-08-11T14:02:51.243991Z","shell.execute_reply.started":"2025-08-11T14:02:51.150254Z","shell.execute_reply":"2025-08-11T14:02:51.243028Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Utilities","metadata":{}},{"cell_type":"code","source":"def set_seed(seed=42):\n    \"\"\"\n    Set seed for reproducibility\n    \"\"\"\n    random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\nset_seed(cfg.seed)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-11T14:02:51.245805Z","iopub.execute_input":"2025-08-11T14:02:51.246039Z","iopub.status.idle":"2025-08-11T14:02:51.254983Z","shell.execute_reply.started":"2025-08-11T14:02:51.246019Z","shell.execute_reply":"2025-08-11T14:02:51.254261Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model Definition (pretrained BEATs)","metadata":{}},{"cell_type":"code","source":"import sys\nsys.path.append(cfg.models_code_path)\n\nfrom BEATs import BEATs, BEATsConfig\n\nclass BEATs_model_classifier(torch.nn.Module):\n    \n    def __init__(self, beats_model_pretrained: BEATs, cfg, hidden_dim: int = 512):\n\n        super().__init__()\n\n        self.cfg = cfg\n        \n        self.beats = beats_model_pretrained\n        beats_dim = beats_model_pretrained.cfg.encoder_embed_dim\n        \n        self.temporal_attention = nn.MultiheadAttention(\n            beats_dim, num_heads=8, dropout=0.5, batch_first=True\n        )\n\n        self.pooling = nn.AdaptiveAvgPool1d(1)\n        \n        self.classifier = nn.Sequential(\n            nn.LayerNorm(beats_dim),\n            nn.Linear(beats_dim, hidden_dim),\n            nn.GELU(),\n            nn.Dropout(0.4),\n            nn.Linear(hidden_dim, hidden_dim // 2),\n            nn.GELU(),\n            nn.Dropout(0.3),\n            nn.Linear(hidden_dim // 2, cfg.num_classes) \n        )\n        \n    \n        for param in self.beats.parameters():\n            param.requires_grad = False\n    \n        \n    def forward(self, waveforms, padding_mask=None):\n\n        beats_features, _ = self.beats.extract_features(waveforms, padding_mask=padding_mask)\n\n        beats_features, features_padding_mask  = self.beats.extract_features(waveforms, padding_mask=padding_mask)\n        \n        #####\n        # Step 2: Apply the temporal self-attention layer\n        # The key_padding_mask ensures attention ignores padded time steps.\n        attended_features, _ = self.temporal_attention(\n            query=beats_features,\n            key=beats_features,\n            value=beats_features,\n            key_padding_mask=features_padding_mask\n        )\n        \n        # Step 3: Perform masked pooling on the *attended* features\n        if features_padding_mask is None:\n            pooled = attended_features.mean(dim=1)\n        else:\n            # Zero out the feature vectors at padded timesteps\n            attended_features = attended_features.masked_fill(\n                features_padding_mask.unsqueeze(-1), 0.0\n            )\n            \n            # Sum the features along the time dimension\n            summed_features = attended_features.sum(dim=1)\n            \n            # Calculate the number of non-padded (actual) timesteps\n            actual_lengths = (~features_padding_mask).sum(dim=1, keepdim=True)\n            actual_lengths = actual_lengths.clamp(min=1e-9)\n            \n            # Divide the sum by the actual lengths to get the true mean\n            pooled = summed_features / actual_lengths\n        # attended, _ = self.multiheadAttentions(\n        #     beats_features, beats_features, beats_features\n        # )\n        \n        # # attended: [B, T, D] -> [B, D, T] -> [B, D, 1] -> [B, D]\n        # pooled = self.pooling(attended.transpose(1, 2)).squeeze(-1)\n        #pooled = self.pooling(beats_features).squeeze(-1)\n        #beats_features = beats_features.mean(dim=1)\n        #print(beats_features.shape)\n        #logits = self.classifier(beats_features)\n        # # Classification\n        logits = self.classifier(pooled)\n        #logits = self.classifier(beats_features)\n        \n        return logits#.squeeze(-1)  # [B] for binary classification\n\n\ndef load_model():\n    # load the pre-trained checkpoints\n    checkpoint = torch.load(cfg.models_weights_path)\n    \n    cfg_model = BEATsConfig(checkpoint['cfg'])\n    BEATs_model = BEATs(cfg_model)\n    BEATs_model.load_state_dict(checkpoint['model'])\n    BEATs_model.eval()\n\n\n    classifier = BEATs_model_classifier(BEATs_model, cfg)\n\n    return classifier","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-11T14:02:51.256428Z","iopub.execute_input":"2025-08-11T14:02:51.256658Z","iopub.status.idle":"2025-08-11T14:02:51.289167Z","shell.execute_reply.started":"2025-08-11T14:02:51.256639Z","shell.execute_reply":"2025-08-11T14:02:51.288557Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Dataset Preparation and Data Augmentations\nWe'll convert audio to mel spectrograms and apply random augmentations with 50% probability each - including time stretching, pitch shifting, and volume adjustments. This randomized approach creates diverse training samples from the same audio files","metadata":{}},{"cell_type":"code","source":"\nclass BirdCLEFDataset(Dataset):\n    def __init__(self, df, cfg, mode=\"train\", max_length=32000):\n        self.df = df\n        self.cfg = cfg\n        self.mode = mode\n        self.max_length = max_length\n        self.sample_rate = 16000\n        \n        # Load taxonomy\n        taxonomy_df = pd.read_csv(self.cfg.taxonomy_csv)\n        self.species_ids = taxonomy_df['primary_label'].tolist()\n        self.num_classes = len(self.species_ids)\n        self.label_to_idx = {label: idx for idx, label in enumerate(self.species_ids)}\n        \n        # Prepare file paths\n        if 'filepath' not in self.df.columns:\n            self.df['filepath'] = self.cfg.train_datadir + '/' + self.df.filename\n        \n        if 'samplename' not in self.df.columns:\n            self.df['samplename'] = self.df.filename.map(\n                lambda x: x.split('/')[0] + '-' + x.split('/')[-1].split('.')[0]\n            )\n        \n        # Debug mode\n        if cfg.debug:\n            self.df = self.df.sample(min(1000, len(self.df)), random_state=cfg.seed).reset_index(drop=True)\n        \n        # Initialize audio augmentations\n        self._init_augmentations()\n    \n    def _init_augmentations(self):\n        \"\"\"Initialize all audio augmentation transforms\"\"\"\n        \n        # Time domain augmentations\n        self.time_shift_prob = 0.5\n        self.time_shift_range = 0.2  # Fraction of audio length\n        \n        self.pitch_shift_prob = 0.3\n        self.pitch_shift_range = 4  # Semitones\n        \n        self.time_stretch_prob = 0.3\n        self.time_stretch_range = (0.8, 1.25)\n        \n        self.noise_injection_prob = 0.4\n        self.noise_levels = (0.001, 0.01)\n        \n        self.gain_prob = 0.5\n        self.gain_range = (0.7, 1.3)\n        \n        # Frequency domain augmentations\n        self.freq_mask_prob = 0.6\n        self.freq_mask_param = 15\n        \n        self.time_mask_prob = 0.6\n        self.time_mask_param = 35\n        \n        # Spectral augmentations\n        self.spec_cutout_prob = 0.3\n        self.mixup_prob = 0.2\n        \n        # Background noise samples for realistic augmentation\n        self.background_noise_prob = 0.3\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        \n        # Load audio\n        waveform, original_sr = torchaudio.load(row['filepath'])\n        \n        # Resample if needed\n        if original_sr != self.sample_rate:\n            resampler = torchaudio.transforms.Resample(original_sr, self.sample_rate)\n            waveform = resampler(waveform)\n        \n        # Convert to mono\n        if waveform.shape[0] > 1:\n            waveform = waveform.mean(dim=0, keepdim=True)\n        \n        # Apply audio augmentations (only during training)\n        if self.mode == \"train\":\n            waveform = self.apply_audio_augmentations(waveform.squeeze(0))\n            waveform = waveform.unsqueeze(0)\n        \n        # Handle length adjustment\n        waveform, padding_mask = self._adjust_length(waveform)\n\n        #target = row.primary_label\n        target = np.array([1 if item == row.primary_label else 0 for item in class_labels])\n        # Encode labels\n        #target = self.encode_label(row['primary_label'])\n        \n        # # Handle secondary labels\n        # if 'secondary_labels' in row and row['secondary_labels'] not in [[''], None, np.nan]:\n        #     if isinstance(row['secondary_labels'], str):\n        #         secondary_labels = eval(row['secondary_labels'])\n        #     else:\n        #         secondary_labels = row['secondary_labels']\n            \n        #     for label in secondary_labels:\n        #         if label in self.label_to_idx:\n        #             target[self.label_to_idx[label]] = 1.0\n        \n        return {\n            'waveform': waveform.squeeze(0),\n            'padding_mask': padding_mask,\n            'target': torch.tensor(target, dtype=torch.float32),\n            'filename': row['filename']\n        }\n    \n    def apply_audio_augmentations(self, waveform):\n        \"\"\"Apply advanced audio augmentations\"\"\"\n        \n        # 1. Time shifting\n        if random.random() < self.time_shift_prob:\n            waveform = self._time_shift(waveform)\n        \n        # 2. Gain/Volume adjustment\n        if random.random() < self.gain_prob:\n            waveform = self._apply_gain(waveform)\n        \n        # 3. Add background noise\n        if random.random() < self.noise_injection_prob:\n            waveform = self._add_noise(waveform)\n        \n        # 4. Pitch shifting (computationally expensive, use sparingly)\n        if random.random() < self.pitch_shift_prob:\n            waveform = self._pitch_shift(waveform)\n        \n        # 5. Time stretching\n        if random.random() < self.time_stretch_prob:\n            waveform = self._time_stretch(waveform)\n        \n        # 6. Band-pass filtering\n        if random.random() < 0.2:\n            waveform = self._band_pass_filter(waveform)\n        \n        # 7. Add realistic background noise\n        if random.random() < self.background_noise_prob:\n            waveform = self._add_background_noise(waveform)\n        \n        return waveform\n    \n    def _time_shift(self, waveform):\n        \"\"\"Shift audio in time\"\"\"\n        shift_amount = int(random.uniform(-self.time_shift_range, self.time_shift_range) * len(waveform))\n        if shift_amount > 0:\n            # Shift right (delay)\n            shifted = torch.cat([torch.zeros(shift_amount), waveform[:-shift_amount]])\n        elif shift_amount < 0:\n            # Shift left (advance)\n            shifted = torch.cat([waveform[-shift_amount:], torch.zeros(-shift_amount)])\n        else:\n            shifted = waveform\n        return shifted\n    \n    def _apply_gain(self, waveform):\n        \"\"\"Apply random gain\"\"\"\n        gain = random.uniform(*self.gain_range)\n        return waveform * gain\n    \n    def _add_noise(self, waveform):\n        \"\"\"Add Gaussian noise\"\"\"\n        noise_level = random.uniform(*self.noise_levels)\n        noise = torch.randn_like(waveform) * noise_level\n        return waveform + noise\n    \n    def _pitch_shift(self, waveform):\n        \"\"\"Pitch shift using librosa\"\"\"\n        waveform_np = waveform.numpy()\n        n_steps = random.uniform(-self.pitch_shift_range, self.pitch_shift_range)\n        \n        # Use librosa for pitch shifting\n        shifted = librosa.effects.pitch_shift(\n            waveform_np, \n            sr=self.sample_rate, \n            n_steps=n_steps\n        )\n        return torch.from_numpy(shifted).float()\n    \n    def _time_stretch(self, waveform):\n        \"\"\"Time stretching using librosa\"\"\"\n        waveform_np = waveform.numpy()\n        rate = random.uniform(*self.time_stretch_range)\n        \n        # Use librosa for time stretching\n        stretched = librosa.effects.time_stretch(waveform_np, rate=rate)\n        \n        # Ensure output length matches input\n        if len(stretched) > len(waveform_np):\n            stretched = stretched[:len(waveform_np)]\n        elif len(stretched) < len(waveform_np):\n            padding = len(waveform_np) - len(stretched)\n            stretched = np.pad(stretched, (0, padding), mode='constant')\n        \n        return torch.from_numpy(stretched).float()\n    \n    def _band_pass_filter(self, waveform):\n        \"\"\"Apply band-pass filter to simulate different recording conditions\"\"\"\n        # Random frequency range for bird calls (typically 1kHz - 8kHz)\n        low_freq = random.uniform(500, 2000)\n        high_freq = random.uniform(4000, 8000)\n        \n        # Ensure high > low\n        if high_freq <= low_freq:\n            high_freq = low_freq + 2000\n        \n        # Create filter\n        nyquist = self.sample_rate / 2\n        low_norm = low_freq / nyquist\n        high_norm = min(high_freq / nyquist, 0.99)\n        \n        # Design butterworth filter\n        b, a = scipy.signal.butter(4, [low_norm, high_norm], btype='band')\n        \n        # Apply filter\n        filtered = scipy.signal.filtfilt(b, a, waveform.numpy())\n        return torch.from_numpy(filtered.copy()).float()\n    \n    def _add_background_noise(self, waveform):\n        \"\"\"Add realistic background noise (wind, rain, traffic)\"\"\"\n        # Generate different types of background noise\n        noise_type = random.choice(['white', 'pink', 'brown', 'wind', 'rain'])\n        \n        if noise_type == 'white':\n            noise = torch.randn_like(waveform)\n        elif noise_type == 'pink':\n            # Pink noise (1/f noise)\n            noise = self._generate_pink_noise(len(waveform))\n        elif noise_type == 'brown':\n            # Brown noise (1/f^2 noise)\n            noise = self._generate_brown_noise(len(waveform))\n        elif noise_type == 'wind':\n            # Low frequency rumble\n            noise = self._generate_wind_noise(len(waveform))\n        else:  # rain\n            # High frequency patter\n            noise = self._generate_rain_noise(len(waveform))\n        \n        # Mix with original signal\n        noise_level = random.uniform(0.01, 0.05)\n        return waveform + noise * noise_level\n    \n    def _generate_pink_noise(self, length):\n        \"\"\"Generate pink noise (1/f spectrum)\"\"\"\n        # Simple approximation of pink noise\n        white_noise = torch.randn(length)\n        # Apply simple low-pass filtering\n        b, a = scipy.signal.butter(1, 0.1, btype='low')\n        pink_noise = scipy.signal.filtfilt(b, a, white_noise.numpy())\n        return torch.from_numpy(pink_noise.copy()).float()\n    \n    def _generate_brown_noise(self, length):\n        \"\"\"Generate brown noise (1/f^2 spectrum)\"\"\"\n        white_noise = torch.randn(length)\n        # Apply stronger low-pass filtering\n        b, a = scipy.signal.butter(2, 0.05, btype='low')\n        brown_noise = scipy.signal.filtfilt(b, a, white_noise.numpy())\n        return torch.from_numpy(brown_noise.copy()).float()\n    \n    def _generate_wind_noise(self, length):\n        \"\"\"Generate wind-like noise (low frequency)\"\"\"\n        # Low frequency rumble\n        t = torch.arange(length) / self.sample_rate\n        wind = 0.5 * torch.sin(2 * np.pi * 3 * t) + 0.3 * torch.sin(2 * np.pi * 7 * t)\n        wind += 0.2 * torch.randn(length)  # Add some randomness\n        return wind\n    \n    def _generate_rain_noise(self, length):\n        \"\"\"Generate rain-like noise (high frequency droplets)\"\"\"\n        # Random impulses to simulate raindrops\n        rain = torch.zeros(length)\n        num_drops = random.randint(length // 100, length // 50)\n        \n        for _ in range(num_drops):\n            pos = random.randint(0, length - 1)\n            intensity = random.uniform(0.1, 0.5)\n            rain[pos] = intensity\n        \n        # Apply some smoothing\n        if len(rain) > 10:\n            rain = F.conv1d(\n                rain.unsqueeze(0).unsqueeze(0),\n                torch.ones(1, 1, 5) / 5,\n                padding=2\n            ).squeeze()\n        \n        return rain\n    \n    def apply_spectral_augmentations(self, spectrogram):\n        \"\"\"Apply spectral augmentations (call this on spectrograms if needed)\"\"\"\n        spec = spectrogram.clone()\n        \n        # Frequency masking\n        if random.random() < self.freq_mask_prob:\n            spec = self._freq_mask(spec)\n        \n        # Time masking\n        if random.random() < self.time_mask_prob:\n            spec = self._time_mask(spec)\n        \n        # Spectral cutout\n        if random.random() < self.spec_cutout_prob:\n            spec = self._spectral_cutout(spec)\n        \n        return spec\n    \n    def _freq_mask(self, spec):\n        \"\"\"Apply frequency masking\"\"\"\n        freq_dim = spec.shape[-2]\n        mask_size = random.randint(1, self.freq_mask_param)\n        mask_start = random.randint(0, freq_dim - mask_size)\n        \n        spec[..., mask_start:mask_start + mask_size, :] = 0\n        return spec\n    \n    def _time_mask(self, spec):\n        \"\"\"Apply time masking\"\"\"\n        time_dim = spec.shape[-1]\n        mask_size = random.randint(1, min(self.time_mask_param, time_dim))\n        mask_start = random.randint(0, time_dim - mask_size)\n        \n        spec[..., :, mask_start:mask_start + mask_size] = 0\n        return spec\n    \n    def _spectral_cutout(self, spec):\n        \"\"\"Apply random rectangular cutouts\"\"\"\n        h, w = spec.shape[-2:]\n        \n        # Random rectangular cutout\n        cut_h = random.randint(5, h // 4)\n        cut_w = random.randint(5, w // 4)\n        \n        start_h = random.randint(0, h - cut_h)\n        start_w = random.randint(0, w - cut_w)\n        \n        spec[..., start_h:start_h + cut_h, start_w:start_w + cut_w] = 0\n        return spec\n    \n    def _adjust_length(self, waveform):\n        \"\"\"Adjust waveform length and create padding mask\"\"\"\n        current_length = waveform.shape[1]\n        \n        if current_length > self.max_length:\n            if self.mode == \"train\":\n                # Random crop during training\n                start = torch.randint(0, current_length - self.max_length + 1, (1,)).item()\n            else:\n                # Center crop during validation/test\n                start = (current_length - self.max_length) // 2\n            waveform = waveform[:, start:start + self.max_length]\n            padding_mask = torch.zeros(self.max_length, dtype=torch.bool)\n            \n        elif current_length < self.max_length:\n            # Pad with zeros\n            padding = self.max_length - current_length\n            waveform = F.pad(waveform, (0, padding))\n            # Create padding mask: True for padded positions\n            padding_mask = torch.zeros(self.max_length, dtype=torch.bool)\n            padding_mask[current_length:] = True\n            \n        else:\n            # Exact length\n            padding_mask = torch.zeros(self.max_length, dtype=torch.bool)\n        \n        return waveform, padding_mask\n    \n    def encode_label(self, label):\n        \"\"\"Encode label to one-hot vector\"\"\"\n        target = np.zeros(self.num_classes)\n        if label in self.label_to_idx:\n            target[self.label_to_idx[label]] = 1.0\n        return target\n    \n    def mixup_data(self, x1, y1, x2, y2, alpha=0.2):\n        \"\"\"Mixup augmentation for combining samples\"\"\"\n        if alpha > 0:\n            lam = np.random.beta(alpha, alpha)\n        else:\n            lam = 1\n        \n        mixed_x = lam * x1 + (1 - lam) * x2\n        mixed_y = lam * y1 + (1 - lam) * y2\n        \n        return mixed_x, mixed_y, lam\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-11T14:02:51.289884Z","iopub.execute_input":"2025-08-11T14:02:51.290127Z","iopub.status.idle":"2025-08-11T14:02:51.319634Z","shell.execute_reply.started":"2025-08-11T14:02:51.290097Z","shell.execute_reply":"2025-08-11T14:02:51.318797Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def collate_fn(batch):\n    \"\"\"Custom collate function to handle different sized spectrograms\"\"\"\n    batch = [item for item in batch if item is not None]\n    if len(batch) == 0:\n        return {}\n        \n    result = {key: [] for key in batch[0].keys()}\n    \n    for item in batch:\n        for key, value in item.items():\n            result[key].append(value)\n    \n    for key in result:\n        if isinstance(result[key][0], torch.Tensor):\n            result[key] = torch.stack(result[key])\n    return result","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-11T14:02:51.320470Z","iopub.execute_input":"2025-08-11T14:02:51.320688Z","iopub.status.idle":"2025-08-11T14:02:51.339458Z","shell.execute_reply.started":"2025-08-11T14:02:51.320659Z","shell.execute_reply":"2025-08-11T14:02:51.338727Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training Utilities\nWe are configuring our optimization strategy with the AdamW optimizer, cosine scheduling, and the BCEWithLogitsLoss criterion.","metadata":{}},{"cell_type":"code","source":"def get_optimizer(model, cfg):\n  \n    if cfg.optimizer == 'Adam':\n        optimizer = optim.Adam(\n            model.parameters(),\n            lr=cfg.lr,\n            weight_decay=cfg.weight_decay\n        )\n    elif cfg.optimizer == 'AdamW':\n        optimizer = optim.AdamW(\n            model.parameters(),\n            lr=cfg.lr,\n            weight_decay=cfg.weight_decay\n        )\n    elif cfg.optimizer == 'SGD':\n        optimizer = optim.SGD(\n            model.parameters(),\n            lr=cfg.lr,\n            momentum=0.9,\n            weight_decay=cfg.weight_decay\n        )\n    else:\n        raise NotImplementedError(f\"Optimizer {cfg.optimizer} not implemented\")\n        \n    return optimizer\n\ndef get_scheduler(optimizer, cfg):\n   \n    if cfg.scheduler == 'CosineAnnealingLR':\n        scheduler = lr_scheduler.CosineAnnealingLR(\n            optimizer,\n            T_max=cfg.T_max,\n            eta_min=cfg.min_lr\n        )\n    elif cfg.scheduler == 'ReduceLROnPlateau':\n        scheduler = lr_scheduler.ReduceLROnPlateau(\n            optimizer,\n            mode='min',\n            factor=0.5,\n            patience=2,\n            min_lr=cfg.min_lr,\n            verbose=True\n        )\n    elif cfg.scheduler == 'StepLR':\n        scheduler = lr_scheduler.StepLR(\n            optimizer,\n            step_size=cfg.epochs // 3,\n            gamma=0.5\n        )\n    elif cfg.scheduler == 'OneCycleLR':\n        scheduler = None  \n    else:\n        scheduler = None\n        \n    return scheduler\n\ndef get_criterion(cfg):\n \n    if cfg.criterion == 'BCEWithLogitsLoss':\n        criterion = nn.BCEWithLogitsLoss()\n    else:\n        raise NotImplementedError(f\"Criterion {cfg.criterion} not implemented\")\n    criterion = nn.CrossEntropyLoss()\n    return criterion","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-11T14:02:51.340193Z","iopub.execute_input":"2025-08-11T14:02:51.340465Z","iopub.status.idle":"2025-08-11T14:02:51.366447Z","shell.execute_reply.started":"2025-08-11T14:02:51.340438Z","shell.execute_reply":"2025-08-11T14:02:51.365737Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training Loop","metadata":{}},{"cell_type":"code","source":"import sys\nsys.path.append('/kaggle/usr/lib/kaggle_metric_utilities')\nsys.path.append('/kaggle/usr/lib/birdclef-roc-auc')\nfrom metric import score\nclass_labels = sorted(os.listdir('../input/birdclef-2025/train_audio/'))\n\ndef cal_score(label, pred):\n    label = np.concatenate(label)\n    pred = np.concatenate(pred)\n\n    label_df = pd.DataFrame(label>0.5, columns=class_labels)\n    pred_df = pd.DataFrame(pred, columns=class_labels)\n    label_df['id'] = np.arange(len(label_df))\n    pred_df['id'] = np.arange(len(pred_df))\n\n    return score(label_df, pred_df, row_id_column_name='id')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-11T14:02:51.367257Z","iopub.execute_input":"2025-08-11T14:02:51.367456Z","iopub.status.idle":"2025-08-11T14:02:51.411012Z","shell.execute_reply.started":"2025-08-11T14:02:51.367439Z","shell.execute_reply":"2025-08-11T14:02:51.410472Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_one_epoch(model, loader, optimizer, criterion, device, scheduler=None):\n    \n    model.train()\n    losses = []\n    all_targets = []\n    all_outputs = []\n    \n    pbar = tqdm(enumerate(loader), total=len(loader), desc=\"Training\")\n    \n    for step, batch in pbar:\n    \n        inputs = batch['waveform'].to(device)\n        padding_masks = batch['padding_mask'].to(device)\n        targets = batch['target'].to(device)\n        \n        optimizer.zero_grad()\n        logits = model(inputs, padding_masks)\n\n        \n        loss = criterion(logits, targets)\n            \n        loss.backward()\n        optimizer.step()\n        \n        #outputs = outputs.detach().cpu().numpy()\n        outputs = torch.softmax(logits,dim=1).detach().cpu().numpy() \n        targets = targets.detach().cpu().numpy()\n        \n        if scheduler is not None and isinstance(scheduler, lr_scheduler.OneCycleLR):\n            scheduler.step()\n\n        \n        #print(outputs.shape)\n        #print(targets.shape)\n        all_outputs.append(outputs)\n        all_targets.append(targets)\n        auc = cal_score(all_targets, all_outputs)\n        \n        losses.append(loss.item())\n        \n        pbar.set_postfix({\n            'train_loss': np.mean(losses[-10:]) if losses else 0,\n            'lr': optimizer.param_groups[0]['lr'],\n            'auc': auc\n        })\n    \n    #all_outputs = np.concatenate(all_outputs)\n    #all_targets = np.concatenate(all_targets)\n    #auc = calculate_auc(all_targets, all_outputs)\n    auc = cal_score(all_targets, all_outputs)\n    avg_loss = np.mean(losses)\n    \n    return avg_loss, auc\n\ndef validate(model, loader, criterion, device):\n   \n    model.eval()\n    losses = []\n    all_targets = []\n    all_outputs = []\n    \n    with torch.no_grad():\n        for batch in tqdm(loader, desc=\"Validation\"):\n                \n           \n            inputs = batch['waveform'].to(device)\n            padding_masks = batch['padding_mask'].to(device)\n            targets = batch['target'].to(device)\n            \n            logits = model(inputs, padding_masks)\n\n            \n            loss = criterion(logits, targets)\n            \n            outputs = torch.softmax(logits, dim=1).detach().cpu().numpy()\n            targets = targets.detach().cpu().numpy()\n            \n            all_outputs.append(outputs)\n            all_targets.append(targets)\n            losses.append(loss.item())\n    \n    #all_outputs = np.concatenate(all_outputs)\n    #all_targets = np.concatenate(all_targets)\n    \n    #auc = calculate_auc(all_targets, all_outputs)\n    auc = cal_score(all_targets, all_outputs)\n    avg_loss = np.mean(losses)\n    \n    return avg_loss, auc\n\ndef calculate_auc(targets, outputs):\n  \n    num_classes = targets.shape[1]\n    aucs = []\n    \n    probs = 1 / (1 + np.exp(-outputs))\n    \n    for i in range(num_classes):\n        \n        if np.sum(targets[:, i]) > 0:\n            class_auc = roc_auc_score(targets[:, i], probs[:, i])\n            aucs.append(class_auc)\n    \n    return np.mean(aucs) if aucs else 0.0","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-11T14:02:51.412568Z","iopub.execute_input":"2025-08-11T14:02:51.412767Z","iopub.status.idle":"2025-08-11T14:02:51.422194Z","shell.execute_reply.started":"2025-08-11T14:02:51.412750Z","shell.execute_reply":"2025-08-11T14:02:51.421364Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training!","metadata":{}},{"cell_type":"code","source":"def run_training(df, cfg):\n    \"\"\"Training function that can either use pre-computed spectrograms or generate them on-the-fly\"\"\"\n\n    taxonomy_df = pd.read_csv(cfg.taxonomy_csv)\n    species_ids = taxonomy_df['primary_label'].tolist()\n    cfg.num_classes = len(species_ids)\n    \n    #if cfg.debug:\n    #    cfg.update_debug_settings()\n\n            \n    skf = StratifiedKFold(n_splits=cfg.n_fold, shuffle=True, random_state=cfg.seed)\n    \n    best_scores = []\n    \n    for fold, (train_idx, val_idx) in enumerate(skf.split(df, df['primary_label'])):\n        if fold not in cfg.selected_folds:\n            continue\n            \n        print(f'\\n{\"=\"*30} Fold {fold} {\"=\"*30}')\n        \n        train_df = df.iloc[train_idx].reset_index(drop=True)\n        val_df = df.iloc[val_idx].reset_index(drop=True)\n        \n        print(f'Training set: {len(train_df)} samples')\n        print(f'Validation set: {len(val_df)} samples')\n        \n        train_dataset = BirdCLEFDataset(train_df, cfg, mode='train')\n        val_dataset = BirdCLEFDataset(val_df, cfg, mode='valid')\n        \n        train_loader = DataLoader(\n            train_dataset, \n            batch_size=cfg.batch_size, \n            shuffle=True, \n            num_workers=cfg.num_workers,\n            pin_memory=True,\n            collate_fn=collate_fn,\n            drop_last=True\n        )\n        \n        val_loader = DataLoader(\n            val_dataset, \n            batch_size=cfg.batch_size, \n            shuffle=False, \n            num_workers=cfg.num_workers,\n            pin_memory=True,\n            collate_fn=collate_fn\n        )\n        \n        #model = BirdCLEFModel(cfg).to(cfg.device)\n        model = load_model().to(cfg.device)\n        optimizer = get_optimizer(model, cfg)\n        criterion = get_criterion(cfg)\n        \n        if cfg.scheduler == 'OneCycleLR':\n            scheduler = lr_scheduler.OneCycleLR(\n                optimizer,\n                max_lr=cfg.lr,\n                steps_per_epoch=len(train_loader),\n                epochs=cfg.epochs,\n                pct_start=0.1\n            )\n        else:\n            scheduler = get_scheduler(optimizer, cfg)\n        \n        best_auc = 0\n        best_epoch = 0\n        \n        for epoch in range(cfg.epochs):\n            print(f\"\\nEpoch {epoch+1}/{cfg.epochs}\")\n            \n            train_loss, train_auc = train_one_epoch(\n                model, \n                train_loader, \n                optimizer, \n                criterion, \n                cfg.device,\n                scheduler if isinstance(scheduler, lr_scheduler.OneCycleLR) else None\n            )\n            \n            val_loss, val_auc = validate(model, val_loader, criterion, cfg.device)\n\n            if scheduler is not None and not isinstance(scheduler, lr_scheduler.OneCycleLR):\n                if isinstance(scheduler, lr_scheduler.ReduceLROnPlateau):\n                    scheduler.step(val_loss)\n                else:\n                    scheduler.step()\n\n            print(f\"Train Loss: {train_loss:.4f}, Train AUC: {train_auc:.4f}\")\n            print(f\"Val Loss: {val_loss:.4f}, Val AUC: {val_auc:.4f}\")\n            \n            if val_auc > best_auc:\n                best_auc = val_auc\n                best_epoch = epoch + 1\n                print(f\"New best AUC: {best_auc:.4f} at epoch {best_epoch}\")\n\n                torch.save({\n                    'model_state_dict': model.state_dict(),\n                    'optimizer_state_dict': optimizer.state_dict(),\n                    'scheduler_state_dict': scheduler.state_dict() if scheduler else None,\n                    'epoch': epoch,\n                    'val_auc': val_auc,\n                    'train_auc': train_auc,\n                    'cfg': cfg\n                }, f\"model_fold{fold}.pth\")\n        \n        best_scores.append(best_auc)\n        print(f\"\\nBest AUC for fold {fold}: {best_auc:.4f} at epoch {best_epoch}\")\n        \n        # Clear memory\n        del model, optimizer, scheduler, train_loader, val_loader\n        torch.cuda.empty_cache()\n        gc.collect()\n    \n    print(\"\\n\" + \"=\"*60)\n    print(\"Cross-Validation Results:\")\n    for fold, score in enumerate(best_scores):\n        print(f\"Fold {cfg.selected_folds[fold]}: {score:.4f}\")\n    print(f\"Mean AUC: {np.mean(best_scores):.4f}\")\n    print(\"=\"*60)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-11T14:02:51.423237Z","iopub.execute_input":"2025-08-11T14:02:51.423524Z","iopub.status.idle":"2025-08-11T14:02:51.440789Z","shell.execute_reply.started":"2025-08-11T14:02:51.423498Z","shell.execute_reply":"2025-08-11T14:02:51.440158Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if __name__ == \"__main__\":\n    import time\n    \n    print(\"\\nLoading training data...\")\n    train_df = pd.read_csv(cfg.train_csv)\n    taxonomy_df = pd.read_csv(cfg.taxonomy_csv)\n\n    print(\"\\nStarting training...\")\n    \n    run_training(train_df, cfg)\n    \n    print(\"\\nTraining complete!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-11T14:02:51.441649Z","iopub.execute_input":"2025-08-11T14:02:51.441943Z"},"scrolled":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}