{"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":31041,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport matplotlib.pyplot as plt\n\nimport gc\nimport random\nfrom pathlib import Path\nfrom tqdm.notebook import tqdm\nimport ast\nimport time\nimport json\n\nimport seaborn as sns\nimport librosa.display\n\nimport librosa\nimport soundfile as sf\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport torchaudio\n\nimport timm\nfrom torch.optim import Adam\n\nfrom sklearn.metrics import accuracy_score, f1_score, precision_score, recall_score\nfrom sklearn.model_selection import train_test_split\n\n\n\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-05-18T16:40:01.965167Z","iopub.execute_input":"2025-05-18T16:40:01.96548Z","iopub.status.idle":"2025-05-18T16:41:04.135117Z","shell.execute_reply.started":"2025-05-18T16:40:01.965457Z","shell.execute_reply":"2025-05-18T16:41:04.134327Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CFG:\n    # Paths\n    TRAIN_AUDIO_DIR = \"/kaggle/input/birdclef-2025/train_audio\"\n    TRAIN_SOUNDSCAPE_DIR = \"/kaggle/input/birdclef-2025/train_soundscapes\"\n    TEST_SOUNDSCAPE_DIR = \"/kaggle/input/birdclef-2025/test_soundscapes\"\n    TRAIN_CSV = \"/kaggle/input/birdclef-2025/train.csv\"\n    TAXONOMY_CSV = \"/kaggle/input/birdclef-2025/taxonomy.csv\"\n    SAMPLE_SUBMISSION_CSV = \"/kaggle/input/birdclef-2025/sample_submission.csv\"\n    RECORDING_LOCATION = \"/kaggle/input/birdclef-2025/recording_location.txt\"\n\n    # Audio Parameters\n    SR = 32000               \n    DURATION = 5             \n    HOP_LENGTH = 512        \n    N_MELS = 128             \n    FMIN = 20                \n    FMAX = SR // 2           \n\n    # Model Parameters\n    MODEL_NAME = \"tf_efficientnet_b0\"  \n    PRETRAINED = True\n    NUM_CLASSES = 206                  \n\n    # Training Parameters\n    EPOCHS = 1\n    BATCH_SIZE = 3\n    LR = 1e-4\n    SEED = 42\n    NUM_WORKERS = 2\n\n    # Model Parameters\n    MODEL_NAME = \"tf_efficientnet_b0\"  \n    PRETRAINED = True\n    NUM_CLASSES = 206                  \n\n    # Training Parameters\n    EPOCHS = 1\n    BATCH_SIZE = 3\n    LR = 1e-4\n    SEED = 42\n    NUM_WORKERS = 2\n\n    # Inference\n    THRESHOLD = 0.5\n    DEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n    # Debug Mode\n    DEBUG = False\n\ncfg = CFG()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-18T16:41:04.13658Z","iopub.execute_input":"2025-05-18T16:41:04.137085Z","iopub.status.idle":"2025-05-18T16:41:04.217434Z","shell.execute_reply.started":"2025-05-18T16:41:04.137053Z","shell.execute_reply":"2025-05-18T16:41:04.216777Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def set_seed(seed=42):\n    \"\"\"Set seed for reproducibility.\"\"\"\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\nset_seed(cfg.SEED)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-18T16:41:04.218332Z","iopub.execute_input":"2025-05-18T16:41:04.218644Z","iopub.status.idle":"2025-05-18T16:41:04.243744Z","shell.execute_reply.started":"2025-05-18T16:41:04.218617Z","shell.execute_reply":"2025-05-18T16:41:04.243075Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_taxonomy = pd.read_csv(cfg.TAXONOMY_CSV)\ndf_train = pd.read_csv(cfg.TRAIN_CSV)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-18T16:41:04.245277Z","iopub.execute_input":"2025-05-18T16:41:04.245473Z","iopub.status.idle":"2025-05-18T16:41:04.427474Z","shell.execute_reply.started":"2025-05-18T16:41:04.245457Z","shell.execute_reply":"2025-05-18T16:41:04.426822Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_train","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-18T16:41:04.428255Z","iopub.execute_input":"2025-05-18T16:41:04.428465Z","iopub.status.idle":"2025-05-18T16:41:04.465073Z","shell.execute_reply.started":"2025-05-18T16:41:04.428449Z","shell.execute_reply":"2025-05-18T16:41:04.464431Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Convert stringified lists to actual Python lists\nfor col in ['secondary_labels', 'type']:\n    df_train[col] = df_train[col].apply(lambda x: ast.literal_eval(x))\n\n# Add full path to audio files\ndf_train['filepath'] = df_train['filename'].apply(lambda x: os.path.join(cfg.TRAIN_AUDIO_DIR, x))\n\n# Preview\ndf_train.sample(5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-18T16:41:04.465859Z","iopub.execute_input":"2025-05-18T16:41:04.466106Z","iopub.status.idle":"2025-05-18T16:41:05.100335Z","shell.execute_reply.started":"2025-05-18T16:41:04.466087Z","shell.execute_reply":"2025-05-18T16:41:05.099567Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Check shape and column details\nprint(\"Shape of training data:\", df_train.shape)\nprint(\"Columns:\", df_train.columns.tolist())\ndf_train.info()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-18T16:41:05.101198Z","iopub.execute_input":"2025-05-18T16:41:05.101479Z","iopub.status.idle":"2025-05-18T16:41:05.136283Z","shell.execute_reply.started":"2025-05-18T16:41:05.10146Z","shell.execute_reply":"2025-05-18T16:41:05.135709Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Check missing values\ndf_train.isnull().sum()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-18T16:41:05.136863Z","iopub.execute_input":"2025-05-18T16:41:05.137118Z","iopub.status.idle":"2025-05-18T16:41:05.157094Z","shell.execute_reply.started":"2025-05-18T16:41:05.1371Z","shell.execute_reply":"2025-05-18T16:41:05.156323Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# How many unique bird labels?\nprint(\"Number of unique bird species:\", df_train['primary_label'].nunique())\n\n# Most common species\ndf_train['primary_label'].value_counts().head(10)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-18T16:41:05.158016Z","iopub.execute_input":"2025-05-18T16:41:05.1583Z","iopub.status.idle":"2025-05-18T16:41:05.175017Z","shell.execute_reply.started":"2025-05-18T16:41:05.158283Z","shell.execute_reply":"2025-05-18T16:41:05.174248Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.figure(figsize=(12, 8))\nsns.countplot(y='primary_label', data=df_train,\n              order=df_train['primary_label'].value_counts().iloc[:30].index)\nplt.title(\"Top 30 Most Frequent Bird Species\")\nplt.xlabel(\"Frequency\")\nplt.ylabel(\"Bird Species\")\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-18T16:41:05.177872Z","iopub.execute_input":"2025-05-18T16:41:05.178078Z","iopub.status.idle":"2025-05-18T16:41:05.607623Z","shell.execute_reply.started":"2025-05-18T16:41:05.178062Z","shell.execute_reply":"2025-05-18T16:41:05.606864Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Count unique values in 'rating'\nrating_counts = df_train['rating'].value_counts().sort_index()\nprint(rating_counts)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-18T16:41:05.60844Z","iopub.execute_input":"2025-05-18T16:41:05.608742Z","iopub.status.idle":"2025-05-18T16:41:05.618862Z","shell.execute_reply.started":"2025-05-18T16:41:05.608707Z","shell.execute_reply":"2025-05-18T16:41:05.618064Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.figure(figsize=(8, 5))\nsns.barplot(x=rating_counts.index, y=rating_counts.values, palette=\"viridis\")\n\nplt.title(\"Distribution of Ratings in Training Data\")\nplt.xlabel(\"Rating\")\nplt.ylabel(\"Count\")\nplt.xticks(rotation=0)\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-18T16:41:05.619601Z","iopub.execute_input":"2025-05-18T16:41:05.619836Z","iopub.status.idle":"2025-05-18T16:41:05.829112Z","shell.execute_reply.started":"2025-05-18T16:41:05.61982Z","shell.execute_reply":"2025-05-18T16:41:05.828355Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Plot class distribution of primary labels\nplt.figure(figsize=(16, 6))\ndf_train[\"primary_label\"].value_counts().plot(kind=\"bar\", color=\"skyblue\")\nplt.title(\"Distribution of Primary Labels in Train Set\")\nplt.xlabel(\"Bird Species (primary_label)\")\nplt.ylabel(\"Number of Samples\")\nplt.xticks(rotation=90)\nplt.grid(axis=\"y\", linestyle=\"--\", alpha=0.7)\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-18T16:41:05.830021Z","iopub.execute_input":"2025-05-18T16:41:05.830308Z","iopub.status.idle":"2025-05-18T16:41:07.026217Z","shell.execute_reply.started":"2025-05-18T16:41:05.830285Z","shell.execute_reply":"2025-05-18T16:41:07.025326Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":" #Plot distribution of primary_label for each rating (0 to 5)\nfor r in range(6):\n    df_rated = df_train[df_train[\"rating\"] == r]\n    \n    if df_rated.empty:\n        print(f\"No data found for rating = {r}\")\n        continue\n    \n    plt.figure(figsize=(16, 5))\n    df_rated[\"primary_label\"].value_counts().plot(kind=\"bar\", color=\"coral\")\n    plt.title(f\"Distribution of Primary Labels (Rating = {r})\")\n    plt.xlabel(\"Bird Species (primary_label)\")\n    plt.ylabel(\"Number of Samples\")\n    plt.xticks(rotation=90)\n    plt.grid(axis=\"y\", linestyle=\"--\", alpha=0.5)\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-18T16:41:07.027113Z","iopub.execute_input":"2025-05-18T16:41:07.027416Z","iopub.status.idle":"2025-05-18T16:41:12.641633Z","shell.execute_reply.started":"2025-05-18T16:41:07.027399Z","shell.execute_reply":"2025-05-18T16:41:12.640802Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Compute duration for each audio file\ndef get_duration(path):\n    f = sf.SoundFile(path)\n    return len(f) / f.samplerate\n\n# Apply to all audio files\ntqdm.pandas()\ndf_train[\"duration\"] = df_train[\"filepath\"].progress_apply(get_duration)\n\n# Plot the distribution\nplt.figure(figsize=(10, 5))\nsns.histplot(df_train[\"duration\"], bins=50, kde=True, color=\"teal\")\nplt.title(\"Distribution of Audio Durations\")\nplt.xlabel(\"Duration (seconds)\")\nplt.ylabel(\"Number of Files\")\nplt.grid(True, linestyle=\"--\", alpha=0.6)\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-18T16:41:12.642404Z","iopub.execute_input":"2025-05-18T16:41:12.642609Z","iopub.status.idle":"2025-05-18T16:45:47.600114Z","shell.execute_reply.started":"2025-05-18T16:41:12.642594Z","shell.execute_reply":"2025-05-18T16:45:47.599402Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Plot Mel Spectrogram\ndef plot_mel_spectrogram(path, sr=cfg.SR, n_mels=cfg.N_MELS, fmin=cfg.FMIN, fmax=cfg.FMAX, hop_length=cfg.HOP_LENGTH):\n    y, sr = librosa.load(path, sr=sr)\n    \n    mel_spec = librosa.feature.melspectrogram(\n        y=y, sr=sr, n_mels=n_mels, fmin=fmin, fmax=fmax, hop_length=hop_length\n    )\n    mel_spec_db = librosa.power_to_db(mel_spec, ref=np.max)\n\n    plt.figure(figsize=(12, 4))\n    librosa.display.specshow(\n        mel_spec_db, sr=sr, hop_length=hop_length, x_axis=\"time\", y_axis=\"mel\", fmax=fmax\n    )\n    plt.colorbar(format=\"%+2.0f dB\")\n    plt.title(\"Mel Spectrogram\")\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-18T16:45:47.600982Z","iopub.execute_input":"2025-05-18T16:45:47.601266Z","iopub.status.idle":"2025-05-18T16:45:47.606524Z","shell.execute_reply.started":"2025-05-18T16:45:47.601243Z","shell.execute_reply":"2025-05-18T16:45:47.605865Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Pick a random audio file\nsample_path = df_train.sample(1)[\"filepath\"].values[0]\nplot_mel_spectrogram(sample_path)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-18T16:45:47.607108Z","iopub.execute_input":"2025-05-18T16:45:47.607321Z","iopub.status.idle":"2025-05-18T16:46:00.418135Z","shell.execute_reply.started":"2025-05-18T16:45:47.607299Z","shell.execute_reply":"2025-05-18T16:46:00.417522Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Plot mel spectrograms for 5 random samples\nfor i, row in df_train.sample(5).iterrows():\n    print(f\"Sample {i+1} - Primary Label: {row['primary_label']}\")\n    plot_mel_spectrogram(row['filepath'])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-18T16:46:00.41897Z","iopub.execute_input":"2025-05-18T16:46:00.419425Z","iopub.status.idle":"2025-05-18T16:46:03.530853Z","shell.execute_reply.started":"2025-05-18T16:46:00.419405Z","shell.execute_reply":"2025-05-18T16:46:03.530188Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def preprocess_audio(path, sr=cfg.SR, duration=cfg.DURATION, n_mels=cfg.N_MELS,\n                     fmin=cfg.FMIN, fmax=cfg.FMAX, hop_length=cfg.HOP_LENGTH):\n    \n    y, sr = librosa.load(path, sr=sr, duration=duration)\n    \n    # Pad if audio is too short\n    expected_length = sr * duration\n    if len(y) < expected_length:\n        y = np.pad(y, (0, expected_length - len(y)))\n    else:\n        y = y[:expected_length]\n    \n    # Create mel spectrogram\n    mel = librosa.feature.melspectrogram(y=y, sr=sr, n_mels=n_mels,\n                                         fmin=fmin, fmax=fmax, hop_length=hop_length)\n    mel_db = librosa.power_to_db(mel, ref=np.max)\n\n    # Normalize to 0-1\n    mel_db -= mel_db.min()\n    mel_db /= mel_db.max()\n\n    return mel_db.astype(np.float32)  # shape: [n_mels, time]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-18T16:46:03.531738Z","iopub.execute_input":"2025-05-18T16:46:03.532002Z","iopub.status.idle":"2025-05-18T16:46:03.538972Z","shell.execute_reply.started":"2025-05-18T16:46:03.531978Z","shell.execute_reply":"2025-05-18T16:46:03.538132Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class BirdClefDataset(Dataset):\n    def __init__(self, df, label2id, transform=None):\n        self.df = df.reset_index(drop=True)\n        self.label2id = label2id\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        filepath = row[\"filepath\"]\n        label = row[\"primary_label\"]\n\n        # Preprocess audio\n        mel = preprocess_audio(filepath)  # shape: [n_mels, time]\n\n        if self.transform:\n            mel = self.transform(mel)\n\n        # Convert to tensor and add channel dimension\n        mel = torch.tensor(mel).unsqueeze(0)  # shape: [1, n_mels, time]\n\n        # Convert label to index\n        label_idx = self.label2id[label]\n\n        return mel, torch.tensor(label_idx, dtype=torch.long)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-18T16:46:03.539908Z","iopub.execute_input":"2025-05-18T16:46:03.540141Z","iopub.status.idle":"2025-05-18T16:46:03.560068Z","shell.execute_reply.started":"2025-05-18T16:46:03.540126Z","shell.execute_reply":"2025-05-18T16:46:03.559398Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Create label-to-index and index-to-label mappings\nunique_labels = sorted(df_train[\"primary_label\"].unique())\nlabel2id = {label: idx for idx, label in enumerate(unique_labels)}\nid2label = {idx: label for label, idx in label2id.items()}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-18T16:46:03.560859Z","iopub.execute_input":"2025-05-18T16:46:03.561081Z","iopub.status.idle":"2025-05-18T16:46:03.5765Z","shell.execute_reply.started":"2025-05-18T16:46:03.561055Z","shell.execute_reply":"2025-05-18T16:46:03.575994Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Create a small sample dataset\nsample_df = df_train.sample(5).reset_index(drop=True)\n\n# Instantiate the dataset\ndataset = BirdClefDataset(sample_df, label2id=label2id)\n\n# Test the first sample\nmel_tensor, label = dataset[0]\n\nprint(f\"Mel spectrogram shape: {mel_tensor.shape}\")  # Expected: [1, n_mels, time]\nprint(f\"Label index: {label} → {list(label2id.keys())[list(label2id.values()).index(label.item())]}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-18T16:46:03.577064Z","iopub.execute_input":"2025-05-18T16:46:03.577285Z","iopub.status.idle":"2025-05-18T16:46:03.622537Z","shell.execute_reply.started":"2025-05-18T16:46:03.577264Z","shell.execute_reply":"2025-05-18T16:46:03.621935Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.figure(figsize=(10, 4))\nplt.imshow(mel_tensor.squeeze(0).numpy(), aspect='auto', origin='lower')\nplt.title(\"Mel Spectrogram Tensor\")\nplt.xlabel(\"Time\")\nplt.ylabel(\"Mel Bins\")\nplt.colorbar(format=\"%+2.0f\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-18T16:46:03.623289Z","iopub.execute_input":"2025-05-18T16:46:03.623532Z","iopub.status.idle":"2025-05-18T16:46:03.909148Z","shell.execute_reply.started":"2025-05-18T16:46:03.623512Z","shell.execute_reply":"2025-05-18T16:46:03.908443Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class BirdCLEFModel(nn.Module):\n    def __init__(self, model_name=cfg.MODEL_NAME, num_classes=cfg.NUM_CLASSES, pretrained=cfg.PRETRAINED):\n        super(BirdCLEFModel, self).__init__()\n        \n        # Use timm backbone with no classifier, include pooling\n        self.backbone = timm.create_model(\n            model_name,\n            pretrained=pretrained,\n            in_chans=1,\n            num_classes=0,\n            global_pool='avg'  # this handles pooling internally\n        )\n        \n        # Add classifier\n        self.classifier = nn.Linear(self.backbone.num_features, num_classes)\n\n    def forward(self, x):\n        x = self.backbone(x)      # shape: (B, features)\n        x = self.classifier(x)    # shape: (B, num_classes)\n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-18T16:46:03.910012Z","iopub.execute_input":"2025-05-18T16:46:03.91024Z","iopub.status.idle":"2025-05-18T16:46:03.915036Z","shell.execute_reply.started":"2025-05-18T16:46:03.910223Z","shell.execute_reply":"2025-05-18T16:46:03.914462Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Create sample dataset and dataloader\nsample_df = df_train.sample(8).reset_index(drop=True)\ndataset = BirdClefDataset(sample_df, label2id=label2id)\ndataloader = DataLoader(dataset, batch_size=4, shuffle=False)\n\n# Instantiate model\nmodel = BirdCLEFModel().to(cfg.DEVICE)\n\n# Get a batch of data\nbatch = next(iter(dataloader))\ninputs, targets = batch\ninputs = inputs.to(cfg.DEVICE)\n\n# Forward pass\noutputs = model(inputs)\n\n# Display shapes\nprint(f\"Input shape      : {inputs.shape}\")   # [B, 1, 128, time_steps]\nprint(f\"Output shape     : {outputs.shape}\")  # [B, num_classes]\nprint(f\"Target labels     : {targets}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-18T16:46:03.915794Z","iopub.execute_input":"2025-05-18T16:46:03.916065Z","iopub.status.idle":"2025-05-18T16:46:05.557027Z","shell.execute_reply.started":"2025-05-18T16:46:03.916041Z","shell.execute_reply":"2025-05-18T16:46:05.556333Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_one_epoch(model, dataloader, optimizer, criterion):\n    model.train()\n    total_loss = 0\n    all_preds, all_labels = [], []\n\n    for inputs, labels in tqdm(dataloader, desc=\"Training\", leave=False):\n        inputs, labels = inputs.to(cfg.DEVICE), labels.to(cfg.DEVICE)\n\n        optimizer.zero_grad()\n        outputs = model(inputs)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n\n        total_loss += loss.item()\n        preds = torch.argmax(outputs, dim=1)\n        all_preds.extend(preds.cpu().numpy())\n        all_labels.extend(labels.cpu().numpy())\n\n    acc = accuracy_score(all_labels, all_preds)\n    return total_loss / len(dataloader), acc\n\n\n\ndef validate_one_epoch(model, dataloader, criterion):\n    model.eval()\n    total_loss = 0\n    all_preds, all_labels = [], []\n\n    with torch.no_grad():\n        for inputs, labels in tqdm(dataloader, desc=\"Validating\", leave=False):\n            inputs, labels = inputs.to(cfg.DEVICE), labels.to(cfg.DEVICE)\n            outputs = model(inputs)\n            loss = criterion(outputs, labels)\n\n            total_loss += loss.item()\n            preds = torch.argmax(outputs, dim=1)\n            all_preds.extend(preds.cpu().numpy())\n            all_labels.extend(labels.cpu().numpy())\n\n    acc = accuracy_score(all_labels, all_preds)\n    return total_loss / len(dataloader), acc","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-18T16:46:05.558115Z","iopub.execute_input":"2025-05-18T16:46:05.558448Z","iopub.status.idle":"2025-05-18T16:46:05.566255Z","shell.execute_reply.started":"2025-05-18T16:46:05.558419Z","shell.execute_reply":"2025-05-18T16:46:05.565445Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Training function with checkpoint saving\ndef train_model(model, train_loader, val_loader, epochs=cfg.EPOCHS, lr=cfg.LR, resume=False, checkpoint_path=None):\n    optimizer = Adam(model.parameters(), lr=lr)\n    criterion = nn.CrossEntropyLoss()\n\n    best_val_acc = 0.0\n    start_epoch = 0\n\n    # Load from checkpoint if resuming\n    if resume and checkpoint_path is not None:\n        checkpoint = torch.load(checkpoint_path)\n        model.load_state_dict(checkpoint['model_state_dict'])\n        optimizer.load_state_dict(checkpoint['optimizer_state_dict'])\n        start_epoch = checkpoint['epoch'] + 1\n        best_val_acc = checkpoint.get('val_acc', 0.0)\n        print(f\"Resumed from checkpoint: {checkpoint_path} | Starting at epoch {start_epoch + 1}\")\n\n    for epoch in range(start_epoch, epochs):\n        print(f\"\\n Epoch {epoch+1}/{epochs}\")\n\n        start = time.time()\n        train_loss, train_acc = train_one_epoch(model, train_loader, optimizer, criterion)\n        val_loss, val_acc = validate_one_epoch(model, val_loader, criterion)\n        end = time.time()\n\n        print(f\" Train Loss: {train_loss:.4f} | Accuracy: {train_acc:.4f}\")\n        print(f\" Val   Loss: {val_loss:.4f} | Accuracy: {val_acc:.4f}\")\n        print(f\" Time: {(end - start):.2f}s\")\n\n         # Save checkpoint every epoch\n        torch.save({\n            'epoch': epoch,\n            'model_state_dict': model.state_dict(),\n            'optimizer_state_dict': optimizer.state_dict(),\n            'val_acc': val_acc\n        }, f\"/kaggle/working/checkpoint_epoch_{epoch+1}.pth\")\n\n        # Save best model and label2id mapping\n        if val_acc > best_val_acc:\n            best_val_acc = val_acc\n\n            # Save model\n            torch.save(model.state_dict(), \"/kaggle/working/best_model.pth\")\n            print(\"Best model saved.\")\n\n            # Save label mapping\n            with open(\"/kaggle/working/label2id.json\", \"w\") as f:\n                json.dump(label2id, f)\n            print(\"label2id mapping saved.\")\n\n    print(f\"\\n Best Validation Accuracy: {best_val_acc:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-18T16:46:05.567048Z","iopub.execute_input":"2025-05-18T16:46:05.567372Z","iopub.status.idle":"2025-05-18T16:46:05.585923Z","shell.execute_reply.started":"2025-05-18T16:46:05.56735Z","shell.execute_reply":"2025-05-18T16:46:05.58523Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Spliting training data\ntrain_df, val_df = train_test_split(df_train, test_size=0.2, stratify=df_train['primary_label'], random_state=cfg.SEED)\n\ntrain_dataset = BirdClefDataset(train_df.reset_index(drop=True), label2id)\nval_dataset = BirdClefDataset(val_df.reset_index(drop=True), label2id)\n\ntrain_loader = DataLoader(train_dataset, batch_size=cfg.BATCH_SIZE, shuffle=True, num_workers=cfg.NUM_WORKERS)\nval_loader = DataLoader(val_dataset, batch_size=cfg.BATCH_SIZE, shuffle=False, num_workers=cfg.NUM_WORKERS)\n\n# Initialize model\nmodel = BirdCLEFModel().to(cfg.DEVICE)\n\n# Launch training (can resume later with resume=True)\ntrain_model(model, train_loader, val_loader, resume=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-18T16:46:05.589035Z","iopub.execute_input":"2025-05-18T16:46:05.589403Z","iopub.status.idle":"2025-05-18T16:53:09.175705Z","shell.execute_reply.started":"2025-05-18T16:46:05.589387Z","shell.execute_reply":"2025-05-18T16:53:09.174928Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_model(model, train_loader, val_loader, resume=True, checkpoint_path=\"checkpoint_epoch_1.pth\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-18T16:53:09.176863Z","iopub.execute_input":"2025-05-18T16:53:09.177105Z","iopub.status.idle":"2025-05-18T16:53:09.415518Z","shell.execute_reply.started":"2025-05-18T16:53:09.177083Z","shell.execute_reply":"2025-05-18T16:53:09.414536Z"}},"outputs":[],"execution_count":null}]}