{"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"}],"dockerImageVersionId":31040,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# BirdCLEF - Mulit Label Audio Classification Train ","metadata":{}},{"cell_type":"markdown","source":"# Data Loader","metadata":{}},{"cell_type":"markdown","source":"## Import","metadata":{}},{"cell_type":"code","source":"import ast\n\nimport torch\nimport torchaudio\nfrom torch.utils.data import Dataset, DataLoader\n\nimport numpy as np\nimport pandas as pd\n\nfrom pathlib import Path","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-05T16:55:31.486576Z","iopub.status.idle":"2025-06-05T16:55:31.486796Z","shell.execute_reply.started":"2025-06-05T16:55:31.486688Z","shell.execute_reply":"2025-06-05T16:55:31.486698Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Dataset Class","metadata":{}},{"cell_type":"code","source":"class BirdCLEFDataset(Dataset):\n    \n    def __init__(\n        self,\n        csv_path: str | Path,          \n        audio_dir: str | Path,\n        sample_rate: int = 32_000,\n        duration_sec: float = 5.0,\n        n_mels: int = 128,\n        n_fft: int = 1024,\n        win_length = 1024,\n        hop_length: int = 512,\n        f_min: int = 50,\n        f_max: int = 16_000,\n    ):\n        super().__init__()\n        \n        meta = (\n            pd.read_csv(csv_path)\n            .rename(columns={\"filename\": \"filepath\", \"primary_label\": \"primary\"})\n        )\n        meta[\"secondary\"] = meta[\"secondary_labels\"].apply(self._parse_secondary)\n        self.meta = meta[[\"filepath\", \"primary\", \"secondary\"]]\n\n        # label mapping\n        unique_labels = set(meta[\"primary\"])\n        for lst in meta[\"secondary\"]:\n            unique_labels.update(lst)\n\n        self.label2idx = {\n            lbl: idx for idx, lbl in enumerate(sorted(unique_labels))\n        }\n        self.idx2label = [lbl for lbl, _ in sorted(self.label2idx.items(), key=lambda x: x[1])]\n        \n        self.n_classes = len(self.label2idx)\n\n        self.audio_dir = Path(audio_dir)\n        self.sr = sample_rate\n        self.samples_per_clip = int(sample_rate * duration_sec)\n\n        # Pre-build transforms so they live on the CPU workers\n        self.mel = torchaudio.transforms.MelSpectrogram(\n            sample_rate    = sample_rate,\n            n_fft          = n_fft,\n            win_length     = win_length,\n            hop_length     = hop_length,\n            n_mels         = n_mels,\n            f_min          = f_min,\n            f_max          = f_max,\n        )\n        self.db  = torchaudio.transforms.AmplitudeToDB(top_db=80.0)\n\n    @staticmethod\n    def _parse_secondary(text: str) -> list[str]:\n        \"\"\"\n        Convert the string form of secondary labels to a list.\n    \n           Args:\n               text (str): List as string.\n           Returns:\n               list: Converted string.\n        \"\"\"\n        if pd.isna(text):\n            return []\n        labels = ast.literal_eval(text)\n        return [lbl for lbl in labels if lbl]  \n\n    def _load_wave(self, wav_path: Path) -> torch.Tensor:\n        \"\"\"\n        Read an audio file and return a mono 5 seconds, fixed-length waveform.\n    \n        Args:\n            wav_path (Path): Absolute or relative path to an audio file\n                readable by `torchaudio.load` (WAV, FLAC, Ogg, MP3 …).\n        Returns:\n            torch.Tensor: A 2-D tensor with shape  \n            `[1, self.samples_per_clip]` and `dtype=torch.float32`.\n    \n                * First dim = 1 (mono channel).  \n                * Second dim = timeline in **samples**.\n        Raises:\n            RuntimeError: If `torchaudio.load` cannot decode the file.   \n        \"\"\"\n        wav, _ = torchaudio.load(wav_path)          # shape [chn, time]\n        wav = torch.mean(wav, dim=0, keepdim=True)   # mono → [1, time]\n\n        # Pad / trim to fixed length\n        n = wav.shape[-1]\n        if n < self.samples_per_clip:                      # pad\n            pad_amt = self.samples_per_clip - n\n            wav = torch.nn.functional.pad(wav, (0, pad_amt))\n        elif n > self.samples_per_clip:                    # crop random (or centred)\n            start = torch.randint(0, n - self.samples_per_clip + 1, (1,)).item()\n            wav = wav[..., start : start+self.samples_per_clip]\n\n        return wav\n\n    def __len__(self) -> int:\n        return len(self.meta)\n\n    def __getitem__(self, idx: int):\n        row = self.meta.iloc[idx]\n        path = self.audio_dir / row[\"filepath\"]\n        wav = self._load_wave(path)               # [1, samples]\n        spec = self.db(self.mel(wav))              # [1, n_mels, time]\n\n        idxs = [self.label2idx[row.primary]] + [\n            self.label2idx[s] for s in row.secondary\n        ]\n        target = torch.zeros(self.n_classes, dtype=torch.float32)\n        target[idxs] = 1.0\n\n        return spec, target\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-05T16:55:31.487517Z","iopub.status.idle":"2025-06-05T16:55:31.487848Z","shell.execute_reply.started":"2025-06-05T16:55:31.487681Z","shell.execute_reply":"2025-06-05T16:55:31.487697Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Data Loader","metadata":{}},{"cell_type":"code","source":"def make_loader(csv_path: str, \n                audio_dir: str,\n                *,\n                batch_size: int = 32,\n                num_workers: int = 4,\n                shuffle: bool = True,\n                **dataset_kwargs) -> DataLoader:          \n    \"\"\"\n    Factory that returns a fully configured ``DataLoader``.\n\n    Args:\n        csv_path (str): Metadata CSV (one row per recording).\n        audio_dir (str): Root directory that contains the audio files.\n        batch_size (int): Items per batch yielded by the loader, default 32. \n        num_workers (int): Number of background worker processes, default 4.\n        shuffle (bool): Whether to reshuffle the dataset each epoch, default True.\n        **dataset_kwargs\n            Any additional arguments accepted by\n            :class:`BirdCLEFDataset` (sample-rate, mel params, etc.).\n    Returns:\n        torch.utils.data.DataLoader\n    \"\"\"\n\n    ds = BirdCLEFDataset(csv_path, audio_dir, **dataset_kwargs)\n\n    return DataLoader(\n        ds,\n        batch_size = batch_size,\n        shuffle = shuffle,\n        num_workers = num_workers,\n        pin_memory = True,\n        persistent_workers = True,\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-05T16:55:31.489300Z","iopub.status.idle":"2025-06-05T16:55:31.489637Z","shell.execute_reply.started":"2025-06-05T16:55:31.489439Z","shell.execute_reply":"2025-06-05T16:55:31.489454Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training\n\nInspired by: [This Notebook](https://www.kaggle.com/code/i2nfinit3y/bird2025-single-sed-model-inference-lb-0-857/notebook)","metadata":{}},{"cell_type":"code","source":"import datetime\nimport tqdm\n\nimport timm\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\n\nfrom typing import Dict, Tuple, Any, Optional","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-05T16:55:31.490436Z","iopub.status.idle":"2025-06-05T16:55:31.490769Z","shell.execute_reply.started":"2025-06-05T16:55:31.490625Z","shell.execute_reply":"2025-06-05T16:55:31.490660Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Own Attention Mechanism for Audio","metadata":{}},{"cell_type":"code","source":"def init_layer(layer: nn.Module): \n    \"\"\"\n    Applies Xavier uniform initialization to the weights.\n    To have better variance for the initial weights which leads to avoiding\n    e.g. sigmoids weaknesses of getting linear for weights close to 0.\n    \"\"\"\n    nn.init.xavier_uniform_(layer.weight)\n\n    # for beeing unbiased in the beginning if we have a bias\n    if hasattr(layer, \"bias\"):\n        if layer.bias is not None:\n            layer.bias.data.fill_(0.0)\n            \n\nclass AttBlockV2(nn.Module):\n    \"\"\"\n    Implements a temporal attention mechanism for audio classification.\n\n    This block takes a sequence of features (frame-level embeddings from a backbone CNN) and\n    computes a weighted average of these features to produce a single, clip-level representation.\n    The attention weights are learned, indicating which parts of the audio sequence are most\n    relevant for each class. This allows the model to selectively focus on informative segments\n    (e.g., where a bird call is present) rather than treating all time segments equally.\n    This is crucial for sound event detection and classification where target sounds might be\n    sparse within a long recording.\n    \"\"\"\n    def __init__(\n        self,\n        in_features: int,\n        out_features: int,\n        activation: str = \"sigmoid\"\n    ) -> None:\n        \"\"\"\n        Initializes the AttBlockV2.\n\n        Args:\n            in_features: Number of input features per time step from the backbone.\n            out_features: Number of output classes.\n            activation: Activation function to apply to the classification branch ('linear' or 'sigmoid').\n        \"\"\"\n        super().__init__()\n\n        self.activation = activation\n        self.att = nn.Conv1d(\n            in_channels=in_features,\n            out_channels=out_features,\n            kernel_size=1,\n            stride=1,\n            padding=0,\n            bias=True,\n        )\n        self.cla = nn.Conv1d(\n            in_channels=in_features,\n            out_channels=out_features,\n            kernel_size=1,\n            stride=1,\n            padding=0,\n            bias=True,\n        )\n\n        self.init_weights()\n\n    def init_weights(self) -> None:\n        \"\"\"\n        Initializes weights of the attention and classification convolutional layers.\n        \"\"\"\n        init_layer(self.att)\n        init_layer(self.cla)\n\n    def forward(self, x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:\n        \"\"\"\n        Performs the forward pass through the attention block.\n\n        Args:\n            x: Input tensor of shape (batch_size, in_features, n_time).\n\n        Returns:\n            A tuple containing:\n            - clipwise_output: (batch_size, out_features) - The clip-level aggregated features for each class.\n            - norm_att: (batch_size, out_features, n_time) - Normalized attention weights across time.\n            - cla_output: (batch_size, out_features, n_time) - Frame-level classification features (before weighted sum).\n        \"\"\"\n        # to learn attention scores \n        norm_att = torch.softmax(torch.tanh(self.att(x)), dim=-1)\n\n        # frame-level class scores\n        cla_output = self.nonlinear_transform(self.cla(x))\n\n        # combine those\n        clipwise_output = torch.sum(norm_att * cla_output, dim=2)\n\n        return clipwise_output, norm_att, cla_output\n\n    def nonlinear_transform(self, x: torch.Tensor) -> torch.Tensor:\n        \"\"\"\n        Applies the specified activation function to the input.\n\n        Args:\n            x: Input tensor.\n\n        Returns:\n            Output tensor after applying activation.\n        \"\"\"\n        if self.activation == \"linear\":\n            return x\n        elif self.activation == \"sigmoid\":\n            return torch.sigmoid(x)\n        else:\n            raise NotImplementedError(f\"Activation '{self.activation}' not supported.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-05T16:55:31.492222Z","iopub.status.idle":"2025-06-05T16:55:31.492473Z","shell.execute_reply.started":"2025-06-05T16:55:31.492361Z","shell.execute_reply":"2025-06-05T16:55:31.492372Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Our BirdCLEF Model","metadata":{}},{"cell_type":"code","source":"def image_delta(x: torch.Tensor) -> torch.Tensor:\n    \"\"\"\n    Computes first and second order differences along the time axis of a spectrogram.\n    This expands a single-channel spectrogram into a 3-channel (original, delta, delta-delta) input.\n    delta --> difference in mel energy betwen time frames\n    delta of delta --> acceleration\n\n    Args:\n        x: Input spectrogram tensor of shape (batch, 1, freq, time).\n\n    Returns:\n        A tensor of shape (batch, 3, freq, time) with original, delta, and delta-delta channels.\n    \"\"\"\n    if x.shape[1] != 1:\n        raise ValueError(\"image_delta expects input with 1 channel for this placeholder.\")\n\n    diff1_raw = x[:, :, :, 1:] - x[:, :, :, :-1]\n    delta1 = F.pad(diff1_raw, (0, 1, 0, 0), 'constant', 0)\n    \n    diff2_raw = delta1[:, :, :, 1:] - delta1[:, :, :, :-1]\n    delta2 = F.pad(diff2_raw, (0, 1, 0, 0), 'constant', 0)\n\n    return torch.cat([x, delta1, delta2], dim=1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-05T16:55:31.493445Z","iopub.status.idle":"2025-06-05T16:55:31.493738Z","shell.execute_reply.started":"2025-06-05T16:55:31.493574Z","shell.execute_reply":"2025-06-05T16:55:31.493587Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class BirdCLEFModel(nn.Module):\n    \"\"\"\n    A trainable PyTorch model for BirdCLEF audio classification.\n\n    This model takes Mel spectrograms as input, processes them through a\n    backbone CNN for feature extraction, and then uses an attention mechanism\n    to aggregate temporal features for clip-level classification.\n    \"\"\"\n    def __init__(\n        self, \n        model_name: str,\n        taxonomy_path: str,\n        n_mels: int,\n        sr: int,\n        duration_train: float,\n        infer_duration: float,\n        bias_for_linear_connection: bool,\n        attention_activation: str,\n        device,\n    ) -> None:\n        \"\"\"\n        Initializes the BirdCLEFModel.\n\n        Args:\n           model_name (str): Name of the base model we want to use\n           taxonomy_path (str): Path which has the meta data for the labels\n           n_mels (int): Number of Mel bins (frequency dimension).\n           SR (float): Sample rate (used for infer_duration calculation).\n           duration_train (float): Duration of training audio segments in seconds.\n           infer_duration (float): Duration of inference audio segments in seconds.\n           device (str): 'cuda' or 'cpu'.\n        \"\"\"\n        super().__init__()\n\n        taxonomy_df = pd.read_csv(taxonomy_path)\n        self.num_classes = len(taxonomy_df) \n\n        self.bn0 = nn.BatchNorm2d(n_mels) \n\n        self.backbone = timm.create_model(\n            model_name,\n            pretrained=True, # false for inference \n            in_chans=3,\n            drop_rate=0.2,\n            drop_path_rate=0.2,\n        )\n\n        layers = list(self.backbone.children())[:-2]\n        self.encoder = nn.Sequential(*layers)\n\n        # Determine the number of output features from the backbone's encoder\n        if \"efficientnet\" in model_name:\n            backbone_out = self.backbone.classifier.in_features\n        elif \"eca\" in model_name:\n            backbone_out = self.backbone.head.fc.in_features\n        elif \"res\" in model_name:\n            backbone_out = self.backbone.fc.in_features\n        else:\n            raise NotImplementedError(f\"Model: '{model_name}' not supported.\")\n\n        self.fc1 = nn.Linear(backbone_out, backbone_out, bias=True)\n\n        self.att_block = AttBlockV2(\n            backbone_out, \n            self.num_classes, \n            activation=attention_activation\n        )\n\n    def extract_feature(self, x: torch.Tensor) -> Tuple[torch.Tensor, int]:\n        \"\"\"\n        Extracts frame-level features from the spectrogram using the backbone encoder.\n\n        Args:\n            x: Input spectrogram tensor of shape (batch_size, channels, n_mels, n_frames).\n\n        Returns:\n            A tuple containing:\n            - Features tensor: (batch_size, backbone_out_features, n_frames_reduced)\n            - Original number of time frames.\n        \"\"\"\n        original_frames_num = x.shape[3]\n\n        x = x.transpose(1, 2)\n        x = self.bn0(x)\n        x = x.transpose(1, 2)\n\n        # The pretrained model from timm without head\n        x = self.encoder(x) \n\n        x = torch.mean(x, dim=2) \n\n        x1 = F.max_pool1d(x, kernel_size=3, stride=1, padding=1)\n        x2 = F.avg_pool1d(x, kernel_size=3, stride=1, padding=1)\n        x = x1 + x2\n\n        x = F.dropout(x, p=0.5, training=self.training)\n\n        x = x.transpose(1, 2)\n        x = F.relu_(self.fc1(x))\n        x = x.transpose(1, 2) \n\n        x = F.dropout(x, p=0.5, training=self.training)\n\n        return x, original_frames_num\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        \"\"\"\n        Forward pass for training.\n\n        Args:\n            x: Input Mel spectrogram tensor of shape (batch_size, channels, n_mels, n_frames).\n               Channels will be 1 for mono, to apply the image detla after it.\n               The model assumes normalization (e.g., to [0,1]) has already been applied.\n\n        Returns:\n            Logits for clip-wise classification output of shape (batch_size, num_classes).\n        \"\"\"\n        if x.shape[1] != 1: \n            raise ValueError(f\"Warning: Model expects 1 channel, but input has {x.shape[1]}.\")\n            \n        x = image_delta(x)\n        features, _ = self.extract_feature(x)\n\n        (clipwise_output, _, _) = self.att_block(features)\n\n        return torch.logit(clipwise_output)\n\n    def attention_infer(self, start: int, end: int, x: torch.Tensor) -> torch.Tensor:\n        \"\"\"\n        Helper function for inference, computing framewise predictions from a slice of features.\n\n        Args:\n            start: Start frame index for the feature slice.\n            end: End frame index for the feature slice.\n            x: Feature tensor of shape (batch_size, features, n_time).\n\n        Returns:\n            Max framewise probabilities for each class (batch_size, num_classes).\n        \"\"\"\n        feat_slice = x[:, :, start:end]\n        # Get framewise probabilities from the classification branch\n        framewise_pred = torch.sigmoid(self.att_block.cla(feat_slice))\n        # Take the maximum probability across frames for each class\n        framewise_pred_max = framewise_pred.max(dim=2)[0]\n        return framewise_pred_max\n\n    def infer(self, x: torch.Tensor, tta_delta: int = 2) -> torch.Tensor:\n        \"\"\"\n        Performs inference with Test Time Augmentation (TTA).\n\n        Args:\n            x: Input Mel spectrogram tensor of shape (batch_size, channels, n_mels, n_frames).\n            tta_delta: Number of frames to shift for TTA.\n\n        Returns:\n            Averaged clip-wise predictions (probabilities) of shape (batch_size, num_classes).\n        \"\"\"\n        with torch.no_grad():\n            if x.shape[1] != 1:\n               raise ValueError(f\"Warning: Model expects 1 channel, but input has {x.shape[1]}.\")\n                \n            x = image_delta(x)\n            features, _ = self.extract_feature(x)\n\n            feat_time = features.size(-1)\n\n            # Calculate central crop start and end points\n            start = int(feat_time / 2 - feat_time * (self.cfg['infer_duration'] / self.cfg['duration_train']) / 2)\n            end = int(start + feat_time * (self.cfg['infer_duration'] / self.cfg['duration_train']))\n\n            # Base --> Get prediction for central crop\n            pred = self.attention_infer(start, end, features)\n\n            # TTA --> shifted\n            start_minus = max(0, start - tta_delta)\n            end_minus = end - tta_delta\n            pred_minus = self.attention_infer(start_minus, end_minus, features)\n\n            start_plus = start + tta_delta\n            end_plus = min(feat_time, end + tta_delta)\n            pred_plus = self.attention_infer(start_plus, end_plus, features)\n\n            # combinbe it\n            final_pred = 0.5 * pred + 0.25 * pred_minus + 0.25 * pred_plus\n            return final_pred","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-05T16:55:31.495352Z","iopub.status.idle":"2025-06-05T16:55:31.495588Z","shell.execute_reply.started":"2025-06-05T16:55:31.495469Z","shell.execute_reply":"2025-06-05T16:55:31.495478Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training Config","metadata":{}},{"cell_type":"code","source":"cfg = {\n    'csv_path': '/kaggle/input/birdclef-2025/train.csv',\n    'audio_path': '/kaggle/input/birdclef-2025/train_audio',\n    'taxonomy_path': '/kaggle/input/birdclef-2025/taxonomy.csv'\n    'SR': 32000,\n    'target_duration': 10.0,\n    'hop_length': 512, #320,\n    'win_length': 1024,\n    'n_mels': 128,\n    'f_min': 20,\n    'f_max': 16000,\n    'n_fft': 1024, #2048,\n    'device': 'cuda' if torch.cuda.is_available() else 'cpu',\n    'in_channels': 1,\n    'model_name': 'efficientnet_b0',\n    'duration_train': 10,\n    'infer_duration': 5,\n    'num_workers': 4,\n    'batch_size': 16,\n    'learning_rate': 1e-4,\n    'num_epochs': 25\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-05T16:55:31.497051Z","iopub.status.idle":"2025-06-05T16:55:31.497420Z","shell.execute_reply.started":"2025-06-05T16:55:31.497229Z","shell.execute_reply":"2025-06-05T16:55:31.497243Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_loader = make_loader(\n    # Dataset kwargs\n    csv_path=cfg['csv_path'],\n    audio_dir=cfg['audio_path'],\n    sample_rate=cfg['SR'],\n    duration_sec=cfg['target_duration'],\n    n_mels=cfg['n_mels'],\n    n_fft=cfg['n_fft'],\n    win_length=cfg['win_length'],\n    hop_length=cfg['hop_length'],\n    f_min=cfg['f_min'],\n    f_max=cfg['f_max'],\n\n    # DataLoder kwargs\n    batch_size=cfg['batch_size'],\n    num_workers=cfg['num_workers'],\n    shuffle=True,\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-05T16:55:31.498641Z","iopub.status.idle":"2025-06-05T16:55:31.498891Z","shell.execute_reply.started":"2025-06-05T16:55:31.498787Z","shell.execute_reply":"2025-06-05T16:55:31.498797Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Possible Loss Class","metadata":{}},{"cell_type":"code","source":"class FocalLoss(nn.Module):\n    def __init__(self, alpha=0.25, gamma=2.0, reduction='mean'):\n        \"\"\"\n        Focal Loss for multi-label classification.\n        alpha: Can be a scalar, or a tensor of shape (num_classes,) for per-class weighting.\n               It often weights the positive class. A common starting point is 0.25 if positive class is rare.\n               Given your case (2 positive vs 204 negative), you might want alpha to effectively\n               upweight the positive examples. If alpha < 0.5, it downweights the class it applies to.\n               Alternatively, for multi-label, it can be tricky. Some use alpha for the positive class\n               and (1-alpha) for the negative class.\n               A simpler approach if you handle the main imbalance (2 vs 204) with alpha might be\n               to set alpha such that positive examples get higher weight (e.g. alpha_for_positives = 0.75).\n               Or, ensure your 'targets' are balanced by pos_weight in BCE_loss first and then apply focal modulation.\n               For simplicity here, let's assume alpha is applied to scale the loss of the positive samples.\n               A common interpretation is alpha for the positive class and 1-alpha for the negative class.\n               If you have few positives, an alpha > 0.5 for positive samples might be desired.\n               Let's use a simpler formulation where alpha weights all loss contributions (can be 1 if not needed)\n               and gamma is the main focusing parameter.\n        gamma: Focusing parameter.\n        reduction: 'mean', 'sum', or 'none'.\n        \"\"\"\n        super(FocalLoss, self).__init__()\n        self.alpha = alpha \n        self.gamma = gamma\n        self.reduction = reduction\n\n    def forward(self, inputs, targets): \n        \n        BCE_loss = F.binary_cross_entropy_with_logits(inputs, targets, reduction='none')\n        probs = torch.sigmoid(inputs)\n        \n        pt = torch.where(targets == 1, probs, 1 - probs)\n        \n        alpha_factor = torch.where(targets == 1, self.alpha, 1 - self.alpha)\n        \n        modulating_factor = (1.0 - pt).pow(self.gamma)\n        \n        focal_loss = alpha_factor * modulating_factor * BCE_loss\n\n        if self.reduction == 'mean':\n            return focal_loss.mean()\n        elif self.reduction == 'sum':\n            return focal_loss.sum()\n        else: # 'none'\n            return focal_loss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-05T16:55:31.499439Z","iopub.status.idle":"2025-06-05T16:55:31.499677Z","shell.execute_reply.started":"2025-06-05T16:55:31.499560Z","shell.execute_reply":"2025-06-05T16:55:31.499570Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Final Train Loop","metadata":{}},{"cell_type":"code","source":"model = BirdCLEFModel(\n    model_name = cfg['model_name'],\n    taxonomy_path = cfg['taxonomy_path'],\n    n_mels = cfg['n_mels'],\n    sr = cfg['sr'],\n    duration_train = cfg['duration_train'],\n    infer_duration = cfg['infer_duration'],\n    bias_for_linear_connection = cfg['bias_for_linear_connection'],\n    attention_activation = cfg['attention_activation'],\n    device= cfg['device'], \n).to(cfg['device'])\n\n# Loss calc based on Binary Cross Entropy with Logits --> multi labes (softmax only single label)\ncriterion = nn.BCEWithLogitsLoss()\n\noptimizer = optim.Adam(model.parameters(), lr=cfg['learning_rate'])\nscheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=cfg['num_epochs'])\n\n# --- Training Loop ---\nprint(\"\\nStarting training loop...\")\nfor epoch in range(cfg['num_epochs']):\n    model.train() \n    running_loss = 0.0\n\n    for i, (spectrograms, labels) in enumerate(train_loader):\n        spectrograms = spectrograms.to(cfg['device'])\n        labels = labels.to(cfg['device'])\n\n        optimizer.zero_grad() \n\n        outputs = model(spectrograms)\n        loss = criterion(outputs, labels)\n\n        loss.backward()\n        optimizer.step()\n\n        running_loss += loss.item()\n\n    scheduler.step()\n\n    print(f\"Epoch {epoch+1}/{cfg['num_epochs']}, Loss: {running_loss / len(train_loader):.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-05T16:55:31.501461Z","iopub.status.idle":"2025-06-05T16:55:31.501828Z","shell.execute_reply.started":"2025-06-05T16:55:31.501631Z","shell.execute_reply":"2025-06-05T16:55:31.501646Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Save Model","metadata":{}},{"cell_type":"code","source":"now = datetime.datetime.now()\ndate = now.strftime('%Y-%m-%d_%H-%M-%S')\nmodel_save_path = f'birdclef_model_{date}.pth'\ntorch.save(model.state_dict(), model_save_path)\nprint(f\"\\nTraining complete. Model saved to {model_save_path}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-05T16:55:31.970809Z","iopub.execute_input":"2025-06-05T16:55:31.971148Z","iopub.status.idle":"2025-06-05T16:55:32.069121Z","shell.execute_reply.started":"2025-06-05T16:55:31.971121Z","shell.execute_reply":"2025-06-05T16:55:32.068494Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}