{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"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":"nvidiaTeslaT4","dataSources":[{"sourceId":91844,"databundleVersionId":11361821,"sourceType":"competition"},{"sourceId":11053663,"sourceType":"datasetVersion","datasetId":6886569}],"dockerImageVersionId":31041,"isInternetEnabled":true,"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 EfficientNetB0 with PyTorch and Timm(for pretrained EffNet). You can check inference and preprocessing notebooks in the following links: \n\n- [EfficientNet B0 Pytorch [Inference] | BirdCLEF'25](https://www.kaggle.com/code/kadircandrisolu/efficientnet-b0-pytorch-inference-birdclef-25)\n\n  \n- [Transforming Audio-to-Mel Spec. | BirdCLEF'25](https://www.kaggle.com/code/kadircandrisolu/transforming-audio-to-mel-spec-birdclef-25)  \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/datasets/kadircandrisolu/birdclef25-effnetb0-starter-weight) I used in the inference notebook was obtained after 10 epochs of training.\n\n**Features**\n* Implement with Pytorch and Timm\n* Flexible audio processing with both pre-computed and on-the-fly mel spectrograms\n* Stratified 5-fold cross-validation with ensemble capability\n* Mixup training for improved generalization\n* Spectrogram augmentations (time/frequency masking, brightness adjustment)\n* AdamW optimizer with Cosine Annealing LR scheduling\n* Debug mode for quick experimentation with smaller datasets\n\n**Pre-computed Spectrograms**\nFor faster training, you can use pre-computed mel spectrograms from [this dataset](https://www.kaggle.com/datasets/kadircandrisolu/birdclef25-mel-spectrograms) by setting `LOAD_DATA = True`","metadata":{}},{"cell_type":"code","source":"# **BirdCLEF 2025 Training Notebook**\n\n# This is a baseline training pipeline for BirdCLEF 2025 using EfficientNetB0 with PyTorch and Timm.\n# Modifications include precision-recall curves, micro/macro metrics, and enhanced visualizations for final results.\n\nimport os\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, f1_score, precision_recall_curve, precision_recall_fscore_support, confusion_matrix, roc_curve, auc\nimport librosa\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nfrom torch.optim import lr_scheduler\n\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom tqdm.auto import tqdm\nimport timm\n\nwarnings.filterwarnings(\"ignore\")\n\n## Configuration\nclass CFG:\n    seed = 42\n    debug = False\n    apex = False\n    print_freq = 100\n    num_workers = 4\n    \n    OUTPUT_DIR = '/kaggle/working/'\n    train_datadir = '/kaggle/input/birdclef-2025/train_audio'\n    train_csv = '/kaggle/input/birdclef-2025/train.csv'\n    submission_csv = '/kaggle/input/birdclef-2025/sample_submission.csv'\n    taxonomy_csv = '/kaggle/input/birdclef-2025/taxonomy.csv'\n    spectrogram_npy = '/kaggle/input/birdclef25-mel-spectrograms/birdclef2025_melspec_5sec_256_256.npy'\n \n    model_name = 'efficientnet_b0'\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_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 = 12\n    batch_size = 32\n    criterion = 'FocalLoss'\n\n    n_fold = 5\n    selected_folds = [0, 1, 2, 3, 4]\n\n    optimizer = 'AdamW'\n    lr = 3e-4\n    weight_decay = 1e-5\n    scheduler = 'CosineAnnealingLR'\n    min_lr = 1e-6\n    T_max = 12\n\n    aug_prob = 0.9\n    mixup_alpha = 1.5\n    \n    def update_debug_settings(self):\n        if self.debug:\n            self.epochs = 2\n            self.selected_folds = [0]\n\ncfg = CFG()\n\n## Utilities\ndef set_seed(seed=42):\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)\n\n## Pre-processing\ndef audio2melspec(audio_data, cfg):\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    mel_spec = librosa.feature.melspectrogram(\n        y=audio_data,\n        sr=cfg.FS,\n        n_fft=cfg.N_FFT,\n        hop_length=cfg.HOP_LENGTH,\n        n_mels=cfg.N_MELS,\n        fmin=cfg.FMIN,\n        fmax=cfg.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\n\ndef process_audio_file(audio_path, cfg, start_time=0):\n    try:\n        audio_data, _ = librosa.load(audio_path, sr=cfg.FS, offset=start_time, duration=cfg.TARGET_DURATION)\n        target_samples = int(cfg.TARGET_DURATION * cfg.FS)\n        if len(audio_data) < target_samples:\n            audio_data = np.pad(audio_data, (0, target_samples - len(audio_data)), mode='constant')\n        mel_spec = audio2melspec(audio_data, cfg)\n        if mel_spec.shape != cfg.TARGET_SHAPE:\n            mel_spec = cv2.resize(mel_spec, cfg.TARGET_SHAPE, interpolation=cv2.INTER_LINEAR)\n        return mel_spec.astype(np.float32)\n    except Exception as e:\n        print(f\"Error processing {audio_path}: {e}\")\n        return None\n\ndef generate_spectrograms(df, cfg):\n    print(\"Generating mel spectrograms...\")\n    start_time = time.time()\n    all_bird_data = {}\n    errors = []\n    for i, row in tqdm(df.iterrows(), total=len(df)):\n        if cfg.debug and i >= 1000:\n            break\n        try:\n            samplename = row['samplename']\n            filepath = row['filepath']\n            mel_spec = process_audio_file(filepath, cfg)\n            if mel_spec is not None:\n                all_bird_data[samplename] = mel_spec\n            else:\n                errors.append((filepath, \"Failed to generate spectrogram\"))\n        except Exception as e:\n            print(f\"Error processing {row.filepath}: {e}\")\n            errors.append((row.filepath, str(e)))\n    end_time = time.time()\n    print(f\"Processing completed in {end_time - start_time:.2f} seconds\")\n    print(f\"Successfully processed {len(all_bird_data)} files out of {len(df)}\")\n    print(f\"Failed to process {len(errors)} files\")\n    return all_bird_data\n\n## Dataset Preparation\nclass BirdCLEFDatasetFromNPY(Dataset):\n    def __init__(self, df, cfg, spectrograms=None, mode=\"train\"):\n        self.df = df.reset_index(drop=True)\n        self.cfg = cfg\n        self.mode = mode\n        self.spectrograms = spectrograms\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        if 'filepath' not in self.df.columns:\n            self.df['filepath'] = self.cfg.train_datadir + '/' + self.df.filename\n        if 'samplename' not in self.df.columns:\n            self.df['samplename'] = self.df.filename.map(lambda x: x.split('/')[0] + '-' + x.split('/')[-1].split('.')[0])\n        sample_names = set(self.df['samplename'])\n        if self.spectrograms:\n            found_samples = sum(1 for name in sample_names if name in self.spectrograms)\n            print(f\"Found {found_samples} matching spectrograms for {mode} dataset out of {len(self.df)} samples\")\n        if cfg.debug:\n            self.df = self.df.sample(min(1000, len(self.df)), random_state=cfg.seed).reset_index(drop=True)\n        print(f\"Dataset size: {len(self.df)} samples\")\n    \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        if idx >= len(self.df):\n            print(f\"Index {idx} out of bounds for dataset of size {len(self.df)}\")\n            return None\n        row = self.df.iloc[idx]\n        samplename = row['samplename']\n        spec = None\n        if self.spectrograms and samplename in self.spectrograms:\n            spec = self.spectrograms[samplename]\n        elif not self.cfg.LOAD_DATA:\n            spec = process_audio_file(row['filepath'], self.cfg)\n        if spec is None:\n            print(f\"Spectrogram for {samplename} not found, returning default\")\n            spec = np.zeros(self.cfg.TARGET_SHAPE, dtype=np.float32)\n        if spec.ndim != 2:\n            print(f\"Invalid spectrogram shape for {samplename}: {spec.shape}\")\n            spec = np.zeros(self.cfg.TARGET_SHAPE, dtype=np.float32)\n        spec = torch.tensor(spec, dtype=torch.float32).unsqueeze(0)\n        if self.mode == \"train\" and random.random() < self.cfg.aug_prob:\n            spec = self.apply_spec_augmentations(spec)\n        target = self.encode_label(row['primary_label'])\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            for label in secondary_labels:\n                if label in self.label_to_idx:\n                    target[self.label_to_idx[label]] = 1.0\n        return {\n            'melspec': spec,\n            'target': torch.tensor(target, dtype=torch.float32),\n            'filename': row['filename']\n        }\n    \n    def apply_spec_augmentations(self, spec):\n        if not isinstance(spec, torch.Tensor) or spec.ndim != 3:\n            print(f\"Invalid spec shape for augmentation: {spec.shape}\")\n            return spec\n        try:\n            if random.random() < 0.6:\n                num_masks = random.randint(1, 5)\n                for _ in range(num_masks):\n                    width = random.randint(5, 25)\n                    start = random.randint(0, spec.shape[2] - width)\n                    spec[:, :, start:start+width] = 0\n            if random.random() < 0.6:\n                num_masks = random.randint(1, 5)\n                for _ in range(num_masks):\n                    height = random.randint(5, 25)\n                    start = random.randint(0, spec.shape[1] - height)\n                    spec[:, start:start+height, :] = 0\n            if random.random() < 0.6:\n                gain = random.uniform(0.7, 1.3)\n                bias = random.uniform(-0.15, 0.15)\n                spec = spec * gain + bias\n                spec = torch.clamp(spec, 0, 1)\n            if random.random() < 0.4:\n                noise = torch.randn_like(spec) * 0.15\n                spec = spec + noise\n                spec = torch.clamp(spec, 0, 1)\n        except Exception as e:\n            print(f\"Error in spec augmentation: {e}\")\n        return spec\n    \n    def encode_label(self, label):\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\ndef collate_fn(batch):\n    batch = [item for item in batch if item is not None and item['melspec'].shape == (1, 256, 256)]\n    if len(batch) == 0:\n        print(\"Empty batch after filtering\")\n        return {}\n    result = {key: [] for key in batch[0].keys()}\n    for item in batch:\n        for key, value in item.items():\n            result[key].append(value)\n    for key in result:\n        try:\n            if key == 'target':\n                result[key] = torch.stack(result[key])\n            elif key == 'melspec':\n                result[key] = torch.stack(result[key])\n        except Exception as e:\n            print(f\"Error stacking {key}: {e}\")\n            return {}\n    return result\n\n## Model Definition\nclass BirdCLEFModel(nn.Module):\n    def __init__(self, cfg):\n        super().__init__()\n        self.cfg = cfg\n        taxonomy_df = pd.read_csv(cfg.taxonomy_csv)\n        cfg.num_classes = len(taxonomy_df)\n        self.backbone = timm.create_model(\n            cfg.model_name,\n            pretrained=cfg.pretrained,\n            in_chans=cfg.in_channels,\n            drop_rate=0.6,\n            drop_path_rate=0.6\n        )\n        if 'efficientnet' in cfg.model_name:\n            backbone_out = self.backbone.classifier.in_features\n            self.backbone.classifier = nn.Identity()\n        elif 'resnet' in cfg.model_name:\n            backbone_out = self.backbone.fc.in_features\n            self.backbone.fc = nn.Identity()\n        else:\n            backbone_out = self.backbone.get_classifier().in_features\n            self.backbone.reset_classifier(0, '')\n        self.pooling = nn.AdaptiveAvgPool2d(1)\n        self.feat_dim = backbone_out\n        self.classifier = nn.Linear(backbone_out, cfg.num_classes)\n        self.mixup_enabled = hasattr(cfg, 'mixup_alpha') and cfg.mixup_alpha > 0\n        if self.mixup_enabled:\n            self.mixup_alpha = cfg.mixup_alpha\n    \n    def forward(self, x, targets=None):\n        if self.training and self.mixup_enabled and targets is not None:\n            mixed_x, targets_a, targets_b, lam = self.mixup_data(x, targets)\n            x = mixed_x\n        else:\n            targets_a, targets_b, lam = None, None, None\n        features = self.backbone(x)\n        if isinstance(features, dict):\n            features = features['features']\n        if len(features.shape) == 4:\n            features = self.pooling(features)\n            features = features.view(features.size(0), -1)\n        logits = self.classifier(features)\n        if self.training and self.mixup_enabled and targets is not None:\n            loss = self.mixup_criterion(F.binary_cross_entropy_with_logits, \n                                       logits, targets_a, targets_b, lam)\n            return logits, loss\n        return logits\n    \n    def mixup_data(self, x, targets):\n        batch_size = x.size(0)\n        lam = np.random.beta(self.mixup_alpha, self.mixup_alpha)\n        indices = torch.randperm(batch_size).to(x.device)\n        mixed_x = lam * x + (1 - lam) * x[indices]\n        return mixed_x, targets, targets[indices], lam\n    \n    def mixup_criterion(self, criterion, pred, y_a, y_b, lam):\n        return lam * criterion(pred, y_a) + (1 - lam) * criterion(pred, y_b)\n\n## Training Utilities\ndef get_optimizer(model, cfg):\n    optimizer = optim.AdamW(\n        model.parameters(),\n        lr=cfg.lr,\n        weight_decay=cfg.weight_decay\n    )\n    return optimizer\n\ndef get_scheduler(optimizer, cfg):\n    scheduler = lr_scheduler.CosineAnnealingLR(\n        optimizer,\n        T_max=cfg.T_max,\n        eta_min=cfg.min_lr\n    )\n    return scheduler\n\nclass FocalLoss(nn.Module):\n    def __init__(self, alpha=1, gamma=3.5):\n        super().__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n    def forward(self, inputs, targets):\n        BCE_loss = F.binary_cross_entropy_with_logits(inputs, targets, reduction='none')\n        pt = torch.exp(-BCE_loss)\n        F_loss = self.alpha * (1-pt)**self.gamma * BCE_loss\n        return F_loss.mean()\n\ndef get_criterion(cfg):\n    criterion = FocalLoss(alpha=1, gamma=3.5)\n    return criterion\n\ndef find_optimal_threshold_per_class(targets, probs):\n    thresholds = np.arange(0.05, 0.95, 0.05)\n    best_thresholds = []\n    for i in range(targets.shape[1]):\n        best_f1 = 0\n        best_thresh = 0.3\n        for thresh in thresholds:\n            preds = (probs[:, i] > thresh).astype(int)\n            f1 = f1_score(targets[:, i], preds, average='macro', zero_division=0)\n            if f1 > best_f1:\n                best_f1 = f1\n                best_thresh = thresh\n        best_thresholds.append(best_thresh)\n    return best_thresholds\n\n## Training and Validation Loops\ndef train_one_epoch(model, loader, optimizer, criterion, device, scheduler=None):\n    model.train()\n    losses = []\n    all_targets = []\n    all_outputs = []\n    pbar = tqdm(enumerate(loader), total=len(loader), desc=\"Training\")\n    \n    for step, batch in pbar:\n        if not batch:\n            continue\n        inputs = batch['melspec'].to(device)\n        targets = batch['target'].to(device)\n        optimizer.zero_grad()\n        outputs = model(inputs)\n        if isinstance(outputs, tuple):\n            outputs, loss = outputs\n        else:\n            loss = criterion(outputs, targets)\n        loss.backward()\n        optimizer.step()\n        outputs = outputs.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        all_outputs.append(outputs)\n        all_targets.append(targets)\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        })\n    \n    all_outputs = np.concatenate(all_outputs)\n    all_targets = np.concatenate(all_targets)\n    auc = calculate_auc(all_targets, all_outputs)\n    probs = 1 / (1 + np.exp(-all_outputs))\n    thresholds = find_optimal_threshold_per_class(all_targets, probs)\n    y_pred = np.zeros_like(probs)\n    for i, thresh in enumerate(thresholds):\n        y_pred[:, i] = (probs[:, i] > thresh).astype(int)\n    f1 = f1_score(all_targets, y_pred, average='macro')\n    avg_loss = np.mean(losses)\n    \n    return avg_loss, auc, f1, thresholds\n\ndef validate(model, loader, criterion, device):\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            if not batch:\n                continue\n            inputs = batch['melspec'].to(device)\n            targets = batch['target'].to(device)\n            outputs = model(inputs)\n            loss = criterion(outputs, targets)\n            outputs = outputs.cpu().numpy()\n            targets = targets.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    probs = 1 / (1 + np.exp(-all_outputs))\n    thresholds = find_optimal_threshold_per_class(all_targets, probs)\n    y_pred = np.zeros_like(probs)\n    for i, thresh in enumerate(thresholds):\n        y_pred[:, i] = (probs[:, i] > thresh).astype(int)\n    f1 = f1_score(all_targets, y_pred, average='macro')\n    avg_loss = np.mean(losses)\n    \n    return avg_loss, auc, f1, thresholds, all_outputs, all_targets\n\ndef calculate_auc(targets, outputs):\n    num_classes = targets.shape[1]\n    aucs = []\n    probs = 1 / (1 + np.exp(-outputs))\n    for i in range(num_classes):\n        if np.sum(targets[:, i]) > 0:\n            class_auc = roc_auc_score(targets[:, i], probs[:, i])\n            aucs.append(class_auc)\n    return np.mean(aucs) if aucs else 0.0\n\n## Visualization Functions\ndef plot_roc_curves(fold_results, thresholds_per_fold):\n    plt.figure(figsize=(8, 6))\n    mean_fpr = np.linspace(0, 1, 100)\n    tprs = []\n    for i, r in enumerate(fold_results):\n        fpr, tpr, _ = roc_curve(r['y_true'].ravel(), r['y_score'].ravel())\n        plt.plot(fpr, tpr, lw=1, alpha=0.6, label=f'Fold {i} (AUC={auc(fpr, tpr):.3f})')\n        tprs.append(np.interp(mean_fpr, fpr, tpr))\n        tprs[-1][0] = 0.0\n        # Mark threshold point (average threshold across classes for visualization)\n        avg_threshold = np.mean(thresholds_per_fold[i])\n        thresh_idx = np.argmin(np.abs(fpr - avg_threshold))\n        plt.scatter(fpr[thresh_idx], tpr[thresh_idx], marker='o', s=100, label=f'Fold {i} Threshold')\n    mean_tpr = np.mean(tprs, axis=0)\n    std_tpr = np.std(tprs, axis=0)\n    mean_auc = auc(mean_fpr, mean_tpr)\n    plt.plot(mean_fpr, mean_tpr, lw=2, color='navy', label=f'Mean ROC (AUC={mean_auc:.3f})')\n    plt.fill_between(mean_fpr, mean_tpr - std_tpr, mean_tpr + std_tpr, color='grey', alpha=0.2)\n    plt.plot([0, 1], [0, 1], '--', color='gray')\n    plt.xlabel('False Positive Rate', fontsize=12)\n    plt.ylabel('True Positive Rate', fontsize=12)\n    plt.title('ROC Curves — 5-Fold CV', fontsize=14)\n    plt.legend(loc='lower right', fontsize=10)\n    plt.grid(True, linestyle='--', alpha=0.7)\n    plt.tight_layout()\n    plt.savefig(os.path.join(cfg.OUTPUT_DIR, 'roc_curves.png'))\n    plt.show()\n\ndef plot_precision_recall_curves(fold_results, species_ids):\n    num_classes = len(species_ids)\n    cols = 3\n    rows = math.ceil(num_classes / cols)\n    fig, axes = plt.subplots(rows, cols, figsize=(cols * 5, rows * 4), constrained_layout=True)\n    axes = axes.ravel() if num_classes > 1 else [axes]\n\n    for idx, species in enumerate(species_ids):\n        mean_precision = np.linspace(0, 1, 100)\n        recalls = []\n        for r in fold_results:\n            precision, recall, _ = precision_recall_curve(r['y_true'][:, idx], r['y_score'][:, idx])\n            recalls.append(np.interp(mean_precision, precision[::-1], recall[::-1]))\n        mean_recall = np.mean(recalls, axis=0)\n        std_recall = np.std(recalls, axis=0)\n        axes[idx].plot(mean_precision, mean_recall, lw=2, label=f'{species} (Mean)')\n        axes[idx].fill_between(mean_precision, mean_recall - std_recall, mean_recall + std_recall, alpha=0.2)\n        axes[idx].set_title(f'Precision-Recall: {species}', fontsize=10)\n        axes[idx].set_xlabel('Precision', fontsize=8)\n        axes[idx].set_ylabel('Recall', fontsize=8)\n        axes[idx].legend(loc='lower left', fontsize=8)\n        axes[idx].grid(True, linestyle='--', alpha=0.7)\n\n    for idx in range(len(species_ids), len(axes)):\n        axes[idx].set_visible(False)\n    \n    plt.suptitle('Precision-Recall Curves per Class', fontsize=14)\n    plt.savefig(os.path.join(cfg.OUTPUT_DIR, 'precision_recall_curves.png'))\n    plt.show()\n\ndef plot_confusion_and_classic_metrics(fold_results):\n    sens_list, spec_list, prec_list, f1_list, acc_list = [], [], [], [], []\n    for i, r in enumerate(fold_results):\n        cm = confusion_matrix(r['y_true'].ravel(), r['y_pred'].ravel())\n        tn, fp, fn, tp = cm.ravel()\n        sens = tp / (tp + fn) if tp + fn > 0 else 0\n        spec = tn / (tn + fp) if tn + fp > 0 else 0\n        prec, recall, f1, _ = precision_recall_fscore_support(\n            r['y_true'].ravel(), r['y_pred'].ravel(), average='binary', zero_division=0)\n        acc = (tp + tn) / (tp + tn + fp + fn) if (tp + tn + fp + fn) > 0 else 0\n        print(f\"\\n––– Fold {i} Metrics –––\")\n        print(f\"Confusion Matrix:\\n{cm}\")\n        print(f\"Accuracy: {acc:.3f}, Sensitivity: {sens:.3f}, Specificity: {spec:.3f}\")\n        print(f\"Precision: {prec:.3f}, Recall: {recall:.3f}, F1: {f1:.3f}\")\n        sens_list.append(sens)\n        spec_list.append(spec)\n        prec_list.append(prec)\n        f1_list.append(f1)\n        acc_list.append(acc)\n    \n    print(\"\\n––– Average Over Folds –––\")\n    print(f\"Accuracy: {np.mean(acc_list):.3f} ± {np.std(acc_list):.3f}\")\n    print(f\"Sensitivity: {np.mean(sens_list):.3f} ± {np.std(sens_list):.3f}\")\n    print(f\"Specificity: {np.mean(spec_list):.3f} ± {np.std(spec_list):.3f}\")\n    print(f\"Precision: {np.mean(prec_list):.3f} ± {np.std(prec_list):.3f}\")\n    print(f\"F1 Score: {np.mean(f1_list):.3f} ± {np.std(f1_list):.3f}\")\n\ndef compute_per_class_metrics(fold_results, species_ids):\n    per_class_metrics = []\n    for idx, species in enumerate(species_ids):\n        precisions, recalls = [], []\n        for r in fold_results:\n            prec, rec, _, _ = precision_recall_fscore_support(\n                r['y_true'][:, idx], r['y_pred'][:, idx], average='binary', zero_division=0)\n            precisions.append(prec)\n            recalls.append(rec)\n        per_class_metrics.append({\n            'species': species,\n            'precision': np.mean(precisions),\n            'recall': np.mean(recalls)\n        })\n        print(f\"Species: {species}, Precision: {np.mean(precisions):.4f}, Recall: {np.mean(recalls):.4f}\")\n    return pd.DataFrame(per_class_metrics)\n\ndef compute_micro_macro_metrics(fold_results, species_ids):\n    micro_precisions, micro_recalls = [], []\n    macro_precisions, macro_recalls = [], []\n    \n    for r in fold_results:\n        prec, rec, _, _ = precision_recall_fscore_support(\n            r['y_true'].ravel(), r['y_pred'].ravel(), average='micro', zero_division=0)\n        micro_precisions.append(prec)\n        micro_recalls.append(rec)\n        prec, rec, _, _ = precision_recall_fscore_support(\n            r['y_true'].ravel(), r['y_pred'].ravel(), average='macro', zero_division=0)\n        macro_precisions.append(prec)\n        macro_recalls.append(rec)\n    \n    print(\"\\n––– Micro and Macro Averages –––\")\n    print(f\"Micro Precision: {np.mean(micro_precisions):.4f} ± {np.std(micro_precisions):.4f}\")\n    print(f\"Micro Recall: {np.mean(micro_recalls):.4f} ± {np.std(micro_recalls):.4f}\")\n    print(f\"Macro Precision: {np.mean(macro_precisions):.4f} ± {np.std(macro_precisions):.4f}\")\n    print(f\"Macro Recall: {np.mean(macro_recalls):.4f} ± {np.std(macro_recalls):.4f}\")\n\n## Training Loop\ndef run_training(df, cfg):\n    print(f\"→ 5-Fold Stratified CV: n_splits={cfg.n_fold}, shuffle=True, seed={cfg.seed}\")\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    spectrograms = None\n    if cfg.LOAD_DATA:\n        print(\"Loading pre-computed spectrograms...\")\n        try:\n            spectrograms = np.load(cfg.spectrogram_npy, allow_pickle=True).item()\n            print(f\"Loaded {len(spectrograms)} spectrograms\")\n        except Exception as e:\n            print(f\"Error loading spectrograms: {e}\")\n            cfg.LOAD_DATA = False\n    \n    if not cfg.LOAD_DATA:\n        if 'filepath' not in df.columns:\n            df['filepath'] = cfg.train_datadir + '/' + df.filename\n        if 'samplename' not in df.columns:\n            df['samplename'] = df.filename.map(lambda x: x.split('/')[0] + '-' + x.split('/')[-1].split('.')[0])\n        spectrograms = generate_spectrograms(df, cfg)\n    \n    best_scores = []\n    all_fold_results = []\n    thresholds_per_fold = []\n    \n    for fold, (train_idx, val_idx) in enumerate(StratifiedKFold(n_splits=cfg.n_fold, shuffle=True, random_state=cfg.seed).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        fold_train_df = df.iloc[train_idx].reset_index(drop=True)\n        fold_val_df = df.iloc[val_idx].reset_index(drop=True)\n        \n        print(f\"Training set: {len(fold_train_df)} samples\")\n        print(f\"Validation set: {len(fold_val_df)} samples\")\n        \n        train_dataset = BirdCLEFDatasetFromNPY(fold_train_df, cfg, spectrograms, mode='train')\n        val_dataset = BirdCLEFDatasetFromNPY(fold_val_df, cfg, spectrograms, mode='valid')\n        \n        class_weights = fold_train_df['primary_label'].value_counts().map(lambda x: 1/(x+1e-6)**0.8).to_dict()\n        sampler_weights = [class_weights[row['primary_label']] for _, row in fold_train_df.iterrows()]\n        if len(sampler_weights) != len(train_dataset):\n            print(f\"Sampler weights size {len(sampler_weights)} does not match dataset size {len(train_dataset)}\")\n            raise ValueError(\"Sampler weights mismatch\")\n        sampler = WeightedRandomSampler(sampler_weights, len(sampler_weights), replacement=True)\n        \n        train_loader = DataLoader(\n            train_dataset, batch_size=cfg.batch_size, sampler=sampler, shuffle=False,\n            num_workers=cfg.num_workers, pin_memory=True, collate_fn=collate_fn, drop_last=True)\n        val_loader = DataLoader(\n            val_dataset, batch_size=cfg.batch_size, shuffle=False,\n            num_workers=cfg.num_workers, pin_memory=True, collate_fn=collate_fn)\n        \n        model = BirdCLEFModel(cfg).to(cfg.device)\n        optimizer = get_optimizer(model, cfg)\n        criterion = get_criterion(cfg)\n        scheduler = get_scheduler(optimizer, cfg)\n        \n        best_f1 = 0\n        best_thresholds = [0.3] * cfg.num_classes\n        best_epoch = 0\n        patience = 10\n        counter = 0\n        \n        for epoch in range(cfg.epochs):\n            print(f\"\\nEpoch {epoch+1}/{cfg.epochs}\")\n            \n            train_loss, train_auc, train_f1, train_thresholds = train_one_epoch(\n                model, train_loader, optimizer, criterion, cfg.device)\n            \n            val_loss, val_auc, val_f1, val_thresholds, all_outputs, all_targets = validate(\n                model, val_loader, criterion, cfg.device)\n            \n            print(f\"Train Loss: {train_loss:.4f}, Train AUC: {train_auc:.4f}, Train F1: {train_f1:.4f}\")\n            print(f\"Val Loss: {val_loss:.4f}, Val AUC: {val_auc:.4f}, Val F1: {val_f1:.4f}\")\n            \n            all_fold_results.append({\n                'y_true': all_targets.astype(int),\n                'y_score': 1 / (1 + np.exp(-all_outputs)),\n                'y_pred': np.zeros_like(all_outputs),\n                'fold': fold\n            })\n            for i, thresh in enumerate(val_thresholds):\n                all_fold_results[-1]['y_pred'][:, i] = (all_fold_results[-1]['y_score'][:, i] > thresh).astype(int)\n            \n            if val_f1 > best_f1:\n                best_f1 = val_f1\n                best_thresholds = val_thresholds\n                best_epoch = epoch + 1\n                print(f\"New best F1: {best_f1:.4f} at epoch {best_epoch}\")\n                torch.save({\n                    'model_state_dict': model.state_dict(),\n                    'epoch': epoch,\n                    'val_f1': val_f1,\n                    'val_thresholds': val_thresholds\n                }, f\"model_fold{fold}.pth\")\n                counter = 0\n            else:\n                counter += 1\n            \n            if scheduler is not None:\n                scheduler.step()\n            \n            if counter >= patience:\n                print(f\"Early stopping at epoch {epoch+1}\")\n                break\n        \n        best_scores.append(best_f1)\n        thresholds_per_fold.append(best_thresholds)\n        print(f\"\\nBest F1 for fold {fold}: {best_f1:.4f} at epoch {best_epoch}\")\n        \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 {fold}: {score:.4f}\")\n    print(f\"Mean F1: {np.mean(best_scores):.4f}\")\n    print(\"=\"*60)\n    \n    # Plot ROC curves with thresholds\n    plot_roc_curves(all_fold_results, thresholds_per_fold)\n    \n    # Plot precision-recall curves\n    plot_precision_recall_curves(all_fold_results, species_ids)\n    \n    # Confusion matrix and classic metrics\n    plot_confusion_and_classic_metrics(all_fold_results)\n    \n    # Per-class metrics\n    print(\"\\nPer-Class Metrics on Validation Set:\")\n    per_class_df = compute_per_class_metrics(all_fold_results, species_ids)\n    per_class_df.to_csv(os.path.join(cfg.OUTPUT_DIR, 'per_class_metrics.csv'), index=False)\n    \n    # Micro and macro metrics\n    compute_micro_macro_metrics(all_fold_results, species_ids)\n\nif __name__ == \"__main__\":\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    print(f\"LOAD_DATA is set to {cfg.LOAD_DATA}\")\n    if cfg.LOAD_DATA:\n        print(\"Using pre-computed mel spectrograms from NPY file\")\n    else:\n        print(\"Will generate spectrograms on-the-fly\")\n    \n    run_training(train_df, cfg)\n    \n    print(\"\\nTraining complete!\")","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}