{"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":[{"sourceType":"competition","sourceId":91844,"databundleVersionId":11361821},{"sourceType":"kernelVersion","sourceId":240162995}],"dockerImageVersionId":31040,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"**EfficientNet B0 Pytorch [Train]**\n\nThis notebook is designed to train a bird song classification model for the BirdCLEF 2025 competition.\n\nBelow is an explanation of the main sections and their content in the notebook:\n\n1. Import Libraries:\n    * Imports necessary libraries such as PyTorch, timm, librosa, etc., which are used for image processing, audio processing, building machine learning models, and data visualization.\n2. Configuration:\n    * The CFG class defines various settings for model training, including seed value, number of epochs, batch size, model name, learning rate, optimizer, scheduler, and more.\n    * It also checks if a GPU is available.\n    * There are settings for debug mode, which reduces the number of epochs and other parameters when enabled.\n3. Dataset Preparation and Data Augmentations:\n    * The `BirdCLEFDatasetFromNPY` class defines a custom dataset.\n    * This dataset loads pre-computed mel spectrograms from a .npy file or generates them from audio files if necessary.\n    * During training, it applies data augmentations to the spectrograms, such as time masking, frequency masking, and random brightness/contrast adjustments.\n    * It also handles encoding labels into a one-hot vector format.\n    * `collate_fn` is a custom function to handle batches with potentially different sized spectrograms (although the notebook seems to use fixed-size spectrograms).\n4. Model Definition:\n    * The `BirdCLEFModel` class defines the classification model.\n    * It uses a pre-trained model from the efficientnet_b0, as the backbone.\n    * An adaptive average pooling layer and a linear classifier are added to the backbone's output.\n    * Mixup, a data augmentation technique, is implemented as part of the model's forward pass during training.\n5. Training Utilities:\n    * The get_optimizer, get_scheduler, and get_criterion functions initialize the optimizer, learning rate scheduler, and loss function based on the defined configuration.\n6. Training Functions:\n    * train_one_epoch executes a single training epoch loop. It calculates the loss for each batch, performs backpropagation, and updates the optimizer.\n    * Mixup is also applied within this function.\n    * `validate` evaluates the model's performance. It calculates the loss and AUC on the validation set.\n    * `calculate_auc` calculates the ROC AUC score.\n7. Load Data:\n    * Loads train.csv and taxonomy.csv as Pandas DataFrames.\n    * Retrieves the list of bird species from taxonomy.csv and sets the number of classes (cfg.num_classes).\n8. Run Training:\n    * Starts the training process according to the configuration (cfg).\n    * Loads the pre-computed mel spectrograms.\n    * Sets up Stratified K-Fold cross-validation. This is a technique to split the data into multiple folds and train and evaluate the model using each fold as the validation set.\n    * For each selected fold, it performs the following:\n    * Splits the data into training and validation sets.\n    * Creates BirdCLEFDatasetFromNPY for each set and sets up DataLoaders.\n    * Initializes the model, optimizer, criterion (loss function), and scheduler.\n    * Iterates through the defined number of epochs for training and validation.\n    * Saves the model weights if the validation AUC improves.\n    * Records the best AUC for each fold.\n9. Cross Validation:\n    * Prints the best AUC for each fold and the average AUC across all folds.\n    * Overall, this notebook constructs a pipeline to efficiently train and evaluate a deep learning model for bird song classification using cross-validation.\n    * It leverages pre-computed spectrograms and incorporates techniques like data augmentation and Mixup.","metadata":{}},{"cell_type":"markdown","source":"# Import Libraries","metadata":{}},{"cell_type":"code","source":"import time\nimport os\nimport random\nimport gc\nimport time\nimport cv2\nimport math\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 matplotlib.pyplot as plt\nimport seaborn as sns\nfrom tqdm.auto import tqdm\n\nimport timm","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-05-19T19:40:01.152332Z","iopub.execute_input":"2025-05-19T19:40:01.152913Z","iopub.status.idle":"2025-05-19T19:40:01.161439Z","shell.execute_reply.started":"2025-05-19T19:40:01.152883Z","shell.execute_reply":"2025-05-19T19:40:01.159983Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Configuration","metadata":{}},{"cell_type":"code","source":"# Check if gpu is available\ntorch.cuda.is_available()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-19T19:40:01.166239Z","iopub.execute_input":"2025-05-19T19:40:01.166651Z","iopub.status.idle":"2025-05-19T19:40:01.194123Z","shell.execute_reply.started":"2025-05-19T19:40:01.166611Z","shell.execute_reply":"2025-05-19T19:40:01.192890Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CFG:\n    seed: int = 42\n    debug: bool = False\n    apex: bool = False\n    print_freq: int = 100\n    num_workers: int = 2\n\n    OUTPUT_DIR: str = \"kaggle/working/\"\n\n    train_data_dir: str = \"/kaggle/input/birdclef-2025/train_audio\"\n    train_csv: str = \"/kaggle/input/birdclef-2025/train.csv\"\n    test_soundscapes: str = \"/kaggle/input/birdclef-2025/test_soundscapes\"\n    submission_csv: str = \"/kaggle/input/birdclef-2025/sample_submission.csv\"\n    taxonomy_csv: str = \"/kaggle/input/birdclef-2025/taxonomy.csv\"\n    spectrogram_npy: str = \"/kaggle/input/transforming-audio-to-mel-spec/birdclef2025_melspec_5sec_256_256.npy\"\n\n    model_name: str = \"efficientnet_b0\"\n    pretrained: bool = True\n    in_channels: int = 1\n\n    TARGET_SHAPE: tuple[int, int] = (256, 256)\n\n    device: str = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    epochs: int = 10\n    batch_size: int = 32\n    criterion: str = \"BCEWithLogitsLoss\"\n\n    n_fold: int = 5\n    selected_folds: list[int] = [0, 1, 2, 3, 4]\n\n    optimizer: str = \"AdamW\"\n    lr: float = 5e-4\n    weight_decay: float = 1e-5\n\n    scheduler: str = \"CosineAnnealingLR\"\n    # scheduler: str = \"ReduceLROnPlateau\"\n    # scheduler: str = \"StepLR\"\n    # scheduler: str = \"OneCycleLR\"\n    min_lr: float = 1e-6\n    T_max: int = epochs\n\n    aug_prob: float = 0.5\n    mixup_alpha: float = 0.5\n\n    def update_debug_settings(self) -> None:\n        if self.debug:\n            self.epochs = 2\n            self.selected_folds = [0]\n\ncfg = CFG()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-19T19:40:01.914760Z","iopub.execute_input":"2025-05-19T19:40:01.915241Z","iopub.status.idle":"2025-05-19T19:40:01.926254Z","shell.execute_reply.started":"2025-05-19T19:40:01.915211Z","shell.execute_reply":"2025-05-19T19:40:01.924096Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Set seed for reproducibility:\n\nrandom.seed(cfg.seed)\nos.environ[\"PYTHONHASHSEED\"] = str(cfg.seed)\nnp.random.seed(cfg.seed)\ntorch.manual_seed(cfg.seed)\ntorch.cuda.manual_seed(cfg.seed)\ntorch.cuda.manual_seed_all(cfg.seed)\ntorch.backends.cudnn.deterministic = True\ntorch.backends.cudnn.benchmark = False","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-19T19:40:04.921952Z","iopub.execute_input":"2025-05-19T19:40:04.922401Z","iopub.status.idle":"2025-05-19T19:40:04.938157Z","shell.execute_reply.started":"2025-05-19T19:40:04.922373Z","shell.execute_reply":"2025-05-19T19:40:04.936519Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Dataset Preparation and Data Augmentations","metadata":{}},{"cell_type":"code","source":"class BirdCLEFDatasetFromNPY(Dataset):\n    \"\"\"\n    Custom PyTorch Dataset class for the BirdCLEF 2025 dataset.\n    It loads pre-computed mel spectrograms for training or validation\n    and applies data augmentations if necessary.\n\n    Args:\n        df (pd.DataFrame): DataFrame containing the data (e.g., train.csv or a subset).\n        cfg (CFG): Configuration object containing hyperparameters and settings.\n        spectrograms (dict, optional): Dictionary of pre-computed mel spectrograms.\n                                       Keys are sample `samplename`, values are the spectrogram numpy arrays.\n                                       Defaults to None.\n        mode (str, optional): Mode of the dataset, either 'train' or 'valid'.\n                              Defaults to 'train'.\n    \"\"\"\n\n    def __init__(self, df: pd.DataFrame, cfg: CFG, spectrograms: [dict[str, np.ndarray] | None]=None, mode: str=\"train\") -> None:\n        \"\"\"\n        Initializes the BirdCLEFDatasetFromNPY.\n        Sets up the DataFrame, configuration, spectrogram data, and label encoding.\n        \"\"\"\n\n        self.df: pd.DataFrame = df\n        self.cfg: CFG = cfg\n        self.mode: str = mode\n\n        self.spectrograms: [dict[str, np.ndarray] | None] = spectrograms\n        \n        taxonomy_df = pd.read_csv(self.cfg.taxonomy_csv)\n        self.species_ids: list[str] = taxonomy_df[\"primary_label\"].tolist()\n        self.num_classes: int = len(self.species_ids)\n        self.label_to_idx: dict[str, int] = {label: idx for idx, label in enumerate(self.species_ids)}\n\n        if \"filepath\" not in self.df.columns:\n            self.df[\"filepath\"] = self.cfg.train_data_dir + \"/\" + self.df.filename\n        \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        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        if cfg.debug:\n            self.df = self.df.sample(min(1000, len(self.df)), random_state=cfg.seed).reset_index(drop=True)\n    \n    def __len__(self) -> int:\n        \"\"\"\n        Returns the number of samples in the dataset.\n        \"\"\"\n\n        return len(self.df)\n    \n    def __getitem__(self, idx) -> dict[str, [torch.Tensor | str]]:\n        \"\"\"\n        Loads an item at the given index, applies preprocessing and augmentations.\n\n        Args:\n            idx (int): Index of the sample.\n\n        Returns:\n            dict: A dictionary containing the following keys:\n                  - 'melspec' (torch.Tensor): The mel spectrogram tensor.\n                  - 'target' (torch.Tensor): The one-hot encoded target label tensor.\n                  - 'filename' (str): The original filename.\n        \"\"\"\n\n        row = self.df.iloc[idx]\n        samplename = row[\"samplename\"]\n        spec = None\n\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\n        if spec is None:\n            spec = np.zeros(self.cfg.TARGET_SHAPE, dtype=np.float32)\n            if self.mode == \"train\":  # Only print warning during training\n                print(f\"Warning: Spectrogram for {samplename} not found and could not be generated\")\n\n        spec = torch.tensor(spec, dtype=torch.float32).unsqueeze(0)  # Add channel dimension\n\n        if self.mode == \"train\" and random.random() < self.cfg.aug_prob:\n            spec = self.apply_spec_augmentations(spec)\n        \n        target = self.encode_label(row[\"primary_label\"])\n        \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    def apply_spec_augmentations(self, spec: torch.Tensor) -> torch.Tensor:\n        \"\"\"\n        Applies data augmentations to the mel spectrogram.\n        Includes Time Masking, Frequency Masking, and random brightness/contrast adjustment.\n\n        Args:\n            spec (torch.Tensor): The mel spectrogram tensor to apply augmentations to.\n\n        Returns:\n            torch.Tensor: The mel spectrogram tensor with augmentations applied.\n        \"\"\"\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\n    \n    def encode_label(self, label: str) -> np.ndarray:\n        \"\"\"\n        Encodes a single label into a one-hot vector.\n\n        Args:\n            label (str): The label to encode.\n\n        Returns:\n            np.ndarray: The one-hot encoded label as a numpy array.\n        \"\"\"\n\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","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-19T19:40:38.128382Z","iopub.execute_input":"2025-05-19T19:40:38.128800Z","iopub.status.idle":"2025-05-19T19:40:38.150232Z","shell.execute_reply.started":"2025-05-19T19:40:38.128761Z","shell.execute_reply":"2025-05-19T19:40:38.149087Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def collate_fn(batch: list[dict[str, [torch.Tensor | str]]]) -> dict[str, [torch.Tensor | list[str]]]:\n    \"\"\"\n    Custom collate function to handle batches from BirdCLEFDatasetFromNPY.\n    This function is particularly useful for handling potential None items\n    returned by __getitem__ and for correctly stacking tensors in the batch.\n\n    Args:\n        batch (list[Optional[dict[str, Union[torch.Tensor, str]]]]): A list of samples from the dataset.\n                                                                      Each sample is a dictionary or None.\n\n    Returns:\n        dict[str, Union[torch.Tensor, list[str]]]: A dictionary where keys are the item names\n                                                  (e.g., 'melspec', 'target', 'filename')\n                                                  and values are either stacked tensors\n                                                  (for 'melspec' and 'target' if shapes are uniform)\n                                                  or a list (for 'filename').\n                                                  Returns an empty dictionary if the input batch is empty\n                                                  or contains only None values.\n    \"\"\"\n\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-19T19:40:52.605013Z","iopub.execute_input":"2025-05-19T19:40:52.605476Z","iopub.status.idle":"2025-05-19T19:40:52.615155Z","shell.execute_reply.started":"2025-05-19T19:40:52.605436Z","shell.execute_reply":"2025-05-19T19:40:52.613747Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model Definition","metadata":{}},{"cell_type":"code","source":"class BirdCLEFModel(nn.Module):\n    \"\"\"Deep learning model for bird song classification using a pre-trained backbone.\"\"\"\n    \n    def __init__(self, cfg) -> None:\n        \"\"\"\n        Initializes the BirdCLEFModel.\n        Sets up the backbone model and the classifier head.\n        \n        Args:\n            cfg (CFG): Configuration object containing model settings and hyperparameters.\n        \"\"\"\n        \n        super().__init__()\n        self.cfg = cfg\n        \n        taxonomy_df = pd.read_csv(cfg.taxonomy_csv)\n        cfg.num_classes = len(taxonomy_df) # Update num_classes in cfg\n        \n        self.backbone = timm.create_model(\n            cfg.model_name,\n            pretrained=cfg.pretrained,\n            in_chans=cfg.in_channels,\n            drop_rate=0.2,\n            drop_path_rate=0.2\n        )\n\n        # Determine the number of features from the backbone's output\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            # Generic approach for models with a get_classifier method\n            backbone_out = self.backbone.get_classifier().in_features\n            self.backbone.reset_classifier(0, \"\")\n        \n        self.pooling = nn.AdaptiveAvgPool2d(1)\n            \n        self.feat_dim = backbone_out\n        \n        self.classifier = nn.Linear(backbone_out, cfg.num_classes)\n        \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: torch.Tensor, targets: [torch.Tensor | None]=None) -> [torch.Tensor | tuple[torch.Tensor, torch.Tensor]]:\n        \"\"\"\n        Forward pass of the model. Optionally applies Mixup during training.\n\n        Args:\n            x (torch.Tensor): Input tensor (spectrogram batch). Shape (batch_size, channels, height, width).\n            targets (Optional[torch.Tensor], optional): Target labels for Mixup.\n                                                      Shape (batch_size, num_classes). Defaults to None.\n\n        Returns:\n            Union[torch.Tensor, tuple[torch.Tensor, torch.Tensor]]: If training with Mixup, returns a tuple\n                                                                  of (logits, loss). Otherwise, returns\n                                                                  the logits tensor.\n        \"\"\"\n\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        \n        features = self.backbone(x)\n\n        # Handle potential dictionary output from some backbones\n        if isinstance(features, dict):\n            features = features['features']\n\n        # Apply pooling if the output is 4D (convolutional features)\n        if len(features.shape) == 4:\n            features = self.pooling(features)\n            features = features.view(features.size(0), -1)\n        \n        logits = self.classifier(features)\n        \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            \n        return logits\n    \n    def mixup_data(self, x: torch.Tensor, targets: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, float]:\n        \"\"\"\n        Applies mixup to the data batch and targets.\n\n        Args:\n            x (torch.Tensor): Input tensor batch. Shape (batch_size, channels, height, width).\n            targets (torch.Tensor): Target labels batch. Shape (batch_size, num_classes).\n\n        Returns:\n            tuple[torch.Tensor, torch.Tensor, torch.Tensor, float]: A tuple containing:\n                                                                  - mixed_x (torch.Tensor): The mixed input tensor.\n                                                                  - targets_a (torch.Tensor): First set of targets.\n                                                                  - targets_b (torch.Tensor): Second set of targets.\n                                                                  - lam (float): The lambda value used for mixing.\n        \"\"\"\n\n        batch_size = x.size(0)\n\n        lam = np.random.beta(self.mixup_alpha, self.mixup_alpha)\n\n        indices = torch.randperm(batch_size).to(x.device)\n\n        mixed_x = lam * x + (1 - lam) * x[indices]\n        \n        return mixed_x, targets, targets[indices], lam\n    \n    def mixup_criterion(self, criterion: nn.Module, pred: torch.Tensor, y_a: torch.Tensor, y_b: torch.Tensor, lam: float) -> torch.Tensor:\n        \"\"\"\n        Applies mixup to the loss function.\n\n        Args:\n            criterion (nn.Module): The loss function (e.g., nn.BCEWithLogitsLoss).\n            pred (torch.Tensor): The model predictions (logits).\n            y_a (torch.Tensor): The first set of targets.\n            y_b (torch.Tensor): The second set of targets.\n            lam (float): The lambda value used for mixing.\n\n        Returns:\n            torch.Tensor: The mixed loss.\n        \"\"\"\n\n        return lam * criterion(pred, y_a) + (1 - lam) * criterion(pred, y_b)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-19T19:40:55.114000Z","iopub.execute_input":"2025-05-19T19:40:55.114366Z","iopub.status.idle":"2025-05-19T19:40:55.134668Z","shell.execute_reply.started":"2025-05-19T19:40:55.114340Z","shell.execute_reply":"2025-05-19T19:40:55.132804Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training Utilities","metadata":{}},{"cell_type":"code","source":"def get_optimizer(model: nn.Module, cfg: CFG) -> optim.Optimizer:\n    \"\"\"\n    Initializes and returns a PyTorch optimizer based on the configuration.\n\n    Args:\n        model (nn.Module): The PyTorch model for which the optimizer is created.\n        cfg (CFG): Configuration object containing optimizer type, learning rate, and weight decay.\n\n    Returns:\n        torch.optim.Optimizer: An instance of a PyTorch optimizer.\n\n    Raises:\n        NotImplementedError: If the optimizer specified in the configuration is not implemented.\n    \"\"\"\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","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-19T19:40:58.746787Z","iopub.execute_input":"2025-05-19T19:40:58.748016Z","iopub.status.idle":"2025-05-19T19:40:58.754862Z","shell.execute_reply.started":"2025-05-19T19:40:58.747975Z","shell.execute_reply":"2025-05-19T19:40:58.753301Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_scheduler(optimizer: optim.Optimizer, cfg: CFG) -> lr_scheduler._LRScheduler | lr_scheduler.ReduceLROnPlateau | None:\n    \"\"\"\n    Initializes and returns a PyTorch learning rate scheduler based on the configuration.\n\n    Args:\n        optimizer (torch.optim.Optimizer): The optimizer for which the scheduler is created.\n        cfg (CFG): Configuration object containing scheduler type and its specific parameters.\n\n    Returns:\n        torch.optim.lr_scheduler._LRScheduler | torch.optim.lr_scheduler.ReduceLROnPlateau | None:\n            An instance of a PyTorch learning rate scheduler, or None if no scheduler is specified\n            or the specified scheduler is 'OneCycleLR' (which might be handled differently).\n    \"\"\"\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","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-19T19:41:00.046052Z","iopub.execute_input":"2025-05-19T19:41:00.046448Z","iopub.status.idle":"2025-05-19T19:41:00.053598Z","shell.execute_reply.started":"2025-05-19T19:41:00.046417Z","shell.execute_reply":"2025-05-19T19:41:00.052357Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_criterion(cfg: CFG) -> nn.Module:\n    \"\"\"\n    Initializes and returns a PyTorch loss function based on the configuration.\n\n    Args:\n        cfg (CFG): Configuration object containing the criterion (loss function) type.\n\n    Returns:\n        torch.nn.Module: An instance of a PyTorch loss function.\n\n    Raises:\n        NotImplementedError: If the criterion specified in the configuration is not implemented.\n    \"\"\"\n    \n    if 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-19T19:41:01.329436Z","iopub.execute_input":"2025-05-19T19:41:01.329806Z","iopub.status.idle":"2025-05-19T19:41:01.338431Z","shell.execute_reply.started":"2025-05-19T19:41:01.329777Z","shell.execute_reply":"2025-05-19T19:41:01.336170Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training Functions","metadata":{}},{"cell_type":"code","source":"def train_one_epoch(\n    model: nn.Module,\n    loader: DataLoader,\n    optimizer: optim.Optimizer,\n    criterion: nn.Module,\n    device: str,\n    scheduler: lr_scheduler._LRScheduler | lr_scheduler.ReduceLROnPlateau | None = None\n) -> tuple[float, float]:\n    \"\"\"\n    Performs one training epoch for the model.\n\n    Iterates through the DataLoader, calculates the loss for each batch,\n    performs backpropagation, and updates the model's weights using the optimizer.\n    Optionally steps the learning rate scheduler if provided.\n\n    Args:\n        model (nn.Module): The PyTorch model to train.\n        loader (DataLoader): DataLoader providing the training data batches.\n        optimizer (torch.optim.Optimizer): The optimizer used for updating model weights.\n        criterion (nn.Module): The loss function.\n        device (str): The device to perform training on ('cuda' or 'cpu').\n        scheduler (torch.optim.lr_scheduler._LRScheduler | torch.optim.lr_scheduler.ReduceLROnPlateau | None, optional):\n            The learning rate scheduler. Expected to be stepped after each batch if it's a OneCycleLR,\n            otherwise stepped outside this function after the epoch. Defaults to None.\n\n    Returns:\n        tuple[float, float]: A tuple containing the average training loss and the\n                             average ROC AUC score for the epoch.\n    \"\"\"\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        # Handle the case where collate_fn might return lists of tensors\n        # (although the current collate_fn primarily stacks fixed-size tensors)\n    \n        if isinstance(batch[\"melspec\"], list):\n            batch_outputs = []\n            batch_losses = []\n            \n            for i in range(len(batch[\"melspec\"])):\n                # Ensure inputs and targets are tensors and on the correct device\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                # Assuming output is logits if mixup is not used in forward for single samples\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            # Standard batch processing\n            inputs = batch[\"melspec\"].to(device)\n            targets = batch[\"target\"].to(device)\n            \n            optimizer.zero_grad()\n            outputs = model(inputs) # Pass targets for potential mixup\n            \n            if isinstance(outputs, tuple):\n                # Model returned (logits, loss) due to mixup\n                outputs, loss = outputs  \n            else:\n                # Model returned logits\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        # Step scheduler if it's a OneCycleLR (stepped after each batch)\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    # Concatenate results from all batches\n    all_outputs = np.concatenate(all_outputs)\n    all_targets = np.concatenate(all_targets)\n\n    # Calculate AUC for the epoch\n    auc = calculate_auc(all_targets, all_outputs)\n    avg_loss = np.mean(losses)\n    \n    return avg_loss, auc","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-19T19:41:02.568962Z","iopub.execute_input":"2025-05-19T19:41:02.569287Z","iopub.status.idle":"2025-05-19T19:41:02.582687Z","shell.execute_reply.started":"2025-05-19T19:41:02.569268Z","shell.execute_reply":"2025-05-19T19:41:02.580981Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def validate(\n    model: nn.Module,\n    loader: DataLoader,\n    criterion: nn.Module,\n    device: str\n) -> tuple[float, float]:\n    \"\"\"\n    Evaluates the model on the validation set.\n\n    Iterates through the DataLoader in evaluation mode, calculates the loss\n    and performance metric (AUC) without backpropagation.\n\n    Args:\n        model (nn.Module): The PyTorch model to evaluate.\n        loader (DataLoader): DataLoader providing the validation data batches.\n        criterion (nn.Module): The loss function.\n        device (str): The device to perform evaluation on ('cuda' or 'cpu').\n\n    Returns:\n        tuple[float, float]: A tuple containing the average validation loss and the\n                             average ROC AUC score for the validation set.\n    \"\"\"\n\n    model.eval() # Set the model to evaluation mode\n    losses = []\n    all_targets = []\n    all_outputs = []\n\n    with torch.no_grad(): # Disable gradient calculation for evaluation\n        for batch in tqdm(loader, desc=\"Validation\"):\n            # Handle the case where collate_fn might return lists of tensors\n            if isinstance(batch[\"melspec\"], list):\n                batch_outputs = []\n                batch_losses = []\n                \n                for i in range(len(batch[\"melspec\"])):\n                    # Ensure inputs and targets are tensors and on the correct device\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) # No targets needed in forward for validation\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    # Concatenate results from all batches\n    all_outputs = np.concatenate(all_outputs)\n    all_targets = np.concatenate(all_targets)\n\n    # Calculate AUC for the validation set\n    auc = calculate_auc(all_targets, all_outputs)\n    avg_loss = np.mean(losses)\n    \n    return avg_loss, auc","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-19T19:41:04.673564Z","iopub.execute_input":"2025-05-19T19:41:04.673863Z","iopub.status.idle":"2025-05-19T19:41:04.685137Z","shell.execute_reply.started":"2025-05-19T19:41:04.673843Z","shell.execute_reply":"2025-05-19T19:41:04.683900Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def calculate_auc(targets: np.ndarray, outputs: np.ndarray) -> float:\n    \"\"\"\n    Calculates the mean ROC AUC score across all classes.\n\n    Computes the ROC AUC score for each class individually where there are\n    positive samples in the target, and then returns the average of these scores.\n\n    Args:\n        targets (np.ndarray): Ground truth labels (one-hot encoded). Shape (num_samples, num_classes).\n        outputs (np.ndarray): Model predictions (logits or probabilities). Shape (num_samples, num_classes).\n\n    Returns:\n        float: The mean ROC AUC score across all classes with at least one positive sample.\n               Returns 0.0 if there are no classes with positive samples.\n    \"\"\"\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-19T19:41:06.635984Z","iopub.execute_input":"2025-05-19T19:41:06.637402Z","iopub.status.idle":"2025-05-19T19:41:06.644240Z","shell.execute_reply.started":"2025-05-19T19:41:06.637364Z","shell.execute_reply":"2025-05-19T19:41:06.642703Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Load Data","metadata":{}},{"cell_type":"code","source":"train_df = pd.read_csv(cfg.train_csv)\ntaxonomy_df = pd.read_csv(cfg.taxonomy_csv)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-19T19:41:08.873691Z","iopub.execute_input":"2025-05-19T19:41:08.874255Z","iopub.status.idle":"2025-05-19T19:41:09.097454Z","shell.execute_reply.started":"2025-05-19T19:41:08.874226Z","shell.execute_reply":"2025-05-19T19:41:09.096064Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-19T19:41:13.805872Z","iopub.execute_input":"2025-05-19T19:41:13.806251Z","iopub.status.idle":"2025-05-19T19:41:13.841100Z","shell.execute_reply.started":"2025-05-19T19:41:13.806225Z","shell.execute_reply":"2025-05-19T19:41:13.839990Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"taxonomy_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-19T19:41:15.251656Z","iopub.execute_input":"2025-05-19T19:41:15.252096Z","iopub.status.idle":"2025-05-19T19:41:15.265389Z","shell.execute_reply.started":"2025-05-19T19:41:15.252065Z","shell.execute_reply":"2025-05-19T19:41:15.263745Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"taxonomy_df = pd.read_csv(cfg.taxonomy_csv)\nspecies_ids = taxonomy_df['primary_label'].tolist()\ncfg.num_classes = len(species_ids)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-19T19:41:16.287857Z","iopub.execute_input":"2025-05-19T19:41:16.288310Z","iopub.status.idle":"2025-05-19T19:41:16.299230Z","shell.execute_reply.started":"2025-05-19T19:41:16.288283Z","shell.execute_reply":"2025-05-19T19:41:16.298014Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"taxonomy_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-19T19:41:17.207826Z","iopub.execute_input":"2025-05-19T19:41:17.208420Z","iopub.status.idle":"2025-05-19T19:41:17.219065Z","shell.execute_reply.started":"2025-05-19T19:41:17.208390Z","shell.execute_reply":"2025-05-19T19:41:17.218138Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Run Training","metadata":{}},{"cell_type":"code","source":"print(cfg.debug)\nif cfg.debug:\n    cfg.update_debug_settings()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-19T19:41:18.634878Z","iopub.execute_input":"2025-05-19T19:41:18.635285Z","iopub.status.idle":"2025-05-19T19:41:18.641643Z","shell.execute_reply.started":"2025-05-19T19:41:18.635261Z","shell.execute_reply":"2025-05-19T19:41:18.640560Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"spectrograms = None","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-19T19:41:19.800296Z","iopub.execute_input":"2025-05-19T19:41:19.800627Z","iopub.status.idle":"2025-05-19T19:41:19.805475Z","shell.execute_reply.started":"2025-05-19T19:41:19.800602Z","shell.execute_reply":"2025-05-19T19:41:19.804420Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"Loading pre-computed mel spectrograms from NPY file...\")\nspectrograms = np.load(cfg.spectrogram_npy, allow_pickle=True).item()\nprint(f\"Loaded {len(spectrograms)} pre-computed mel spectrograms\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-19T19:41:22.389524Z","iopub.execute_input":"2025-05-19T19:41:22.389844Z","iopub.status.idle":"2025-05-19T19:42:20.054649Z","shell.execute_reply.started":"2025-05-19T19:41:22.389820Z","shell.execute_reply":"2025-05-19T19:42:20.051675Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"skf = StratifiedKFold(n_splits=cfg.n_fold, shuffle=True, random_state=cfg.seed)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-19T19:42:27.028188Z","iopub.execute_input":"2025-05-19T19:42:27.028779Z","iopub.status.idle":"2025-05-19T19:42:27.047487Z","shell.execute_reply.started":"2025-05-19T19:42:27.028743Z","shell.execute_reply":"2025-05-19T19:42:27.045549Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"best_scores = []","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-19T19:42:27.722786Z","iopub.execute_input":"2025-05-19T19:42:27.723179Z","iopub.status.idle":"2025-05-19T19:42:27.730318Z","shell.execute_reply.started":"2025-05-19T19:42:27.723149Z","shell.execute_reply":"2025-05-19T19:42:27.728329Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = train_df.copy()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-19T19:42:28.338231Z","iopub.execute_input":"2025-05-19T19:42:28.338964Z","iopub.status.idle":"2025-05-19T19:42:28.361791Z","shell.execute_reply.started":"2025-05-19T19:42:28.338913Z","shell.execute_reply":"2025-05-19T19:42:28.357348Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"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 = BirdCLEFDatasetFromNPY(train_df, cfg, spectrograms=spectrograms, mode='train')\n    val_dataset = BirdCLEFDatasetFromNPY(val_df, cfg, spectrograms=spectrograms, 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    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()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-19T19:42:34.924192Z","iopub.execute_input":"2025-05-19T19:42:34.924518Z","iopub.status.idle":"2025-05-19T19:43:09.546681Z","shell.execute_reply.started":"2025-05-19T19:42:34.924493Z","shell.execute_reply":"2025-05-19T19:43:09.544697Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Cross Validation","metadata":{}},{"cell_type":"code","source":"print(\"Cross-Validation Results:\")\nfor fold, score in enumerate(best_scores):\n    print(f\"Fold {cfg.selected_folds[fold]}: {score:.4f}\")\nprint(f\"Mean AUC: {np.mean(best_scores):.4f}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}