{"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":"none","dataSources":[{"sourceId":91844,"databundleVersionId":11361821,"sourceType":"competition"},{"sourceId":11702705,"sourceType":"datasetVersion","datasetId":7339703}],"dockerImageVersionId":31011,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport pandas as pd\nimport librosa\nimport numpy as np\nimport sys\nimport torch\nimport cv2\nimport math\nimport time\nimport logging\nimport random\nimport gc\nfrom sklearn.model_selection import train_test_split\nfrom tqdm.notebook import tqdm\nimport matplotlib.pyplot as plt\n\nimport torch\nimport torch.nn as nn\nimport torchaudio\nimport torchaudio.transforms as AT\nfrom torch.utils.data import DataLoader, Dataset\nfrom torchvision import models\nimport torchvision\n\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.optim import lr_scheduler\n\nimport seaborn as sns\nfrom sklearn.metrics import roc_auc_score\n\nimport timm\nfrom sklearn.model_selection import StratifiedKFold\n\n\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-05-08T06:37:56.109413Z","iopub.execute_input":"2025-05-08T06:37:56.109692Z","iopub.status.idle":"2025-05-08T06:38:11.080835Z","shell.execute_reply.started":"2025-05-08T06:37:56.109661Z","shell.execute_reply":"2025-05-08T06:38:11.080091Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#!pip install timm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-08T06:38:11.081570Z","iopub.execute_input":"2025-05-08T06:38:11.081794Z","iopub.status.idle":"2025-05-08T06:38:11.085702Z","shell.execute_reply.started":"2025-05-08T06:38:11.081777Z","shell.execute_reply":"2025-05-08T06:38:11.085040Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DEBUG_MODE = False\n\nOUTPUT_DIR = '/kaggle/working/'\nDATA_ROOT = '/kaggle/input/birdclef-2025'\nTRAIN_DIR = '/kaggle/input/birdclef-2025/train_audio'\n\nTAXONOMY_CSV = '/kaggle/input/birdclef-2025/taxonomy.csv'\nTRAIN_CSV = '/kaggle/input/birdclef-2025/train.csv'\n\nFS = 32000     # tần số để cắt file ogg\n    \n# Mel spectrogram parameters\nN_FFT = 1024\nHOP_LENGTH = 512\nN_MELS = 128\nFMIN = 50\nFMAX = 14000\n    \nTARGET_DURATION = 5.0\nTARGET_SHAPE = (256, 256)  \n\nAUG_PROB = 0.5               # xác suất sẽ augment data cho tập train\n    \nN_MAX = 50 if DEBUG_MODE else None  \n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-08T06:38:11.087842Z","iopub.execute_input":"2025-05-08T06:38:11.088626Z","iopub.status.idle":"2025-05-08T06:38:11.111352Z","shell.execute_reply.started":"2025-05-08T06:38:11.088605Z","shell.execute_reply":"2025-05-08T06:38:11.110550Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"Loading taxonomy data...\")\ntaxonomy_df = pd.read_csv(f'{DATA_ROOT}/taxonomy.csv')\nspecies_class_map = dict(zip(taxonomy_df['primary_label'], taxonomy_df['class_name']))    # dict map tên loài sang class name\n\n# load dữ liệu training\nprint(\"Loading training metadata...\")\ntrain_df = pd.read_csv(f'{DATA_ROOT}/train.csv')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-08T06:38:11.112220Z","iopub.execute_input":"2025-05-08T06:38:11.112526Z","iopub.status.idle":"2025-05-08T06:38:11.341322Z","shell.execute_reply.started":"2025-05-08T06:38:11.112501Z","shell.execute_reply":"2025-05-08T06:38:11.340592Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class BirdCLEFDataset(Dataset):\n    # def __init__(self, spectrograms=None, df, taxonomy_csv, target_shape=(128, 313), mode=\"train\"):\n    def __init__(self, df, spectrograms=None, mode=\"train\"):\n        self.df = df\n        self.mode = mode\n        self.spectrograms = spectrograms\n\n        # Load taxonomy để tạo label mapping từ id loài sang one hot code\n        taxonomy_df = pd.read_csv(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        # thêm cột đường dẫn đến ogg nếu df chưa có\n        if 'filepath' not in self.df.columns:\n            self.df['filepath'] = TRAIN_DIR + '/' + self.df['filename']\n\n        # thêm cột sample name để lấy kết quả từ npz\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                    \n        sample_names = set(self.df['samplename'])\n        # in ra số lượng spectrogrames lấy được trong npy với các sample name của từng tập train, val\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\n    \n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        samplename = row['samplename']\n        spec = self.spectrograms[samplename]\n\n        # nếu ko có spectrogram thì tạo ảnh trắng\n        if spec is None:\n            spec = np.zeros(TARGET_SHAPE, dtype=np.float32)\n            if self.mode == \"train\":  \n                print(f\"Warning: Spectrogram for {samplename} not found and could not be generated\")\n\n        \n        # Add channel dimension: [1, H, W]\n        spec = torch.tensor(spec, dtype=torch.float32).unsqueeze(0)\n\n        # có thể thêm data augmentation cho tập train, aug_prob = 0.5\n        if self.mode == \"train\" and random.random() < AUG_PROB:\n            spec = self.apply_spec_augmentations(spec)\n\n        # One-hot encode label\n        target = np.zeros(self.num_classes, dtype=np.float32)\n        label = row['primary_label']\n        if label in self.label_to_idx:\n            target[self.label_to_idx[label]] = 1.0\n\n        # nếu có label thứ 2\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            'melspec': spec,\n            'target': torch.tensor(target, dtype=torch.float32),\n            'filename': row['filename']\n        }\n\n    \n    def apply_spec_augmentations(self, spec):\n        \"\"\"Apply augmentations to spectrogram\"\"\"\n    \n        # Time masking (horizontal stripes)\n        if random.random() < 0.5:\n            num_masks = random.randint(1, 3)\n            for _ in range(num_masks):\n                width = random.randint(5, 20)\n                start = random.randint(0, spec.shape[2] - width)\n                spec[0, :, start:start+width] = 0\n        \n        # Frequency masking (vertical stripes)\n        if random.random() < 0.5:\n            num_masks = random.randint(1, 3)\n            for _ in range(num_masks):\n                height = random.randint(5, 20)\n                start = random.randint(0, spec.shape[1] - height)\n                spec[0, start:start+height, :] = 0\n        \n        # Random brightness/contrast\n        if random.random() < 0.5:\n            gain = random.uniform(0.8, 1.2)\n            bias = random.uniform(-0.1, 0.1)\n            spec = spec * gain + bias\n            spec = torch.clamp(spec, 0, 1) \n            \n        return spec","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-08T06:38:11.342184Z","iopub.execute_input":"2025-05-08T06:38:11.342449Z","iopub.status.idle":"2025-05-08T06:38:11.357549Z","shell.execute_reply.started":"2025-05-08T06:38:11.342430Z","shell.execute_reply":"2025-05-08T06:38:11.356762Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"spectrograms = np.load(\"/kaggle/input/birdclef-melspec-data/bird_mel_spectrograms.npy\", allow_pickle=True).item()\n\ntrain_meta = pd.read_csv('/kaggle/input/birdclef-2025/train.csv')\ntrain_df, val_df = train_test_split(train_meta, test_size=0.2, random_state=42)\n\ntrain_dataset = BirdCLEFDataset(train_df, spectrograms, mode='train')\ntrain_loader = DataLoader(train_dataset, batch_size=24, shuffle=True, num_workers=2,drop_last=True)\n\nval_dataset = BirdCLEFDataset(val_df, spectrograms, mode='val')\nval_loader = DataLoader(val_dataset, batch_size=24, shuffle=False, num_workers=1,drop_last=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-08T06:38:11.358302Z","iopub.execute_input":"2025-05-08T06:38:11.358640Z","iopub.status.idle":"2025-05-08T06:39:01.539086Z","shell.execute_reply.started":"2025-05-08T06:38:11.358620Z","shell.execute_reply":"2025-05-08T06:39:01.538205Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CFG:\n    epochs = 8\n    batch_size = 32  \n    criterion = 'AsymmetricLossMultiLabel'\n\n    n_fold = 5\n    selected_folds = [0, 1, 2, 3, 4]   \n\n    optimizer = 'AdamW'\n    lr = 5e-4 \n    weight_decay = 1e-5\n  \n    scheduler = 'CosineAnnealingLR'\n    min_lr = 1e-6\n    T_max = epochs\n\n    LOAD_DATA = True  \n    \n    def update_debug_settings(self):\n        if self.debug:\n            self.epochs = 2\n            self.selected_folds = [0]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-08T06:39:01.539930Z","iopub.execute_input":"2025-05-08T06:39:01.540159Z","iopub.status.idle":"2025-05-08T06:39:01.545205Z","shell.execute_reply.started":"2025-05-08T06:39:01.540140Z","shell.execute_reply":"2025-05-08T06:39:01.544476Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class BirdCLEFModel(nn.Module):\n    def __init__(self, num_classes, pretrained=True):\n        super().__init__()\n        #num_classes = len(taxonomy_df)\n\n        self.patch_size = 16\n        self.embed_dim = 768\n        self.depth = 12\n        self.num_heads = 12\n        self.spec_height = 256\n        self.spec_width = 256\n        \n        \n        # Create Audio Spectrogram Transformer backbone\n        # For this we use ViT but adapt it for audio spectrogram input\n        self.ast = timm.create_model(\n            'vit_base_patch16_224',  # Base ViT model\n            pretrained=pretrained,\n            img_size=(256, 256),  # Spectrogram dimensions\n            in_chans=1,  # Usually 1 for mel spectrograms\n            patch_size=self.patch_size,\n            embed_dim=self.embed_dim,\n            depth=self.depth,\n            num_heads=self.num_heads,\n            drop_path_rate=0.2,\n            drop_rate=0.3\n        )\n        \n        # Get the output feature dimension\n        backbone_out = self.ast.head.in_features\n        self.ast.head = nn.Identity()  # Remove the classification head\n        \n        # Secondary feature extractor - can be a CNN for local features\n        self.backbone2 = timm.create_model(\n            'efficientnet_b0',\n            pretrained=pretrained,\n            in_chans=1,\n            drop_rate=0.3,\n            drop_path_rate=0.2\n        )\n        \n        # Get output features for backbone 2\n        if 'efficientnet' in 'regnety_008':\n            backbone2_out = self.backbone2.classifier.in_features\n            self.backbone2.classifier = nn.Identity()\n        elif 'resnet' in 'regnety_008':\n            backbone2_out = self.backbone2.fc.in_features\n            self.backbone2.fc = nn.Identity()\n        elif 'convnext' in 'regnety_008':\n            backbone2_out = self.backbone2.head.fc.in_features\n            self.backbone2.head.fc = nn.Identity()\n        else:\n            backbone2_out = self.backbone2.get_classifier().in_features\n            self.backbone2.reset_classifier(0, '')\n        \n        # Global pooling for CNN backbone\n        self.pooling2 = nn.AdaptiveAvgPool2d(1)\n        \n        # Feature dimensions\n        self.feat_dim1 = backbone_out\n        self.feat_dim2 = backbone2_out\n        \n        # Feature fusion layers (to combine transformer and CNN outputs)\n        self.fusion = nn.Sequential(\n            nn.Linear(backbone_out + backbone2_out, backbone_out),\n            nn.BatchNorm1d(backbone_out),\n            nn.SiLU(inplace=True),\n            nn.Dropout(0.3)\n        )\n        \n        # Classifier head\n        self.classifier = nn.Linear(backbone_out, num_classes)\n        \n        # Mixup and other augmentations\n        self.mixup_alpha = 0.5\n            \n    def forward(self, x, targets=None):\n        # Apply mixup if enabled and in training mode\n        if self.training 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        \n        # Extract features from AST\n        features1 = self.ast(x)\n        \n        # Extract features from CNN backbone\n        features2 = self.backbone2(x)\n        \n        # Handle feature maps if necessary for backbone 2\n        if len(features2.shape) == 4:\n            features2 = self.pooling2(features2)\n            features2 = features2.view(features2.size(0), -1)\n        \n        # Concatenate features from both backbones\n        combined_features = torch.cat([features1, features2], dim=1)\n        \n        # Fuse the features\n        fused_features = self.fusion(combined_features)\n        \n        # Get logits from classifier\n        logits = self.classifier(fused_features)\n            \n        return logits\n    \n    def mixup_data(self, x, targets):\n        \"\"\"Applies mixup to the data batch\"\"\"\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        \"\"\"Applies mixup to the loss function\"\"\"\n        return lam * criterion(pred, y_a) + (1 - lam) * criterion(pred, y_b)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-08T06:39:01.546281Z","iopub.execute_input":"2025-05-08T06:39:01.546548Z","iopub.status.idle":"2025-05-08T06:39:01.567274Z","shell.execute_reply.started":"2025-05-08T06:39:01.546529Z","shell.execute_reply":"2025-05-08T06:39:01.566457Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class AsymmetricLossMultiLabel(nn.Module):\n    def __init__(\n        self,\n        gamma_neg=4,\n        gamma_pos=1,\n        clip=0.05,\n        eps=1e-8,\n        disable_torch_grad_focal_loss=False,\n        reduction=\"mean\",\n    ):\n        super().__init__()\n\n        self.gamma_neg = gamma_neg\n        self.gamma_pos = gamma_pos\n        self.clip = clip\n        self.disable_torch_grad_focal_loss = disable_torch_grad_focal_loss\n        self.eps = eps\n        self.reduction = reduction\n\n    def forward(self, x, y):\n        \"\"\" \"\n        Parameters\n        ----------\n        x: input logits\n        y: targets (multi-label binarized vector)\n        \"\"\"\n\n        # Calculating Probabilities\n        x_sigmoid = torch.sigmoid(x)\n        xs_pos = x_sigmoid\n        xs_neg = 1 - x_sigmoid\n\n        # Asymmetric Clipping\n        if self.clip is not None and self.clip > 0:\n            xs_neg = (xs_neg + self.clip).clamp(max=1)\n\n        # Basic CE calculation\n        los_pos = y * torch.log(xs_pos.clamp(min=self.eps))\n        los_neg = (1 - y) * torch.log(xs_neg.clamp(min=self.eps))\n        loss = los_pos + los_neg\n\n        # Asymmetric Focusing\n        if self.gamma_neg > 0 or self.gamma_pos > 0:\n            if self.disable_torch_grad_focal_loss:\n                torch._C.set_grad_enabled(False)\n            pt0 = xs_pos * y\n            pt1 = xs_neg * (1 - y)  # pt = p if t > 0 else 1-p\n            pt = pt0 + pt1\n            one_sided_gamma = self.gamma_pos * y + self.gamma_neg * (1 - y)\n            one_sided_w = torch.pow(1 - pt, one_sided_gamma)\n            if self.disable_torch_grad_focal_loss:\n                torch._C.set_grad_enabled(True)\n            loss *= one_sided_w\n\n        if self.reduction == \"mean\":\n            return -loss.mean()\n        if self.reduction == \"sum\":\n            return -loss.sum()\n\n        return -loss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-08T06:39:01.569471Z","iopub.execute_input":"2025-05-08T06:39:01.569749Z","iopub.status.idle":"2025-05-08T06:39:01.592488Z","shell.execute_reply.started":"2025-05-08T06:39:01.569730Z","shell.execute_reply":"2025-05-08T06:39:01.591603Z"}},"outputs":[],"execution_count":null},{"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    if cfg.criterion == 'AsymmetricLossMultiLabel':\n        criterion = AsymmetricLossMultiLabel(\n            gamma_neg=4,\n            gamma_pos=1,\n            clip=0.05,\n            eps=1e-8,\n            disable_torch_grad_focal_loss=False,\n            reduction=\"mean\"\n        )\n    elif cfg.criterion == 'BCEWithLogitsLoss':\n        criterion = nn.BCEWithLogitsLoss()\n    else:\n        raise NotImplementedError(f\"Criterion {cfg.criterion} not implemented\")\n        \n    return criterion","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-08T06:39:01.594334Z","iopub.execute_input":"2025-05-08T06:39:01.594664Z","iopub.status.idle":"2025-05-08T06:39:01.610193Z","shell.execute_reply.started":"2025-05-08T06:39:01.594635Z","shell.execute_reply":"2025-05-08T06:39:01.609237Z"}},"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        if isinstance(batch['melspec'], list):\n            batch_outputs = []\n            batch_losses = []\n            \n            for i in range(len(batch['melspec'])):\n                inputs = batch['melspec'][i].unsqueeze(0).to(device)\n                target = batch['target'][i].unsqueeze(0).to(device)\n                \n                optimizer.zero_grad()\n                output = model(inputs)\n                loss = criterion(output, target)\n                loss.backward()\n                \n                batch_outputs.append(output.detach().cpu())\n                batch_losses.append(loss.item())\n            \n            optimizer.step()\n            outputs = torch.cat(batch_outputs, dim=0).numpy()\n            loss = np.mean(batch_losses)\n            targets = batch['target'].numpy()\n            \n        else:\n            inputs = batch['melspec'].to(device)\n            targets = batch['target'].to(device)\n            \n            optimizer.zero_grad()\n            outputs = model(inputs)\n            \n            if isinstance(outputs, tuple):\n                outputs, loss = outputs  \n            else:\n                loss = criterion(outputs, targets)\n                \n            loss.backward()\n            optimizer.step()\n            \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 if isinstance(loss, float) else 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    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            if isinstance(batch['melspec'], list):\n                batch_outputs = []\n                batch_losses = []\n                \n                for i in range(len(batch['melspec'])):\n                    inputs = batch['melspec'][i].unsqueeze(0).to(device)\n                    target = batch['target'][i].unsqueeze(0).to(device)\n                    \n                    output = model(inputs)\n                    loss = criterion(output, target)\n                    \n                    batch_outputs.append(output.detach().cpu())\n                    batch_losses.append(loss.item())\n                \n                outputs = torch.cat(batch_outputs, dim=0).numpy()\n                loss = np.mean(batch_losses)\n                targets = batch['target'].numpy()\n                \n            else:\n                inputs = batch['melspec'].to(device)\n                targets = batch['target'].to(device)\n                \n                outputs = model(inputs)\n                loss = criterion(outputs, targets)\n                \n                outputs = outputs.detach().cpu().numpy()\n                targets = targets.detach().cpu().numpy()\n            \n            all_outputs.append(outputs)\n            all_targets.append(targets)\n            losses.append(loss if isinstance(loss, float) else 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    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-05-08T06:39:01.611160Z","iopub.execute_input":"2025-05-08T06:39:01.611533Z","iopub.status.idle":"2025-05-08T06:39:01.647311Z","shell.execute_reply.started":"2025-05-08T06:39:01.611509Z","shell.execute_reply":"2025-05-08T06:39:01.646271Z"}},"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 key == 'target' and isinstance(result[key][0], torch.Tensor):\n            result[key] = torch.stack(result[key])\n        elif key == 'melspec' and isinstance(result[key][0], torch.Tensor):\n            shapes = [t.shape for t in result[key]]\n            if len(set(str(s) for s in shapes)) == 1:\n                result[key] = torch.stack(result[key])\n    \n    return result","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-08T06:39:01.648476Z","iopub.execute_input":"2025-05-08T06:39:01.649460Z","iopub.status.idle":"2025-05-08T06:39:01.669149Z","shell.execute_reply.started":"2025-05-08T06:39:01.649431Z","shell.execute_reply":"2025-05-08T06:39:01.668186Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"num_classes = len(taxonomy_df['primary_label'])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-08T06:39:01.669931Z","iopub.execute_input":"2025-05-08T06:39:01.670252Z","iopub.status.idle":"2025-05-08T06:39:01.688287Z","shell.execute_reply.started":"2025-05-08T06:39:01.670226Z","shell.execute_reply":"2025-05-08T06:39:01.687528Z"}},"outputs":[],"execution_count":null},{"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    species_ids = taxonomy_df['primary_label'].tolist()\n    num_classes = len(species_ids)\n    \n    if DEBUG_MODE:\n        cfg.update_debug_settings()\n        \n        model = BirdCLEFModel(num_classes=num_classes, pretrained=True).to(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                device,\n                scheduler if isinstance(scheduler, lr_scheduler.OneCycleLR) else None\n            )\n            \n            val_loss, val_auc = validate(model, val_loader, criterion, 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        print(f\"{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(\"=\"*60)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-08T06:39:01.689205Z","iopub.execute_input":"2025-05-08T06:39:01.689536Z","iopub.status.idle":"2025-05-08T06:39:01.705706Z","shell.execute_reply.started":"2025-05-08T06:39:01.689516Z","shell.execute_reply":"2025-05-08T06:39:01.704974Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import time\ncfg = CFG()\nrun_training(train_df, cfg)\n    \nprint(\"\\nTraining complete!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-08T06:39:21.315898Z","iopub.execute_input":"2025-05-08T06:39:21.316669Z","iopub.status.idle":"2025-05-08T06:39:21.321409Z","shell.execute_reply.started":"2025-05-08T06:39:21.316643Z","shell.execute_reply":"2025-05-08T06:39:21.320413Z"}},"outputs":[],"execution_count":null}]}