{"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"},{"sourceId":234320181,"sourceType":"kernelVersion"},{"sourceId":15853,"sourceType":"modelInstanceVersion","modelInstanceId":2739,"modelId":319}],"dockerImageVersionId":31040,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-23T17:47:28.219924Z","iopub.execute_input":"2025-05-23T17:47:28.22043Z","iopub.status.idle":"2025-05-23T17:47:28.224765Z","shell.execute_reply.started":"2025-05-23T17:47:28.220403Z","shell.execute_reply":"2025-05-23T17:47:28.224016Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport pandas as pd\nimport librosa\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\nimport os\nimport torch\nimport torchaudio\nimport torchaudio.transforms as T\nfrom torchaudio.transforms import MelSpectrogram\nfrom pathlib import Path\nfrom torchvision import transforms\nfrom tqdm import tqdm\nimport time\nimport cv2\nimport math\nfrom sklearn.model_selection import train_test_split\nfrom torch.utils.data import Dataset\nimport torchvision.transforms as transforms\nfrom torch.utils.data import DataLoader\nimport torch.nn.functional as F\nimport torch.nn as nn\nfrom torchvision.models.regnet import regnet_y_800mf\nimport pickle\nfrom sklearn.model_selection import StratifiedKFold\nimport os\nos.environ['TF_CPP_MIN_LOG_LEVEL'] = '3'\nimport tensorflow as tf\n\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-05-23T17:47:06.205183Z","iopub.execute_input":"2025-05-23T17:47:06.205758Z","iopub.status.idle":"2025-05-23T17:47:28.218873Z","shell.execute_reply.started":"2025-05-23T17:47:06.205729Z","shell.execute_reply":"2025-05-23T17:47:28.218243Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Config:\n\n    DEBUG_MODE = True\n\n    DATA_ROOT = \"/kaggle/input/birdclef-2025\"\n    FS = 32000\n\n    # Mel spectrogram parameters\n    N_FFT = 1024\n    HOP_LENGTH = 512\n    N_MELS = 128\n    FMIN = 50\n    FMAX = 14000\n\n    TARGET_DURATION = 5.0\n    TARGET_SHAPE = (256, 256)\n\n    N_MAX = 2800 if DEBUG_MODE else None\n\nconfig = Config()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(f\"Debug mode: {'ON' if config.DEBUG_MODE else 'OFF'}\")\nprint(f\"Max samples to process: {config.N_MAX if config.N_MAX is not None else 'ALL'}\")\n\nprint(\"Loading taxonomy data...\")\ntaxonomy_df = pd.read_csv(f'{config.DATA_ROOT}/taxonomy.csv')\nspecies_class_map = dict(zip(taxonomy_df['primary_label'], taxonomy_df['class_name']))\n\nprint(\"Loading training metadata...\")\ntrain_df = pd.read_csv(f'{config.DATA_ROOT}/train.csv')\n\n# human voice removal dataset\nprint(\"Loading voice summary data...\")\nwith open(\"/kaggle/input/bc25-separation-voice-from-data/train_voice_summary.txt\") as f:\n    voice_audio_list = f.read().split(\"\\n\")\n\nprint(\"Loading voice data...\")\nwith open(\"/kaggle/input/bc25-separation-voice-from-data/train_voice_data.pkl\", \"rb\") as f:\n    voice_data = pickle.load(f)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-23T17:47:28.225883Z","iopub.execute_input":"2025-05-23T17:47:28.226493Z","iopub.status.idle":"2025-05-23T17:47:28.578997Z","shell.execute_reply.started":"2025-05-23T17:47:28.226465Z","shell.execute_reply":"2025-05-23T17:47:28.578347Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"label_list = sorted(train_df['primary_label'].unique())\nlabel_id_list = list(range(len(label_list)))\nlabel2id = dict(zip(label_list, label_id_list))\nid2label = dict(zip(label_id_list, label_list))\n\nprint(f'Found {len(label_list)} unique species')\nworking_df = train_df[['primary_label', 'secondary_labels', 'rating', 'filename']].copy()\nworking_df['target'] = working_df.primary_label.map(label2id)\nworking_df['filepath'] = config.DATA_ROOT + '/train_audio/' + working_df.filename\nworking_df['samplename'] = working_df.filename.map(lambda x: x.split('/')[0] + '-' + x.split('/')[-1].split('.')[0])\nworking_df['class'] = working_df.primary_label.map(lambda x: species_class_map.get(x, 'Unknown'))\ntotal_samples = min(len(working_df), config.N_MAX or len(working_df))\nprint(f'Total samples to process: {total_samples} out of {len(working_df)} available')\nprint(f'Samples by class:')\nprint(working_df['class'].value_counts())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-23T17:47:31.592899Z","iopub.execute_input":"2025-05-23T17:47:31.593181Z","iopub.status.idle":"2025-05-23T17:47:31.650891Z","shell.execute_reply.started":"2025-05-23T17:47:31.59316Z","shell.execute_reply":"2025-05-23T17:47:31.650269Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(label2id)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-22T23:58:27.297545Z","iopub.execute_input":"2025-05-22T23:58:27.297872Z","iopub.status.idle":"2025-05-22T23:58:27.302821Z","shell.execute_reply.started":"2025-05-22T23:58:27.297847Z","shell.execute_reply":"2025-05-22T23:58:27.301743Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# perch model prediction as well defining class labels for it\nimport scipy.special \nPERCH_MODEL_PATH = \"/kaggle/input/bird-vocalization-classifier/tensorflow2/bird-vocalization-classifier/8\"\n\n# Load label.csv into a list of class labels\nlabel_path = \"/kaggle/input/bird-vocalization-classifier/tensorflow2/bird-vocalization-classifier/8/assets/label.csv\"\nlabel_df = pd.read_csv(label_path, header=None)\ngbvc_labels = label_df[0].tolist()[1:]\n\n# Load model from saved_model.pb\nmodel = tf.saved_model.load(PERCH_MODEL_PATH)\ninfer = model.signatures[\"serving_default\"]\n\ndef run_gbvc(audio, sample_rate=32000):\n    # Assumes audio is a float32 numpy array of shape (samples,)\n    audio_tensor = tf.convert_to_tensor(audio[np.newaxis, :], dtype=tf.float32)  # shape: (1, samples)\n    \n    # Run inference\n    outputs = infer(audio_tensor)\n    # print(\"Model output keys:\", outputs.keys())\n    # Get logits or probabilities\n    probs = scipy.special.softmax(outputs['label'].numpy()[0])  # shape: (1524,)\n    return probs","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-23T17:47:35.621907Z","iopub.execute_input":"2025-05-23T17:47:35.622185Z","iopub.status.idle":"2025-05-23T17:47:51.113631Z","shell.execute_reply.started":"2025-05-23T17:47:35.622163Z","shell.execute_reply":"2025-05-23T17:47:51.113055Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# single file testing\n# Load a 5-second audio file\naudio_path = \"/kaggle/input/birdclef-2025/train_audio/cotfly1/XC955594.ogg\"\naudio, _ = librosa.load(audio_path, sr=32000, duration=5.0)\n\n# Pad if too short\nif len(audio) < 5 * 32000:\n    audio = np.pad(audio, (0, 5 * 32000 - len(audio)), mode='constant')\n\n# Run prediction\nprobs = run_gbvc(audio)\n\n# # Top-5 predictions (index only)\ntop_indices = probs.argsort()[-5:][::-1]\nprint(np.argmax(probs))\nprint(probs.argsort()[-5:][::-1])\nprint(\"Top 5 prediction indices:\", top_indices)\nprint(\"Top 5 probabilities:\", probs[top_indices])\n# print(\"Top 5 predictions as percentages:\")\n# for idx in top_indices:\n#     label = gbvc_labels[idx]\n#     percent = probs[idx] * 100\n#     print(f\"{label}: {percent:.2f}%\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-23T18:22:50.832851Z","iopub.execute_input":"2025-05-23T18:22:50.833563Z","iopub.status.idle":"2025-05-23T18:22:58.113678Z","shell.execute_reply.started":"2025-05-23T18:22:50.833541Z","shell.execute_reply":"2025-05-23T18:22:58.112922Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# print(probs[1402]*100)\nidx=10603\npercent = probs[idx] * 100\nlabel=gbvc_labels[idx]\nprint(f\"{label}: {percent:.2f}%\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-23T18:24:35.453471Z","iopub.execute_input":"2025-05-23T18:24:35.454033Z","iopub.status.idle":"2025-05-23T18:24:35.457942Z","shell.execute_reply.started":"2025-05-23T18:24:35.454011Z","shell.execute_reply":"2025-05-23T18:24:35.457257Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Show top 5 predictions with class names and percentages\nprint(\"Top 5 predictions as percentages:\")\nfor idx in top_indices:\n    label = gbvc_labels[idx]\n    percent = probs[idx] * 100\n    print(f\"{label}: {percent:.2f}%\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-23T18:25:18.633643Z","iopub.execute_input":"2025-05-23T18:25:18.633924Z","iopub.status.idle":"2025-05-23T18:25:18.638714Z","shell.execute_reply.started":"2025-05-23T18:25:18.633907Z","shell.execute_reply":"2025-05-23T18:25:18.637895Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import hashlib\n# for checking duplicate files\n# def compute_file_hash(filepath):\n#     \"\"\"Return a SHA256 hash of the file contents.\"\"\"\n#     with open(filepath, 'rb') as f:\n#         return hashlib.sha256(f.read()).hexdigest()\n\n# # Add hash column\n# print(\"Computing file hashes to detect duplicates...\")\n# working_df['file_hash'] = working_df['filepath'].map(compute_file_hash)\n\n# # Find duplicates (keep first occurrence)\n# duplicates_df = working_df[working_df.duplicated(subset='file_hash', keep='first')]\n# clean_df = working_df.drop_duplicates(subset='file_hash', keep='first')\n\n# print(f\"Found {len(duplicates_df)} duplicate audio files.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-22T00:33:39.632342Z","iopub.execute_input":"2025-05-22T00:33:39.632654Z","iopub.status.idle":"2025-05-22T00:37:14.652349Z","shell.execute_reply.started":"2025-05-22T00:33:39.632634Z","shell.execute_reply":"2025-05-22T00:37:14.651067Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def audio2melspec(audio_data):\n    if np.isnan(audio_data).any():\n        mean_signal = np.nanmean(audio_data)\n        audio_data = np.nan_to_num(audio_data, nan=mean_signal)\n\n    audio_tensor = torch.tensor(audio_data, dtype=torch.float32).to(\"cuda\").unsqueeze(0)  # (1, N)\n\n    mel_transform = T.MelSpectrogram(\n    sample_rate=config.FS,\n    n_fft=config.N_FFT,\n    hop_length=config.HOP_LENGTH,\n    n_mels=config.N_MELS,\n    f_min=config.FMIN,\n    f_max=config.FMAX,\n    power=2.0\n    ).to(\"cuda\")\n\n    db_transform = T.AmplitudeToDB(stype=\"power\").to(\"cuda\")\n    \n    mel_spec = mel_transform(audio_tensor)\n    mel_spec_db = db_transform(mel_spec)\n\n    mel_spec_db -= mel_spec_db.min()\n    mel_spec_db /= mel_spec_db.max() + 1e-8\n\n    return mel_spec_db.squeeze(0).cpu().numpy().astype(np.float32)  # Return to CPU\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-23T18:52:00.636906Z","iopub.execute_input":"2025-05-23T18:52:00.637621Z","iopub.status.idle":"2025-05-23T18:52:00.643192Z","shell.execute_reply.started":"2025-05-23T18:52:00.637573Z","shell.execute_reply":"2025-05-23T18:52:00.642416Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def remove_voice(audio_data, voice_segment, config):\n    # creating a list of silent sections with no voice in it\n    non_voice_segments = []\n    prev_end = 0\n\n    for st in voice_segment:\n        start, end = math.floor(st[\"start\"] * config.FS), math.ceil(st[\"end\"] * config.FS)\n        if prev_end < start:\n            non_voice_segments.append((prev_end, start))\n        prev_end = max(prev_end, end)\n\n    if prev_end < len(audio_data):\n        non_voice_segments.append((prev_end, len(audio_data)))\n\n    # finding longest segment\n    longest_segment = None\n    max_duration = 0\n    for start, end in non_voice_segments:\n        segment = audio_data[start:end]\n        if len(segment) > max_duration:\n            max_duration = len(segment)\n            longest_segment = segment\n    return longest_segment","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-23T17:47:51.114723Z","iopub.execute_input":"2025-05-23T17:47:51.114962Z","iopub.status.idle":"2025-05-23T17:47:51.12028Z","shell.execute_reply.started":"2025-05-23T17:47:51.114945Z","shell.execute_reply":"2025-05-23T17:47:51.119499Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(working_df.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-23T03:06:14.951293Z","iopub.execute_input":"2025-05-23T03:06:14.951551Z","iopub.status.idle":"2025-05-23T03:06:14.96244Z","shell.execute_reply.started":"2025-05-23T03:06:14.951533Z","shell.execute_reply":"2025-05-23T03:06:14.961875Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# gpu based preprocessing\nprint(\"Starting gpu based audio processing...\")\nprint(f\"{'DEBUG MODE - Processing only ' + str(config.N_MAX) + ' samples' if config.DEBUG_MODE else 'FULL MODE - Processing all samples'}\")\nstart_time = time.time()\n\nall_bird_data = {}\nerrors = []\ndropped_samples = 0\nworking_df['soft_label'] = None\nbad_indices = [] \n\nfor i, row in tqdm(working_df.iterrows(), total=total_samples):\n    if config.N_MAX is not None and i >= config.N_MAX:\n        break\n\n    try:\n        primary = row['primary_label']\n        flag=1\n\n        # Parse secondary labels\n        secondary_raw = row.get(\"secondary_labels\", \"\")\n        try:\n            secondary = ast.literal_eval(secondary_raw) if isinstance(secondary_raw, str) else []\n        except:\n            secondary = []\n        secondary = [s for s in secondary if isinstance(s, str)]\n        \n        # useing torchaudio.load as an alternative\n        # audio_data, _ = librosa.load(row.filepath, sr=config.FS)\n        waveform, sr = torchaudio.load(row.filepath)\n        waveform = torchaudio.functional.resample(waveform, sr, config.FS)\n        audio_data = waveform.numpy().squeeze()\n\n        if row.filepath in voice_audio_list:\n            voice_segment = voice_data[row.filepath]\n            audio_data = remove_voice(audio_data, voice_segment, config)\n\n        target_samples = int(config.TARGET_DURATION * config.FS)\n        if len(audio_data) < target_samples:\n            n_copy = math.ceil(target_samples / len(audio_data))\n            if n_copy > 1:\n                audio_data = np.concatenate([audio_data] * n_copy)\n\n        start_idx = max(0, int(len(audio_data) / 2 - target_samples / 2))\n        end_idx = min(len(audio_data), start_idx + target_samples)\n        center_audio = audio_data[start_idx:end_idx]\n\n        if len(center_audio) < target_samples:\n            center_audio = np.pad(center_audio, (0, target_samples - len(center_audio)), mode='constant')\n\n        label_dict = {}\n\n        # Case: primary label starts with number (non-bird)\n        if primary[0].isdigit():\n            if not secondary:\n                label_dict[label2id[primary]] = 1.0\n            else:\n                label_dict[label2id[primary]] = 0.5\n                for s in secondary:\n                    label_dict[label2id[s]] = 0.5 / len(secondary)\n        \n        # Case: bird label — filter with Perch/GBVC\n        else:\n            with tf.device('/GPU:0'):\n                probs = run_gbvc(center_audio)\n            top_label = gbvc_labels[np.argmax(probs)]\n        \n            if top_label != primary and top_label not in secondary:\n                dropped_samples += 1\n                bad_indices.append(i)\n                flag=0\n                continue  # skip bad sample\n        \n            if not secondary:\n                label_dict[label2id[primary]] =  probs[gbvc_labels.index(primary)]\n                \n            else:\n                if top_label in secondary:\n                    primary = top_label\n                    secondary = [s for s in secondary if s != top_label]\n                label_dict[label2id[primary]] = probs[gbvc_labels.index(primary)] * 0.5\n                for s in secondary:\n                    prob_s = probs[gbvc_labels.index(s)] if s in gbvc_labels else 1.0\n                    label_dict[label2id[s]] = prob_s * (0.5 / len(secondary)) #label2id[s]\n        \n            # Add 0.05 * GBVC pseudo labels, only for matching BirdCLEF labels\n            for gbvc_label, prob in zip(gbvc_labels, probs):\n                if gbvc_label in label2id and gbvc_label != primary and prob >= 0.10 and gbvc_label not in secondary:\n                    label_dict[label2id[gbvc_label]] =  0.05 * prob\n        \n            # label_vec = np.clip(label_vec, 0.0, 1.0)\n\n        # Save to working_df and creating \n        if flag==1:\n            working_df.at[i, 'soft_label'] = label_dict\n            mel_spec = audio2melspec(center_audio)\n\n            if mel_spec.shape != config.TARGET_SHAPE:\n                mel_spec = cv2.resize(mel_spec, config.TARGET_SHAPE, interpolation=cv2.INTER_LINEAR)\n\n            all_bird_data[row.samplename] = mel_spec\n\n    except Exception as e:\n        print(f\"Error processing {row.filepath}: {e}\")\n        errors.append((row.filepath, str(e)))\n\n# Remove bad samples\nworking_df = working_df.drop(index=bad_indices).reset_index(drop=True)\nend_time = time.time()\nprint(f\"Processing completed in {end_time - start_time:.2f} seconds\")\nprint(f\"Successfully processed {len(working_df)} files out of {total_samples} total\")\nprint(f\"Dropped {dropped_samples} files\")\nprint(f\"Failed to process {len(errors)} files\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-23T19:21:29.039302Z","iopub.execute_input":"2025-05-23T19:21:29.039939Z","iopub.status.idle":"2025-05-23T19:35:41.691696Z","shell.execute_reply.started":"2025-05-23T19:21:29.039917Z","shell.execute_reply":"2025-05-23T19:35:41.690855Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Save the cleaned DataFrame to a CSV file in the /kaggle/working/ directory\nworking_df.to_csv('/kaggle/working/cleaned_working_df.csv', index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-23T23:31:54.605789Z","iopub.execute_input":"2025-05-23T23:31:54.606342Z","iopub.status.idle":"2025-05-23T23:31:54.736464Z","shell.execute_reply.started":"2025-05-23T23:31:54.60632Z","shell.execute_reply":"2025-05-23T23:31:54.735576Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(working_df.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-23T23:32:17.241166Z","iopub.execute_input":"2025-05-23T23:32:17.241692Z","iopub.status.idle":"2025-05-23T23:32:17.256764Z","shell.execute_reply.started":"2025-05-23T23:32:17.241669Z","shell.execute_reply":"2025-05-23T23:32:17.256175Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(f\"Successfully processed {len(working_df)} files out of {total_samples} total\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T00:10:56.068556Z","iopub.execute_input":"2025-05-24T00:10:56.069061Z","iopub.status.idle":"2025-05-24T00:10:56.073959Z","shell.execute_reply.started":"2025-05-24T00:10:56.069032Z","shell.execute_reply":"2025-05-24T00:10:56.073187Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nfrom torch.utils.data import Dataset\n\nclass BirdCLEFDataset(Dataset):\n    def __init__(self, df, mel_data, num_classes, image_size):\n        self.df = df\n        self.mel_data = mel_data\n        self.num_classes = num_classes\n        self.image_size = image_size\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        mel = self.mel_data[row['samplename']]  # shape: (H, W)\n\n        # Resize to match EfficientNetB2 input size (260x260)\n        mel = cv2.resize(mel, (self.image_size, self.image_size))\n        mel = np.stack([mel] * 3, axis=0) \n        mel_tensor = torch.tensor(mel, dtype=torch.float32)  # shape: (1, H, W)\n\n        # Normalize mel spectrogram (mean/std)\n        mel_tensor = (mel_tensor - mel_tensor.mean()) / (mel_tensor.std() + 1e-6)\n\n        # Build soft-label vector\n        label = torch.zeros(self.num_classes, dtype=torch.float32)\n        for class_id, weight in row['soft_label'].items():\n            # print('class_id')\n            # print(class_id)\n            # print('weight')\n            # print(weight)\n            label[int(class_id)] = float(weight)\n        # label = torch.tensor(label, dtype=torch.float32)\n        # print('tensor')\n        # print(mel_tensor.shape, label.shape)\n        return mel_tensor, label\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T00:31:16.247418Z","iopub.execute_input":"2025-05-24T00:31:16.247732Z","iopub.status.idle":"2025-05-24T00:31:16.254882Z","shell.execute_reply.started":"2025-05-24T00:31:16.247707Z","shell.execute_reply":"2025-05-24T00:31:16.254188Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\nfrom torch.utils.data import DataLoader\n\nNUM_CLASSES = 206\nIMAGE_SIZE = 224  # for b2\n\ntrain_df, val_df = train_test_split(working_df, test_size=0.2, random_state=42)\n\ntrain_dataset = BirdCLEFDataset(train_df, all_bird_data, NUM_CLASSES, IMAGE_SIZE)\nval_dataset = BirdCLEFDataset(val_df, all_bird_data, NUM_CLASSES, IMAGE_SIZE)\n\ntrain_loader = DataLoader(train_dataset, batch_size=32, shuffle=True, num_workers=4)\nval_loader = DataLoader(val_dataset, batch_size=32, shuffle=False, num_workers=4)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T00:31:18.139454Z","iopub.execute_input":"2025-05-24T00:31:18.139762Z","iopub.status.idle":"2025-05-24T00:31:18.153892Z","shell.execute_reply.started":"2025-05-24T00:31:18.139738Z","shell.execute_reply":"2025-05-24T00:31:18.1531Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import timm\nimport torch.nn as nn\n\ndef get_model():\n    model = timm.create_model(\"regnety_008\", pretrained=True, in_chans=3, num_classes=NUM_CLASSES)\n    return model\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T00:31:20.240095Z","iopub.execute_input":"2025-05-24T00:31:20.24037Z","iopub.status.idle":"2025-05-24T00:31:20.24431Z","shell.execute_reply.started":"2025-05-24T00:31:20.240347Z","shell.execute_reply":"2025-05-24T00:31:20.243655Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn.functional as F\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel = get_model().to(device)\noptimizer = torch.optim.Adam(model.parameters(), lr=1e-4)\n\ndef train_one_epoch():\n    model.train()\n    total_loss = 0\n    for mel, label in train_loader:\n        mel, label = mel.to(device), label.to(device)\n        optimizer.zero_grad()\n        logits = model(mel)\n        loss = F.binary_cross_entropy_with_logits(logits, label)\n        loss.backward()\n        optimizer.step()\n        total_loss += loss.item()\n    return total_loss / len(train_loader)\n\ndef validate():\n    model.eval()\n    total_loss = 0\n    with torch.no_grad():\n        for mel, label in val_loader:\n            mel, label = mel.to(device), label.to(device)\n            logits = model(mel)\n            loss = F.binary_cross_entropy_with_logits(logits, label)\n            total_loss += loss.item()\n    return total_loss / len(val_loader)\n\nfrom sklearn.metrics import f1_score\n\n# def validation_f1(threshold=0.5):\n#     model.eval()\n#     y_true = []\n#     y_pred_bin = []\n\n#     with torch.no_grad():\n#         for mel, label in val_loader:\n#             mel = mel.to(device)\n#             logits = model(mel)\n#             probs = torch.sigmoid(logits).cpu().numpy()  # convert logits to probs\n#             true = label.numpy()\n\n#             # Binarize predictions\n#             pred_bin = (probs >= threshold).astype(int)\n\n#             y_true.append(true)\n#             y_pred_bin.append(pred_bin)\n\n#     y_true = np.vstack(y_true)\n#     y_pred_bin = np.vstack(y_pred_bin)\n\n#     # Binarize predictions\n#     # y_pred_bin = (y_pred >= threshold).astype(int)\n\n#     # Macro-averaged F1 score\n#     f1 = f1_score(y_true, y_pred_bin, average='macro', zero_division=0)\n#     return f1\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T00:31:24.631474Z","iopub.execute_input":"2025-05-24T00:31:24.632001Z","iopub.status.idle":"2025-05-24T00:31:24.82001Z","shell.execute_reply.started":"2025-05-24T00:31:24.631977Z","shell.execute_reply":"2025-05-24T00:31:24.819271Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"EPOCHS = 10\n\nfor epoch in range(EPOCHS):\n    train_loss = train_one_epoch()\n    val_loss = validate()\n    # val_f1 = validation_f1()\n\n    print(f\"Epoch {epoch+1}/{EPOCHS} - Train Loss: {train_loss:.4f} | Val Loss: {val_loss:.4f} \")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T00:31:26.262867Z","iopub.execute_input":"2025-05-24T00:31:26.263537Z","iopub.status.idle":"2025-05-24T00:44:05.891945Z","shell.execute_reply.started":"2025-05-24T00:31:26.263513Z","shell.execute_reply":"2025-05-24T00:44:05.890993Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.save(model.state_dict(), 'regnety008_trained_xx.pth')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-24T00:44:20.298451Z","iopub.execute_input":"2025-05-24T00:44:20.299252Z","iopub.status.idle":"2025-05-24T00:44:20.37337Z","shell.execute_reply.started":"2025-05-24T00:44:20.299218Z","shell.execute_reply":"2025-05-24T00:44:20.372793Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}