{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":91844,"databundleVersionId":11361821,"sourceType":"competition"}],"dockerImageVersionId":31012,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Imports","metadata":{}},{"cell_type":"code","source":"# General\nimport os\nimport gc\nimport random\nimport numpy as np\nimport pandas as pd\nfrom pathlib import Path\nfrom tqdm.notebook import tqdm\nimport ast\nimport time\nimport json\n\n# Visualization\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport librosa.display\n\n# Audio processing\nimport librosa\nimport soundfile as sf\n\n# Machine Learning / Deep Learning\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport torchaudio\n\n# Model and Optimizer\nimport timm\nfrom torch.optim import Adam\n\n# Metrics\nfrom sklearn.metrics import accuracy_score, f1_score, precision_score, recall_score\nfrom sklearn.model_selection import train_test_split\n\n# Ignore warnings\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\n# Set display options\npd.set_option('display.max_columns', 100)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-05T12:26:32.701424Z","iopub.execute_input":"2025-05-05T12:26:32.701750Z","iopub.status.idle":"2025-05-05T12:26:48.052068Z","shell.execute_reply.started":"2025-05-05T12:26:32.701698Z","shell.execute_reply":"2025-05-05T12:26:48.051164Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Config Setup","metadata":{}},{"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    # 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-05T12:26:54.907395Z","iopub.execute_input":"2025-05-05T12:26:54.907745Z","iopub.status.idle":"2025-05-05T12:26:54.917512Z","shell.execute_reply.started":"2025-05-05T12:26:54.907694Z","shell.execute_reply":"2025-05-05T12:26:54.916472Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Helper Function","metadata":{}},{"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-05T12:26:59.536768Z","iopub.execute_input":"2025-05-05T12:26:59.537139Z","iopub.status.idle":"2025-05-05T12:26:59.549356Z","shell.execute_reply.started":"2025-05-05T12:26:59.537115Z","shell.execute_reply":"2025-05-05T12:26:59.548442Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Loading Dataset","metadata":{}},{"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-05T12:27:02.021364Z","iopub.execute_input":"2025-05-05T12:27:02.022109Z","iopub.status.idle":"2025-05-05T12:27:02.244852Z","shell.execute_reply.started":"2025-05-05T12:27:02.022074Z","shell.execute_reply":"2025-05-05T12:27:02.243687Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_train","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-05T12:27:05.532823Z","iopub.execute_input":"2025-05-05T12:27:05.533129Z","iopub.status.idle":"2025-05-05T12:27:05.573465Z","shell.execute_reply.started":"2025-05-05T12:27:05.533104Z","shell.execute_reply":"2025-05-05T12:27:05.572495Z"}},"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-05T12:27:08.620087Z","iopub.execute_input":"2025-05-05T12:27:08.620400Z","iopub.status.idle":"2025-05-05T12:27:09.338852Z","shell.execute_reply.started":"2025-05-05T12:27:08.620376Z","shell.execute_reply":"2025-05-05T12:27:09.338002Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Basic Analysis","metadata":{}},{"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-05T12:27:12.361250Z","iopub.execute_input":"2025-05-05T12:27:12.361546Z","iopub.status.idle":"2025-05-05T12:27:12.402054Z","shell.execute_reply.started":"2025-05-05T12:27:12.361525Z","shell.execute_reply":"2025-05-05T12:27:12.401186Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Check missing values\ndf_train.isnull().sum()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-05T12:27:15.092257Z","iopub.execute_input":"2025-05-05T12:27:15.092871Z","iopub.status.idle":"2025-05-05T12:27:15.117017Z","shell.execute_reply.started":"2025-05-05T12:27:15.092846Z","shell.execute_reply":"2025-05-05T12:27:15.116305Z"}},"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-05T12:27:18.171323Z","iopub.execute_input":"2025-05-05T12:27:18.171622Z","iopub.status.idle":"2025-05-05T12:27:18.183025Z","shell.execute_reply.started":"2025-05-05T12:27:18.171598Z","shell.execute_reply":"2025-05-05T12:27:18.182208Z"}},"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-05T12:27:20.976438Z","iopub.execute_input":"2025-05-05T12:27:20.977011Z","iopub.status.idle":"2025-05-05T12:27:21.517538Z","shell.execute_reply.started":"2025-05-05T12:27:20.976986Z","shell.execute_reply":"2025-05-05T12:27:21.516621Z"}},"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-05T12:27:24.968791Z","iopub.execute_input":"2025-05-05T12:27:24.969105Z","iopub.status.idle":"2025-05-05T12:27:24.981127Z","shell.execute_reply.started":"2025-05-05T12:27:24.969083Z","shell.execute_reply":"2025-05-05T12:27:24.979894Z"}},"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-05T12:27:27.405009Z","iopub.execute_input":"2025-05-05T12:27:27.405304Z","iopub.status.idle":"2025-05-05T12:27:27.653564Z","shell.execute_reply.started":"2025-05-05T12:27:27.405282Z","shell.execute_reply":"2025-05-05T12:27:27.652766Z"}},"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-05T12:27:30.765490Z","iopub.execute_input":"2025-05-05T12:27:30.766189Z","iopub.status.idle":"2025-05-05T12:27:32.228040Z","shell.execute_reply.started":"2025-05-05T12:27:30.766162Z","shell.execute_reply":"2025-05-05T12:27:32.227126Z"}},"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-05T12:27:36.290884Z","iopub.execute_input":"2025-05-05T12:27:36.291182Z","iopub.status.idle":"2025-05-05T12:27:42.977405Z","shell.execute_reply.started":"2025-05-05T12:27:36.291161Z","shell.execute_reply":"2025-05-05T12:27:42.976168Z"}},"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-05T12:27:44.551411Z","iopub.execute_input":"2025-05-05T12:27:44.552432Z","iopub.status.idle":"2025-05-05T12:32:32.645266Z","shell.execute_reply.started":"2025-05-05T12:27:44.552403Z","shell.execute_reply":"2025-05-05T12:32:32.644190Z"}},"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-05T12:32:50.288462Z","iopub.execute_input":"2025-05-05T12:32:50.288835Z","iopub.status.idle":"2025-05-05T12:32:50.295869Z","shell.execute_reply.started":"2025-05-05T12:32:50.288807Z","shell.execute_reply":"2025-05-05T12:32:50.294700Z"}},"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-05T12:32:53.158237Z","iopub.execute_input":"2025-05-05T12:32:53.159153Z","iopub.status.idle":"2025-05-05T12:33:09.576352Z","shell.execute_reply.started":"2025-05-05T12:32:53.159123Z","shell.execute_reply":"2025-05-05T12:33:09.575427Z"}},"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-05T12:33:14.125852Z","iopub.execute_input":"2025-05-05T12:33:14.126877Z","iopub.status.idle":"2025-05-05T12:33:17.747193Z","shell.execute_reply.started":"2025-05-05T12:33:14.126847Z","shell.execute_reply":"2025-05-05T12:33:17.746299Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Preprocessing Function – Audio → Mel Spectrogram (Fixed Size)","metadata":{}},{"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-05T12:33:21.732500Z","iopub.execute_input":"2025-05-05T12:33:21.733287Z","iopub.status.idle":"2025-05-05T12:33:21.739449Z","shell.execute_reply.started":"2025-05-05T12:33:21.733259Z","shell.execute_reply":"2025-05-05T12:33:21.738562Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Build Dataset Class for Training","metadata":{}},{"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-05T12:33:26.755230Z","iopub.execute_input":"2025-05-05T12:33:26.756177Z","iopub.status.idle":"2025-05-05T12:33:26.762952Z","shell.execute_reply.started":"2025-05-05T12:33:26.756143Z","shell.execute_reply":"2025-05-05T12:33:26.761799Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Create Label Mapping","metadata":{}},{"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-05T12:33:30.543848Z","iopub.execute_input":"2025-05-05T12:33:30.544159Z","iopub.status.idle":"2025-05-05T12:33:30.550739Z","shell.execute_reply.started":"2025-05-05T12:33:30.544136Z","shell.execute_reply":"2025-05-05T12:33:30.549763Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Try Dataset for Clerity","metadata":{}},{"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-05T12:33:33.694904Z","iopub.execute_input":"2025-05-05T12:33:33.695954Z","iopub.status.idle":"2025-05-05T12:33:33.756079Z","shell.execute_reply.started":"2025-05-05T12:33:33.695917Z","shell.execute_reply":"2025-05-05T12:33:33.755212Z"}},"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-05T12:33:36.629219Z","iopub.execute_input":"2025-05-05T12:33:36.629529Z","iopub.status.idle":"2025-05-05T12:33:36.988051Z","shell.execute_reply.started":"2025-05-05T12:33:36.629502Z","shell.execute_reply":"2025-05-05T12:33:36.987208Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model","metadata":{}},{"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-05T12:33:40.681062Z","iopub.execute_input":"2025-05-05T12:33:40.681365Z","iopub.status.idle":"2025-05-05T12:33:40.687576Z","shell.execute_reply.started":"2025-05-05T12:33:40.681343Z","shell.execute_reply":"2025-05-05T12:33:40.686658Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Test Dataset + Model Compatibility","metadata":{}},{"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-05T12:33:44.108309Z","iopub.execute_input":"2025-05-05T12:33:44.108604Z","iopub.status.idle":"2025-05-05T12:33:45.255005Z","shell.execute_reply.started":"2025-05-05T12:33:44.108583Z","shell.execute_reply":"2025-05-05T12:33:45.254165Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training Loop","metadata":{}},{"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\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-05T12:33:49.908589Z","iopub.execute_input":"2025-05-05T12:33:49.909182Z","iopub.status.idle":"2025-05-05T12:33:49.917641Z","shell.execute_reply.started":"2025-05-05T12:33:49.909151Z","shell.execute_reply":"2025-05-05T12:33:49.916944Z"}},"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-05T12:33:55.021818Z","iopub.execute_input":"2025-05-05T12:33:55.022097Z","iopub.status.idle":"2025-05-05T12:33:55.031245Z","shell.execute_reply.started":"2025-05-05T12:33:55.022080Z","shell.execute_reply":"2025-05-05T12:33:55.030183Z"}},"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)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-05T12:34:11.907499Z","iopub.execute_input":"2025-05-05T12:34:11.907843Z","iopub.status.idle":"2025-05-05T13:21:16.058612Z","shell.execute_reply.started":"2025-05-05T12:34:11.907814Z","shell.execute_reply":"2025-05-05T13:21:16.057529Z"}},"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-04T15:16:41.094698Z","iopub.execute_input":"2025-05-04T15:16:41.095010Z","iopub.status.idle":"2025-05-04T15:16:41.221659Z","shell.execute_reply.started":"2025-05-04T15:16:41.094986Z","shell.execute_reply":"2025-05-04T15:16:41.219549Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}