{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","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,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":12688641,"sourceType":"datasetVersion","datasetId":8018520},{"sourceId":509081,"sourceType":"modelInstanceVersion","isSourceIdPinned":false,"modelInstanceId":403798,"modelId":421720},{"sourceId":511848,"sourceType":"modelInstanceVersion","isSourceIdPinned":false,"modelInstanceId":405209,"modelId":421951}],"dockerImageVersionId":30918,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# **BirdCLEF 2025 Inference Notebook**\nThis notebook runs inference on BirdCLEF 2025 test soundscapes and generates a submission file. You can find the pre-processing and training processes in the following notebooks:\n\n- [BirdCLEF'25 | Transfer learning BEATs](https://www.kaggle.com/code/hubfor/birdclef-25-transfer-learning-beats)  \n","metadata":{}},{"cell_type":"code","source":"import os\nimport gc\nimport warnings\nimport logging\nimport time\nimport math\nimport cv2\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\nimport librosa\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport timm\nfrom tqdm.auto import tqdm\n\nwarnings.filterwarnings(\"ignore\")\nlogging.basicConfig(level=logging.ERROR)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-07T21:26:40.012818Z","iopub.execute_input":"2025-08-07T21:26:40.013196Z","iopub.status.idle":"2025-08-07T21:26:40.021528Z","shell.execute_reply.started":"2025-08-07T21:26:40.013167Z","shell.execute_reply":"2025-08-07T21:26:40.019229Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CFG:\n \n    test_soundscapes = '/kaggle/input/birdclef-2025/test_soundscapes'\n    submission_csv = '/kaggle/input/birdclef-2025/sample_submission.csv'\n    taxonomy_csv = '/kaggle/input/birdclef-2025/taxonomy.csv'\n    model_path = '/kaggle/input/finetuned_beats_birdclef/pytorch/2/1/model_fold3.pth'  \n    models_code_path = '/kaggle/input/microsoft-beats-model'\n    models_weights_path = '/kaggle/input/microsoft-pretrained-beats-iter3-plus-as2m/pytorch/default/1/BEATs_iter3_plus_AS2M.pt'\n    # Audio parameters\n    FS = 32000  \n    WINDOW_SIZE = 5  \n    \n    # Mel spectrogram parameters\n    N_FFT = 1024\n    HOP_LENGTH = 512\n    N_MELS = 128\n    FMIN = 50\n    FMAX = 14000\n    TARGET_SHAPE = (256, 256)\n    \n    device = 'cuda' if torch.cuda.is_available() else 'cpu'\n    \n    # Inference parameters\n    batch_size = 16\n    use_tta = False  \n    tta_count = 3   \n    threshold = 0.5\n    \n    use_specific_folds = False  # If False, use all found models\n    folds = [0, 1]  # Used only if use_specific_folds is True\n    \n    debug = False\n    debug_count = 3\n\ncfg = CFG()\ncfg.num_classes = 206","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-07T21:26:40.022793Z","iopub.execute_input":"2025-08-07T21:26:40.023236Z","iopub.status.idle":"2025-08-07T21:26:40.081553Z","shell.execute_reply.started":"2025-08-07T21:26:40.023193Z","shell.execute_reply":"2025-08-07T21:26:40.080218Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(f\"Using device: {cfg.device}\")\nprint(f\"Loading taxonomy data...\")\ntaxonomy_df = pd.read_csv(cfg.taxonomy_csv)\nspecies_ids = taxonomy_df['primary_label'].tolist()\nnum_classes = len(species_ids)\nprint(f\"Number of classes: {num_classes}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-07T21:26:40.084636Z","iopub.execute_input":"2025-08-07T21:26:40.085250Z","iopub.status.idle":"2025-08-07T21:26:40.134220Z","shell.execute_reply.started":"2025-08-07T21:26:40.085214Z","shell.execute_reply":"2025-08-07T21:26:40.132775Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import sys\nsys.path.append(cfg.models_code_path)\n\nfrom BEATs import BEATs, BEATsConfig\n\nclass BEATs_model_classifier(torch.nn.Module):\n    \n    def __init__(self, beats_model_pretrained: BEATs, cfg, hidden_dim: int = 512):\n\n        super().__init__()\n\n        self.cfg = cfg\n        \n        self.beats = beats_model_pretrained\n        beats_dim = beats_model_pretrained.cfg.encoder_embed_dim\n        \n        # self.multiheadAttentions = nn.MultiheadAttention(\n        #     beats_dim, num_heads=8, dropout=0.5, batch_first=True\n        # )\n\n        self.pooling = nn.AdaptiveAvgPool1d(1)\n        \n        self.classifier = nn.Sequential(\n            nn.LayerNorm(beats_dim),\n            nn.Linear(beats_dim, hidden_dim),\n            nn.GELU(),\n            nn.Dropout(0.4),\n            nn.Linear(hidden_dim, hidden_dim // 2),\n            nn.GELU(),\n            nn.Dropout(0.3),\n            nn.Linear(hidden_dim // 2, cfg.num_classes) \n        )\n        \n    \n        for param in self.beats.parameters():\n            param.requires_grad = False\n    \n        \n    def forward(self, waveforms, padding_mask=None):\n\n        beats_features, _ = self.beats.extract_features(waveforms, padding_mask=padding_mask)\n        \n        # attended, _ = self.multiheadAttentions(\n        #     beats_features, beats_features, beats_features\n        # )\n        \n        # # attended: [B, T, D] -> [B, D, T] -> [B, D, 1] -> [B, D]\n        # pooled = self.pooling(attended.transpose(1, 2)).squeeze(-1)\n        #pooled = self.pooling(beats_features).squeeze(-1)\n        beats_features = beats_features.mean(dim=1)\n        #print(beats_features.shape)\n        logits = self.classifier(beats_features)\n        # # Classification\n        #logits = self.classifier(pooled)\n        #logits = self.classifier(beats_features)\n        \n        return logits#.squeeze(-1)  # [B] for binary classification\n\n\n\ndef load_model(cfg):\n    # load the pre-trained checkpoints\n    checkpoint = torch.load(cfg.models_weights_path)\n    \n    cfg_model = BEATsConfig(checkpoint['cfg'])\n    BEATs_model = BEATs(cfg_model)\n    BEATs_model.load_state_dict(checkpoint['model'])\n\n    classifier = BEATs_model_classifier(BEATs_model, cfg)\n\n    model_state = torch.load(cfg.model_path, map_location=torch.device('cpu'))\n    classifier.load_state_dict(model_state['model_state_dict'])\n    classifier.eval()\n    return classifier","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-07T21:28:37.219641Z","iopub.execute_input":"2025-08-07T21:28:37.220070Z","iopub.status.idle":"2025-08-07T21:28:37.230094Z","shell.execute_reply.started":"2025-08-07T21:28:37.220038Z","shell.execute_reply":"2025-08-07T21:28:37.228678Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def predict(audio_path, model, cfg, species_ids):\n    \"\"\"Process a single audio file and predict species presence for each 5-second segment using dataset approach\"\"\"\n    predictions = []\n    row_ids = []\n    soundscape_id = Path(audio_path).stem\n    \n    try:\n        print(f\"Processing {soundscape_id}\")\n        \n        # Load audio using torchaudio (same as dataset)\n        waveform, sample_rate = torchaudio.load(audio_path)\n        \n        # Resample to 16kHz if needed (same as dataset)\n        if sample_rate != 16000:\n            resampler = torchaudio.transforms.Resample(sample_rate, 16000)\n            waveform = resampler(waveform)\n        \n        # Convert to mono (same as dataset)\n        if waveform.shape[0] > 1:\n            waveform = waveform.mean(dim=0, keepdim=True)\n        \n        # Calculate segment parameters\n        segment_length = int(16000 * cfg.WINDOW_SIZE)  # 5 seconds at 16kHz\n        total_length = waveform.shape[1]\n        total_segments = int(total_length / segment_length)\n        \n        for segment_idx in range(total_segments):\n            start_sample = segment_idx * segment_length\n            end_sample = start_sample + segment_length\n            \n            # Extract segment\n            segment_waveform = waveform[:, start_sample:end_sample]\n            \n            # Handle padding/cropping (same logic as dataset)\n            current_length = segment_waveform.shape[1]\n            if current_length < segment_length:\n                # Pad with zeros\n                padding = segment_length - current_length\n                segment_waveform = F.pad(segment_waveform, (0, padding))\n                padding_mask = torch.zeros(segment_length, dtype=torch.bool)\n                padding_mask[current_length:] = True\n            elif current_length > segment_length:\n                # Crop to exact length\n                segment_waveform = segment_waveform[:, :segment_length]\n                padding_mask = torch.zeros(segment_length, dtype=torch.bool)\n            else:\n                # Exact length\n                padding_mask = torch.zeros(segment_length, dtype=torch.bool)\n            \n            # Remove channel dimension (same as dataset)\n            segment_waveform = segment_waveform.squeeze(0)\n            \n            # Create row_id\n            end_time_sec = (segment_idx + 1) * cfg.WINDOW_SIZE\n            row_id = f\"{soundscape_id}_{end_time_sec}\"\n            row_ids.append(row_id)\n            \n            # Prepare input for model (add batch dimension)\n            segment_waveform = segment_waveform.unsqueeze(0).to(cfg.device)\n            padding_mask = padding_mask.unsqueeze(0).to(cfg.device)\n            \n\n            # Get predictions\n            with torch.no_grad():\n                outputs = model(segment_waveform, padding_mask=padding_mask)\n                final_preds = torch.sigmoid(outputs).cpu().numpy().squeeze()\n            \n            predictions.append(final_preds)\n            \n    except Exception as e:\n        print(f\"Error processing {audio_path}: {e}\")\n    \n    return row_ids, predictions","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-07T21:26:41.284177Z","iopub.execute_input":"2025-08-07T21:26:41.284504Z","iopub.status.idle":"2025-08-07T21:26:41.296768Z","shell.execute_reply.started":"2025-08-07T21:26:41.284475Z","shell.execute_reply":"2025-08-07T21:26:41.295237Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torchaudio\ndef run_inference(cfg, model, species_ids):\n    \"\"\"Run inference on all test soundscapes\"\"\"\n    test_files = list(Path(cfg.test_soundscapes).glob('*.ogg'))\n    #test_files = list(Path(\"/kaggle/input/birdclef-2025/train_soundscapes\").glob('*.ogg'))\n        \n    if cfg.debug:\n        print(f\"Debug mode enabled, using only {cfg.debug_count} files\")\n        test_files = test_files[:cfg.debug_count]\n    \n    print(f\"Found {len(test_files)} test soundscapes\")\n    \n    all_row_ids = []\n    all_predictions = []\n    \n    for audio_path in tqdm(test_files):\n        row_ids, predictions = predict(str(audio_path), model, cfg, species_ids)\n        all_row_ids.extend(row_ids)\n        all_predictions.extend(predictions)\n    \n    return all_row_ids, all_predictions","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-07T21:26:41.297895Z","iopub.execute_input":"2025-08-07T21:26:41.298207Z","iopub.status.idle":"2025-08-07T21:26:41.328056Z","shell.execute_reply.started":"2025-08-07T21:26:41.298180Z","shell.execute_reply":"2025-08-07T21:26:41.326729Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def create_submission(row_ids, predictions, species_ids, cfg):\n    \"\"\"Create submission dataframe\"\"\"\n    print(\"Creating submission dataframe...\")\n\n    submission_dict = {'row_id': row_ids}\n    \n    for i, species in enumerate(species_ids):\n        submission_dict[species] = [pred[i] for pred in predictions]\n\n    submission_df = pd.DataFrame(submission_dict)\n\n    submission_df.set_index('row_id', inplace=True)\n\n    sample_sub = pd.read_csv(cfg.submission_csv, index_col='row_id')\n\n    missing_cols = set(sample_sub.columns) - set(submission_df.columns)\n    if missing_cols:\n        print(f\"Warning: Missing {len(missing_cols)} species columns in submission\")\n        for col in missing_cols:\n            submission_df[col] = 0.0\n\n    submission_df = submission_df[sample_sub.columns]\n\n    submission_df = submission_df.reset_index()\n    \n    return submission_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-07T21:26:41.329504Z","iopub.execute_input":"2025-08-07T21:26:41.329922Z","iopub.status.idle":"2025-08-07T21:26:41.355035Z","shell.execute_reply.started":"2025-08-07T21:26:41.329883Z","shell.execute_reply":"2025-08-07T21:26:41.353752Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def main():\n    start_time = time.time()\n    print(\"Starting BirdCLEF-2025 inference...\")\n\n    model = load_model(cfg).to(cfg.device)\n    \n    if not model:\n        print(\"No models found! Please check model paths.\")\n        return\n\n    row_ids, predictions = run_inference(cfg, model, species_ids)\n\n    submission_df = create_submission(row_ids, predictions, species_ids, cfg)\n\n    submission_path = 'submission.csv'\n    submission_df.to_csv(submission_path, index=False)\n    print(f\"Submission saved to {submission_path}\")\n    \n    end_time = time.time()\n    print(f\"Inference completed in {(end_time - start_time)/60:.2f} minutes\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-07T21:26:41.356158Z","iopub.execute_input":"2025-08-07T21:26:41.356505Z","iopub.status.idle":"2025-08-07T21:26:41.386428Z","shell.execute_reply.started":"2025-08-07T21:26:41.356476Z","shell.execute_reply":"2025-08-07T21:26:41.385293Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-07T21:28:41.250374Z","iopub.execute_input":"2025-08-07T21:28:41.250900Z","iopub.status.idle":"2025-08-07T21:28:45.245306Z","shell.execute_reply.started":"2025-08-07T21:28:41.250855Z","shell.execute_reply":"2025-08-07T21:28:45.243915Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}