{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":91844,"databundleVersionId":11361821,"sourceType":"competition"},{"sourceId":424730,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":346179,"modelId":367451}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"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)\n\n#logging.basicConfig(level=logging.INFO)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#REGNET MODEL\n\"\"\"\n# BirdCLEF 2025 - Fixed RegNet Inference Code\n\nThis notebook loads a trained RegNet model and creates predictions for the BirdCLEF 2025 test data.\nIt fixes the previous issues with the submission format to ensure correct evaluation.\n\"\"\"\n# Configuration class\nclass CFG_regnet:\n    # Paths\n    test_soundscapes = '/kaggle/input/birdclef-2025/train_audio'\n    sample_submission = '/kaggle/input/birdclef-2025/sample_submission.csv'\n    taxonomy_csv = '/kaggle/input/birdclef-2025/taxonomy.csv'\n    model_path = '/kaggle/input/regnetmodel2/pytorch/default/1/best_model.pth'\n\n    \n    # Audio parameters\n    sample_rate = 32000\n    duration = 5  # seconds\n    \n    # Mel spectrogram parameters - must match training settings\n    n_fft = 1024\n    hop_length = 512\n    n_mels = 64\n    fmin = 50\n    fmax = 14000\n    \n    # Image parameters\n    img_size = 224\n    \n    # Model parameters\n    model_name = 'regnety_008'  # Must match training model\n    \n    # Inference parameters\n    batch_size = 32\n    threshold = 0.5  # Confidence threshold for positive predictions\n    \n    # Test-time augmentation\n    use_tta = True  # Enable TTA for better performance\n    tta_steps = 3\n    \n    # Debug\n    debug_mode = False    # Set to True to process only a few soundscapes\n    debug_count = 3       # Number of soundscapes to process in debug mode\n    \n    # Device\n    device = \"cpu\"\n\n# RegNet model definition - must match training model architecture\nclass BirdCLEFModel(nn.Module):\n    def __init__(self, model_name, num_classes, in_channels=1):\n        super().__init__()\n        \n        # Load the RegNet model\n        self.backbone = timm.create_model(\n            model_name,\n            pretrained=False,  # Not using pretrained weights for inference\n            in_chans=in_channels,\n            num_classes=0      # Remove classifier head\n        )\n        \n        # Get feature dimension automatically\n        with torch.no_grad():\n            dummy_input = torch.zeros(1, in_channels, cfg_regnet.img_size, cfg_regnet.img_size)\n            features = self.backbone(dummy_input)\n            feature_dim = features.shape[1]\n        \n        # Create classifier head\n        self.classifier = nn.Sequential(\n            nn.Dropout(0.2),\n            nn.Linear(feature_dim, num_classes)\n        )\n        \n    def forward(self, x):\n        features = self.backbone(x)\n        output = self.classifier(features)\n        return output\n\n# Audio processing functions - identical to training\ndef audio_to_melspec(audio, cfg_regnet):\n    \"\"\"Convert audio data to mel spectrogram\"\"\"\n    # Handle NaN values\n    if np.isnan(audio).any():\n        audio = np.nan_to_num(audio)\n    \n    # Generate mel spectrogram\n    mel_spec = librosa.feature.melspectrogram(\n        y=audio,\n        sr=cfg_regnet.sample_rate,\n        n_fft=cfg_regnet.n_fft,\n        hop_length=cfg_regnet.hop_length,\n        n_mels=cfg_regnet.n_mels,\n        fmin=cfg_regnet.fmin,\n        fmax=cfg_regnet.fmax,\n        power=2.0\n    )\n    \n    # Convert to dB scale\n    mel_spec_db = librosa.power_to_db(mel_spec, ref=np.max)\n    \n    # Normalize\n    mel_spec_norm = (mel_spec_db + 80) / 80  # Typical dB range\n    \n    return np.clip(mel_spec_norm, 0, 1)  # Clip to [0, 1]\n\ndef apply_tta(mel_spec, step):\n    \"\"\"Apply test-time augmentation\"\"\"\n    if step == 0:\n        # Original spectrogram\n        return mel_spec\n    elif step == 1:\n        # Horizontal flip (time axis)\n        return np.flip(mel_spec, axis=1)\n    elif step == 2:\n        # Vertical flip (frequency axis)\n        return np.flip(mel_spec, axis=0)\n    else:\n        return mel_spec\n\ndef load_model_and_species():\n    \"\"\"Load trained model from checkpoint and get species list\"\"\"\n    print(f\"Loading model from {cfg_regnet.model_path}\")\n    \n    # First load the sample submission to get expected column names\n    sample_sub = pd.read_csv(cfg_regnet.sample_submission)\n    species_columns = [col for col in sample_sub.columns if col != 'row_id']\n    num_species = len(species_columns)\n    print(f\"Sample submission has {num_species} species columns\")\n    \n    # Load checkpoint\n    checkpoint = torch.load(cfg_regnet.model_path, map_location=cfg_regnet.device)\n    \n    # Create model with same number of outputs as species\n    model = BirdCLEFModel(cfg_regnet.model_name, num_species, in_channels=1)\n    \n    # Try to load model state\n    try:\n        if 'model_state_dict' in checkpoint:\n            model.load_state_dict(checkpoint['model_state_dict'])\n            print(\"Loaded model state from 'model_state_dict'\")\n        else:\n            for key in ['state_dict', 'model']:\n                if key in checkpoint:\n                    model.load_state_dict(checkpoint[key])\n                    print(f\"Loaded model state from '{key}'\")\n                    break\n    except Exception as e:\n        print(f\"WARNING: Error loading model weights: {e}\")\n    \n    model = model.to(cfg_regnet.device)\n    model.eval()\n    \n    # Create direct mapping from model outputs to species column names\n    species_map = {i: species for i, species in enumerate(species_columns)}\n    \n    return model, species_map, species_columns\n\ndef predict_on_soundscape(model, audio_path, species_map):\n    \"\"\"Process a soundscape file and generate predictions for each 5-second segment\"\"\"\n    soundscape_id = Path(audio_path).stem\n    \n    try:\n        # Load audio\n        audio_data, _ = librosa.load(\n            audio_path, \n            sr=cfg_regnet.sample_rate,\n            res_type='kaiser_fast'  # Faster resampling\n        )\n        \n        # Calculate total segments\n        segment_samples = cfg_regnet.sample_rate * cfg_regnet.duration\n        total_segments = int(len(audio_data) / segment_samples)\n        \n        # Initialize lists for results\n        row_ids = []\n        all_predictions = []\n        \n        # Process each segment\n        for segment_idx in range(total_segments):\n            # Extract segment\n            start_sample = segment_idx * segment_samples\n            end_sample = start_sample + segment_samples\n            segment_audio = audio_data[start_sample:end_sample]\n            \n            # Create row ID in required format\n            end_time_sec = (segment_idx + 1) * cfg_regnet.duration\n            row_id = f\"{soundscape_id}_{end_time_sec}\"\n            row_ids.append(row_id)\n            \n            if cfg_regnet.use_tta:\n                # Apply test-time augmentation\n                segment_preds = []\n                \n                for tta_step in range(cfg_regnet.tta_steps):\n                    # Process audio to mel spectrogram\n                    mel_spec = audio_to_melspec(segment_audio, cfg_regnet)\n                    mel_spec = apply_tta(mel_spec, tta_step)\n                    \n                    # Resize\n                    mel_spec = cv2.resize(mel_spec, (cfg_regnet.img_size, cfg_regnet.img_size))\n                    \n                    # Convert to tensor\n                    mel_spec = torch.tensor(mel_spec, dtype=torch.float32).unsqueeze(0).unsqueeze(0)\n                    mel_spec = mel_spec.to(cfg_regnet.device)\n                    \n                    # Get predictions\n                    with torch.no_grad():\n                        outputs = model(mel_spec)\n                        probs = torch.sigmoid(outputs).cpu().numpy().squeeze()\n                        segment_preds.append(probs)\n                \n                # Average predictions from all TTA steps\n                final_preds = np.mean(segment_preds, axis=0)\n            else:\n                # Process audio without TTA\n                mel_spec = audio_to_melspec(segment_audio, cfg_regnet)\n                \n                # Resize\n                mel_spec = cv2.resize(mel_spec, (cfg_regnet.img_size, cfg_regnet.img_size))\n                \n                # Convert to tensor\n                mel_spec = torch.tensor(mel_spec, dtype=torch.float32).unsqueeze(0).unsqueeze(0)\n                mel_spec = mel_spec.to(cfg_regnet.device)\n                \n                # Get predictions\n                with torch.no_grad():\n                    outputs = model(mel_spec)\n                    final_preds = torch.sigmoid(outputs).cpu().numpy().squeeze()\n            \n            # Create a prediction dictionary\n            pred_dict = {}\n            \n            # Handle case where final_preds might be a single value\n            if np.isscalar(final_preds):\n                # If model only has one output class\n                if len(species_map) > 0:\n                    species_id = list(species_map.values())[0]\n                    pred_dict[species_id] = float(final_preds)\n            else:\n                # Map model outputs to species IDs\n                for i, prob in enumerate(final_preds):\n                    if i in species_map:\n                        species_id = species_map[i]\n                        pred_dict[species_id] = float(prob)\n            \n            all_predictions.append(pred_dict)\n        \n        return row_ids, all_predictions\n    \n    except Exception as e:\n        print(f\"Error processing {audio_path}: {e}\")\n        import traceback\n        traceback.print_exc()\n        return [], []\n\ndef create_submission_file(all_row_ids, all_predictions):\n    \"\"\"Create submission file in the required format\"\"\"\n    print(\"Creating submission file...\")\n    \n    # Load sample submission\n    sample_sub = pd.read_csv(cfg_regnet.sample_submission)\n    \n    # Initialize submission with row_ids\n    submission_df = pd.DataFrame({'row_id': all_row_ids})\n    \n    # Add each species column from sample submission\n    for col in sample_sub.columns:\n        if col != 'row_id':\n            # Default all values to 0.0\n            submission_df[col] = 0.0\n            \n            # Update with actual predictions if available\n            for i, pred_dict in enumerate(all_predictions):\n                if col in pred_dict and i < len(submission_df):\n                    submission_df.loc[i, col] = pred_dict[col]\n    \n    # Ensure column order matches sample submission exactly\n    submission_df = submission_df[sample_sub.columns]\n    \n    # Save to CSV\n    submission_df.to_csv('submission_regnet.csv', index=False, float_format='%.6f')\n    print(f\"Submission file created with {len(submission_df)} predictions\")\n\ndef run_inference():\n    \"\"\"Main inference function\"\"\"\n    start_time = time.time()\n    print(f\"Starting inference using device: {cfg_regnet.device}\")\n    \n    # Load model with species mapping\n    model, species_map, _ = load_model_and_species()\n    \n    # Find test files\n    test_files = list(Path(cfg_regnet.test_soundscapes).glob('*/*.ogg'))\n    \n    if cfg_regnet.debug_mode:\n        print(f\"Debug mode: processing only {cfg_regnet.debug_count} files\")\n        test_files = test_files[:cfg_regnet.debug_count]\n    \n    print(f\"Found {len(test_files)} test soundscapes\")\n    \n    # Process each soundscape\n    all_row_ids = []\n    all_predictions = []\n    \n    for audio_path in tqdm(test_files, desc=\"Processing test soundscapes\"):\n        row_ids, predictions = predict_on_soundscape(model, str(audio_path), species_map)\n        all_row_ids.extend(row_ids)\n        all_predictions.extend(predictions)\n    \n    # Create submission file\n    create_submission_file(all_row_ids, all_predictions)\n    \n    print(f\"Inference completed in {(time.time() - start_time)/60:.2f} minutes\")\n\ndef run_regnet():\n    try:\n        run_inference()\n        \n        # Verify submission\n        try:\n            sample_sub = pd.read_csv(cfg_regnet.sample_submission)\n            submission = pd.read_csv('submission_regnet.csv')\n            \n            print(f\"Submission file: {submission.shape}, Sample: {sample_sub.shape}\")\n            if set(submission.columns) != set(sample_sub.columns):\n                print(\"WARNING: Column mismatch!\")\n            else:\n                print(\"Submission has correct columns ✓\")\n        except Exception as e:\n            print(f\"Error verifying submission: {e}\")\n        \n    except Exception as e:\n        print(f\"Error during inference: {e}\")\n        import traceback\n        traceback.print_exc()\ncfg_regnet = CFG_regnet()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"run_regnet()","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}