{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceType":"competition","sourceId":70203,"databundleVersionId":8068726,"isSourceIdPinned":false},{"sourceType":"kernelVersion","sourceId":314404207,"isSourceIdPinned":false}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Day 2: Zero-Shot Evaluation\nThis notebook loads the embeddings and inferences extracted from `02_extract_embeddings` and computes the validation AUC. It calculates both the BirdNET zero-shot baseline and the Perch Prototypical Head baseline.","metadata":{}},{"cell_type":"code","source":"!pip install -q scikit-learn pandas numpy pyarrow","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-26T05:07:23.435340Z","iopub.execute_input":"2026-04-26T05:07:23.435646Z","iopub.status.idle":"2026-04-26T05:07:28.360467Z","shell.execute_reply.started":"2026-04-26T05:07:23.435612Z","shell.execute_reply":"2026-04-26T05:07:28.359519Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport ast\nfrom pathlib import Path\nfrom sklearn.metrics import roc_auc_score\nfrom sklearn.model_selection import train_test_split\n\nINPUT_DIR = Path('/kaggle/input')\nDATA_OUT = Path('./data')\nDATA_OUT.mkdir(exist_ok=True)\nREPORTS_OUT = Path('./reports')\nREPORTS_OUT.mkdir(exist_ok=True)\n\nprint(\"Searching for parquet files...\")\nperch_files = list(INPUT_DIR.rglob('perch_train_*.parquet'))\nbirdnet_files = list(INPUT_DIR.rglob('birdnet_train_*.parquet'))\n\nif not perch_files or not birdnet_files:\n    print(\"WARNING: Could not find extracted Parquet files. Did you attach the output of your Extraction Notebook as Data?\")\nelse:\n    print(f\"Found {len(perch_files)} Perch parquets and {len(birdnet_files)} BirdNET parquets.\")\n    \n    # Load all parquets into dataframes\n    perch_df = pd.concat([pd.read_parquet(f) for f in perch_files]).reset_index(drop=True)\n    birdnet_df = pd.concat([pd.read_parquet(f) for f in birdnet_files]).reset_index(drop=True)\n    print(f\"Loaded {len(perch_df)} Perch embeddings and {len(birdnet_df)} BirdNET predictions.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-26T05:07:28.362400Z","iopub.execute_input":"2026-04-26T05:07:28.362753Z","iopub.status.idle":"2026-04-26T05:09:10.522735Z","shell.execute_reply.started":"2026-04-26T05:07:28.362722Z","shell.execute_reply":"2026-04-26T05:09:10.521785Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if 'perch_df' in locals() and len(perch_df) > 0:\n    # 1. Train / Validation Split (80/20) based on species\n    print(\"Creating Train/Validation split...\")\n    focus_species = perch_df['species_code'].unique().tolist()\n    print(f\"Found {len(focus_species)} unique species.\")\n    \n    # We want at least 2 samples per class to stratify, let's filter classes with < 2 samples\n    counts = perch_df['species_code'].value_counts()\n    valid_species = counts[counts >= 2].index.tolist()\n    valid_perch_df = perch_df[perch_df['species_code'].isin(valid_species)].copy()\n    \n    train_df, val_df = train_test_split(\n        valid_perch_df, \n        test_size=0.2, \n        stratify=valid_perch_df['species_code'],\n        random_state=42\n    )\n    print(f\"Split: {len(train_df)} train, {len(val_df)} validation clips.\")\n    val_file_ids = set(val_df['file_id'])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-26T05:09:10.523719Z","iopub.execute_input":"2026-04-26T05:09:10.524275Z","iopub.status.idle":"2026-04-26T05:09:10.565066Z","shell.execute_reply.started":"2026-04-26T05:09:10.524245Z","shell.execute_reply":"2026-04-26T05:09:10.564095Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if 'perch_df' in locals() and len(perch_df) > 0:\n    # 2. Perch Prototypical Head Evaluation\n    print(\"\\n--- Computing Perch Prototype Baseline ---\")\n    prototypes = {}\n    for sp in valid_species:\n        sp_embs = np.stack(train_df[train_df['species_code'] == sp]['emb'].values)\n        prototypes[sp] = np.mean(sp_embs, axis=0)\n        \n    P_matrix = np.stack([prototypes[sp] for sp in valid_species]) # Shape: (N_species, 1536)\n    \n    y_true_perch = []\n    y_pred_perch = []\n    \n    for idx, row in val_df.iterrows():\n        emb = row['emb']\n        # Cosine similarity\n        sims = np.dot(P_matrix, emb) / (np.linalg.norm(P_matrix, axis=1) * np.linalg.norm(emb) + 1e-12)\n        # Softmax with temperature 0.05\n        sims_scaled = sims / 0.05\n        probs = np.exp(sims_scaled - np.max(sims_scaled)) / np.sum(np.exp(sims_scaled - np.max(sims_scaled)))\n        \n        y_t = np.zeros(len(valid_species))\n        y_t[valid_species.index(row['species_code'])] = 1.0\n        \n        y_true_perch.append(y_t)\n        y_pred_perch.append(probs)\n        \n    y_true_perch = np.array(y_true_perch)\n    y_pred_perch = np.array(y_pred_perch)\n    \n    perch_auc_dict = {}\n    for i, sp in enumerate(valid_species):\n        if np.sum(y_true_perch[:, i]) > 0:\n            perch_auc_dict[sp] = roc_auc_score(y_true_perch[:, i], y_pred_perch[:, i])\n            \n    perch_macro_auc = np.mean(list(perch_auc_dict.values()))\n    print(f\"Perch Prototype Macro Validation AUC: {perch_macro_auc:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-26T05:09:10.566912Z","iopub.execute_input":"2026-04-26T05:09:10.567439Z","iopub.status.idle":"2026-04-26T05:09:12.058156Z","shell.execute_reply.started":"2026-04-26T05:09:10.567411Z","shell.execute_reply":"2026-04-26T05:09:12.057434Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if 'birdnet_df' in locals() and len(birdnet_df) > 0:\n    # 3. BirdNET Zero-Shot Evaluation\n    print(\"\\n--- Computing BirdNET Zero-Shot Baseline ---\")\n    \n    # We need to map scientific names to species codes. Let's try to find the taxonomy file.\n    taxa_csv = next(INPUT_DIR.rglob('eBird_Taxonomy_v2021.csv'), None)\n    taxa_mapping = {}\n    if taxa_csv:\n        taxa_df = pd.read_csv(taxa_csv)\n        taxa_mapping = dict(zip(taxa_df['SCI_NAME'].str.lower(), taxa_df['SPECIES_CODE']))\n    else:\n        print(\"WARNING: eBird_Taxonomy_v2021.csv not found! BirdNET predictions might not map correctly without it.\")\n        \n    # Isolate the validation set for BirdNET\n    birdnet_val_df = birdnet_df[birdnet_df['file_id'].isin(val_file_ids)].copy()\n    \n    y_true_birdnet = []\n    y_pred_birdnet = []\n    \n    for idx, row in birdnet_val_df.iterrows():\n        preds = np.zeros(len(valid_species))\n        \n        y_t = np.zeros(len(valid_species))\n        if row['species_code'] in valid_species:\n            y_t[valid_species.index(row['species_code'])] = 1.0\n        \n        try:\n            dets = ast.literal_eval(row['detections'])\n            for d in dets:\n                sci_name = d.get('scientific_name', '').lower()\n                conf = d.get('confidence', 0.0)\n                sp_code = taxa_mapping.get(sci_name, None)\n                if sp_code in valid_species:\n                    s_idx = valid_species.index(sp_code)\n                    preds[s_idx] = max(preds[s_idx], conf)\n        except:\n            pass # Failed to parse detections string\n            \n        y_true_birdnet.append(y_t)\n        y_pred_birdnet.append(preds)\n        \n    y_true_birdnet = np.array(y_true_birdnet)\n    y_pred_birdnet = np.array(y_pred_birdnet)\n    \n    birdnet_auc_dict = {}\n    for i, sp in enumerate(valid_species):\n        if np.sum(y_true_birdnet[:, i]) > 0:\n            # Only calculate if there's at least one positive example and one negative\n            if len(np.unique(y_true_birdnet[:, i])) > 1:\n                birdnet_auc_dict[sp] = roc_auc_score(y_true_birdnet[:, i], y_pred_birdnet[:, i])\n            \n    if birdnet_auc_dict:\n        birdnet_macro_auc = np.mean(list(birdnet_auc_dict.values()))\n        print(f\"BirdNET Zero-Shot Macro Validation AUC: {birdnet_macro_auc:.4f}\")\n    else:\n        birdnet_macro_auc = np.nan\n        print(\"BirdNET Zero-Shot Macro Validation AUC: N/A (Mapping failed or missing data)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-26T05:09:12.059258Z","iopub.execute_input":"2026-04-26T05:09:12.059711Z","iopub.status.idle":"2026-04-26T05:09:14.940151Z","shell.execute_reply.started":"2026-04-26T05:09:12.059680Z","shell.execute_reply":"2026-04-26T05:09:14.939325Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if 'perch_df' in locals() and len(perch_df) > 0:\n    # 4. Save Results to CSV\n    results_data = []\n    for sp in valid_species:\n        results_data.append({\n            'species': sp,\n            'perch_prototype_auc': perch_auc_dict.get(sp, np.nan),\n            'birdnet_zeroshot_auc': birdnet_auc_dict.get(sp, np.nan) if 'birdnet_auc_dict' in locals() else np.nan\n        })\n        \n    results_df = pd.DataFrame(results_data)\n    results_df.to_csv(REPORTS_OUT / 'per_species_zero_shot_auc.csv', index=False)\n    \n    # Summary Table\n    summary_df = pd.DataFrame([{\n        'system': 'Perch Prototype (Mean Emb)',\n        'macro_auc_val': perch_macro_auc\n    }, {\n        'system': 'BirdNET Zero-Shot',\n        'macro_auc_val': birdnet_macro_auc if 'birdnet_macro_auc' in locals() else np.nan\n    }])\n    summary_df.to_csv(REPORTS_OUT / 'main_results_day2.csv', index=False)\n    \n    print(\"\\n--- Zero-Shot Evaluation Complete ---\")\n    print(\"Results saved to reports/per_species_zero_shot_auc.csv and reports/main_results_day2.csv\")\n    display(summary_df)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-26T05:09:14.941086Z","iopub.execute_input":"2026-04-26T05:09:14.941773Z","iopub.status.idle":"2026-04-26T05:09:14.975302Z","shell.execute_reply.started":"2026-04-26T05:09:14.941741Z","shell.execute_reply":"2026-04-26T05:09:14.974667Z"}},"outputs":[],"execution_count":null}]}