{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.11.11"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":91844,"databundleVersionId":11361821,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":11252438,"sourceType":"datasetVersion","datasetId":7031717}],"dockerImageVersionId":31041,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true},"papermill":{"default_parameters":{},"duration":7379.999251,"end_time":"2025-05-22T08:33:28.254184","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2025-05-22T06:30:28.254933","version":"2.6.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Pseudo-Label Generation\n\n\n**Goal:** For each 5-second segment of each soundscape, we assign a pseudo-label *only if the predicted bird species has confidence $\\geq$ a non-linear dynamic threshold*.  \nThe threshold is **lower for rare species** and **higher for common ones**, based on their frequency in the original training data.\n\n**Training notebook:** [EfficientNet-B0 Train w/ Pseudo-Label](https://www.kaggle.com/code/sietze535/efficientnet-b0-train-w-pseudo-label)\n\n","metadata":{"papermill":{"duration":0.002095,"end_time":"2025-05-22T06:30:32.402109","exception":false,"start_time":"2025-05-22T06:30:32.400014","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport librosa\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom pathlib import Path\n\nclass CFG:\n    test_soundscapes = \"/kaggle/input/birdclef-2025/train_soundscapes\"\n    taxonomy_csv     = \"/kaggle/input/birdclef-2025/taxonomy.csv\"\n    model_dir        = \"/kaggle/input/bird-clef-2025-eficientnetv2-b0-02-02-10epoch\"\n    SR = 32000\n    WINDOW_SIZE = 5\n    N_FFT = 1024\n    HOP_LENGTH = 512\n    N_MELS = 128\n    FMIN = 50\n    FMAX = 14000\n    TARGET_SHAPE = (256, 256)\n    model_name = \"tf_efficientnetv2_b0.in1k\"\n    use_all_folds = False\n    in_channels = 1\n    device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    pseudo_label_threshold = 0.93 # for fallback in case dynamic threshold fails\n\ncfg = CFG()\n\ntaxonomy_df = pd.read_csv(cfg.taxonomy_csv)\nspecies_ids = taxonomy_df['primary_label'].tolist()\nnum_classes = len(species_ids)\nprint(f\"Loaded taxonomy with {num_classes} species classes.\")\nprint(f\"Using device: {cfg.device}\")\n","metadata":{"papermill":{"duration":5.271753,"end_time":"2025-05-22T06:30:37.675504","exception":false,"start_time":"2025-05-22T06:30:32.403751","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-06-03T11:35:33.530581Z","iopub.execute_input":"2025-06-03T11:35:33.530771Z","iopub.status.idle":"2025-06-03T11:35:39.951084Z","shell.execute_reply.started":"2025-06-03T11:35:33.530753Z","shell.execute_reply":"2025-06-03T11:35:39.950279Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Step 1: Count number of training samples per class\ntrain_df_original = pd.read_csv(\"/kaggle/input/birdclef-2025/train.csv\")\nlabel_counts = train_df_original['primary_label'].value_counts().to_dict()\n\n# Step 2: Set dynamic thresholds (less common → lower threshold)\nmin_thresh = 0.55\nmax_thresh = 0.99\n\nmax_count = max(label_counts.values())\nmin_count = min(label_counts.values())\n\nper_class_thresholds = {}\nfor label, count in label_counts.items():\n    commonness = (count - min_count) / (max_count - min_count + 1e-8)  # 0 = rare, 1 = common\n    per_class_thresholds[label] = min_thresh + (commonness ** 0.5) * (max_thresh - min_thresh)  # Apply square root\n\n# Sanity check\nprint(\"Dynamic threshold example:\")\nfor label in sorted(label_counts, key=label_counts.get)[:5]:  # 5 least frequent\n    print(f\"{label} (rare): {per_class_thresholds[label]:.3f}\")\nfor label in sorted(label_counts, key=label_counts.get, reverse=True)[:5]:  # 5 most frequent\n    print(f\"{label} (common): {per_class_thresholds[label]:.3f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-03T11:35:39.951954Z","iopub.execute_input":"2025-06-03T11:35:39.952270Z","iopub.status.idle":"2025-06-03T11:35:40.151410Z","shell.execute_reply.started":"2025-06-03T11:35:39.952225Z","shell.execute_reply":"2025-06-03T11:35:40.150565Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import timm\nfrom pathlib import Path\nimport torch\nimport torch.nn as nn\n\nclass BirdCLEFModel(nn.Module):\n    def __init__(self, cfg, num_classes):\n        super().__init__()\n        self.backbone = timm.create_model(\n            cfg.model_name,\n            pretrained=False,\n            in_chans=cfg.in_channels,\n            drop_rate=0.0,\n            drop_path_rate=0.0\n        )\n        \n        if 'efficientnet' in cfg.model_name:\n            backbone_out = self.backbone.classifier.in_features\n            self.backbone.classifier = nn.Identity()\n        elif 'resnet' in cfg.model_name:\n            backbone_out = self.backbone.fc.in_features\n            self.backbone.fc = nn.Identity()\n        else:\n            backbone_out = self.backbone.get_classifier().in_features\n            self.backbone.reset_classifier(0, '')  # use '' for newer timm models\n            \n        self.pooling = nn.AdaptiveAvgPool2d(1)\n        self.classifier = nn.Linear(backbone_out, num_classes)\n    \n    def forward(self, x):\n        features = self.backbone(x)\n        if isinstance(features, dict):\n            features = features.get('features', features.get('out', features))\n        if features.dim() == 4:\n            features = self.pooling(features)\n            features = features.view(features.size(0), -1)\n        logits = self.classifier(features)\n        return logits\n\n# Find model files\nmodel_files = list(Path(cfg.model_dir).rglob(\"*.pth\")) + list(Path(cfg.model_dir).rglob(\"*.pt\"))\nassert len(model_files) > 0, \"No model files found.\"\n\nmodels = []\nif cfg.use_all_folds:\n    # Load all model files (existing behavior)\n    for mfile in model_files:\n        print(f\"Loading: {mfile.name}\")\n        model = BirdCLEFModel(cfg, num_classes=num_classes)\n        \n        ckpt = torch.load(mfile, map_location=cfg.device, weights_only=False)\n        \n        if 'model_state_dict' in ckpt:\n            model.load_state_dict(ckpt['model_state_dict'])\n        else:\n            model.load_state_dict(ckpt)\n            \n        model.to(cfg.device)\n        model.eval()\n        models.append(model)\nelse:\n    # Load only the best model file\n    best_model_file = None\n    for mfile in model_files:\n        if 'best' in mfile.name.lower():  # Prioritize file with 'best' in name\n            best_model_file = mfile\n            break\n    if not best_model_file:\n        best_model_file = model_files[0]  # Fallback to first file if no 'best' found\n    print(f\"Loading best model: {best_model_file.name}\")\n    model = BirdCLEFModel(cfg, num_classes=num_classes)\n    \n    ckpt = torch.load(best_model_file, map_location=cfg.device, weights_only=False)\n    \n    if 'model_state_dict' in ckpt:\n        model.load_state_dict(ckpt['model_state_dict'])\n    else:\n        model.load_state_dict(ckpt)\n        \n    model.to(cfg.device)\n    model.eval()\n    models.append(model)","metadata":{"papermill":{"duration":15.246156,"end_time":"2025-05-22T06:30:52.923584","exception":false,"start_time":"2025-05-22T06:30:37.677428","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-06-03T11:35:40.152840Z","iopub.execute_input":"2025-06-03T11:35:40.153063Z","iopub.status.idle":"2025-06-03T11:35:49.112960Z","shell.execute_reply.started":"2025-06-03T11:35:40.153047Z","shell.execute_reply":"2025-06-03T11:35:49.112341Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import cv2\nfrom tqdm import tqdm\n\ndef audio_to_melspec(audio_data, cfg):\n    if np.isnan(audio_data).any():\n        mean_val = np.nanmean(audio_data)\n        audio_data = np.nan_to_num(audio_data, nan=mean_val)\n    mel_spec = librosa.feature.melspectrogram(\n        y=audio_data, sr=cfg.SR, n_fft=cfg.N_FFT,\n        hop_length=cfg.HOP_LENGTH, n_mels=cfg.N_MELS,\n        fmin=cfg.FMIN, fmax=cfg.FMAX, power=2.0\n    )\n    mel_spec_db = librosa.power_to_db(mel_spec, ref=np.max)\n    mel_spec_norm = (mel_spec_db - mel_spec_db.min()) / (mel_spec_db.max() - mel_spec_db.min() + 1e-8)\n    return mel_spec_norm\n\ndef process_audio_segment(audio_segment, cfg):\n    segment_length = cfg.SR * cfg.WINDOW_SIZE\n    if len(audio_segment) < segment_length:\n        audio_segment = np.pad(audio_segment, (0, segment_length - len(audio_segment)), mode='constant')\n    mel = audio_to_melspec(audio_segment, cfg)\n    if mel.shape != cfg.TARGET_SHAPE:\n        mel = cv2.resize(mel, cfg.TARGET_SHAPE, interpolation=cv2.INTER_LINEAR)\n    return mel.astype(np.float32)\n","metadata":{"papermill":{"duration":7350.583315,"end_time":"2025-05-22T08:33:23.509033","exception":false,"start_time":"2025-05-22T06:30:52.925718","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-06-03T11:35:49.113851Z","iopub.execute_input":"2025-06-03T11:35:49.114309Z","iopub.status.idle":"2025-06-03T11:35:49.389594Z","shell.execute_reply.started":"2025-06-03T11:35:49.114284Z","shell.execute_reply":"2025-06-03T11:35:49.388754Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from collections import defaultdict\n\n# temp storage for all predictions per class\nclass_preds = defaultdict(list)  # label → list of (prob, segment_name, mel)\n\npseudo_mels = {}\ntotal_segments = 0\naudio_files = sorted(Path(cfg.test_soundscapes).glob(\"*.ogg\"))\n\nprint(f\"Processing {len(audio_files)} files...\")\nwith torch.no_grad():\n    for audio_path in tqdm(audio_files, desc=\"Audio files\", dynamic_ncols=True):\n        audio_name = audio_path.name\n        y, sr = librosa.load(audio_path, sr=cfg.SR)\n        if sr != cfg.SR:\n            y = librosa.resample(y, orig_sr=sr, target_sr=cfg.SR)\n\n        seg_samples = cfg.SR * cfg.WINDOW_SIZE\n        n_segments = int(len(y) / seg_samples)\n        total_segments += n_segments\n\n        TOP_N = 3  # max per segment (still useful)\n        \n        for seg_idx in range(n_segments):\n            start = seg_idx * seg_samples\n            end = start + seg_samples\n            segment = y[start:end]\n            mel = process_audio_segment(segment, cfg)\n            tensor = torch.tensor(mel).unsqueeze(0).unsqueeze(0).to(cfg.device)\n        \n            if len(models) == 1:\n                probs = torch.sigmoid(models[0](tensor)).cpu().numpy().squeeze()\n            else:\n                preds = [torch.sigmoid(m(tensor)).cpu().numpy().squeeze() for m in models]\n                probs = np.mean(np.stack(preds, axis=0), axis=0)\n        \n            top_n_indices = np.argsort(probs)[-TOP_N:][::-1]\n        \n            for idx in top_n_indices:\n                prob = float(probs[idx])\n                label = species_ids[idx]\n                threshold = per_class_thresholds.get(label, cfg.pseudo_label_threshold)\n\n                if prob >= threshold:\n                    segment_name = f\"{audio_name.replace('.ogg', '')}_{seg_idx * cfg.WINDOW_SIZE}.ogg\"\n                    base = segment_name.replace('.ogg', '')\n                    class_preds[label].append((prob, segment_name, mel))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-03T11:35:49.390684Z","iopub.execute_input":"2025-06-03T11:35:49.391180Z","execution_failed":"2025-06-03T11:36:19.933Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Keep only top-K per class\nresults = []\nresults_with_conf = []  # for confidence analysis\nMAX_PER_CLASS = 400  # configurable cap\n\nfor label, entries in class_preds.items():\n    top_entries = sorted(entries, key=lambda x: -x[0])[:MAX_PER_CLASS]\n    for prob, segment_name, mel in top_entries:\n        results.append([segment_name, label, prob])\n        results_with_conf.append({\n            \"filename\": segment_name,\n            \"primary_label\": label,\n            \"confidence\": prob,\n            \"threshold\": per_class_thresholds.get(label, cfg.pseudo_label_threshold)\n        })\n        pseudo_mels[segment_name.replace('.ogg', '')] = mel\n\nprint(f\"Done. Collected {len(results)} pseudo-labels from {total_segments} segments.\")\n\n# DEBUG: show that our keys line up with the filenames in results\nprint(\"First 10 filenames from results (with .ogg):\")\nprint([row[0] for row in results[:10]])\nprint(\"First 10 keys in pseudo_mels dict:\")\nprint(list(pseudo_mels.keys())[:10])\n\nprint(\"\\nCheck each of the first 10:\")\nfor fn, lbl, prob in results[:10]:\n    k = fn.replace(\".ogg\",\"\")\n    print(f\"{fn}  → key='{k}'  in dict? {k in pseudo_mels}\")\n\nnp.save(\"pseudo_mels.npy\", pseudo_mels)\n\n# Save confidence info separately for analysis\nconf_df = pd.DataFrame(results_with_conf)\nconf_df.to_csv(\"pseudo_label_confidences.csv\", index=False)\nprint(f\"Saved {len(conf_df)} entries with confidence scores to 'pseudo_label_confidences.csv'\")\n\n# Save train.csv-compatible format\npseudo_train = pd.DataFrame({\n    \"filename\": [row[0] for row in results],\n    \"primary_label\": [row[1] for row in results],\n    \"secondary_labels\": [[] for _ in results],\n    \"latitude\": [None] * len(results),\n    \"longitude\": [None] * len(results),\n    \"author\": [\"pseudo\"] * len(results),\n    \"rating\": [0] * len(results),\n    \"collection\": [\"pseudo\"] * len(results)\n})\n\npseudo_train.to_csv(\"pseudo_train.csv\", index=False)\nprint(f\"Saved {len(pseudo_train)} pseudo-labeled samples to 'pseudo_train.csv'\")\npseudo_train.head()\n","metadata":{"papermill":{"duration":0.413836,"end_time":"2025-05-22T08:33:24.295015","exception":false,"start_time":"2025-05-22T08:33:23.881179","status":"completed"},"tags":[],"trusted":true,"execution":{"execution_failed":"2025-06-03T11:36:19.934Z"}},"outputs":[],"execution_count":null}]}