{"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":"none","dataSources":[{"sourceId":91844,"databundleVersionId":11361821,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":393191,"sourceType":"modelInstanceVersion","isSourceIdPinned":false,"modelInstanceId":323580,"modelId":344368}],"dockerImageVersionId":31040,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"\"\"\"\n# BirdCLEF 2025 - AST Model Inference Notebook\n\nThis notebook loads a trained Audio Spectrogram Transformer (AST) model and creates predictions for the BirdCLEF 2025 test data.\n\"\"\"\n\nimport os\nimport gc\nimport warnings\nimport logging\nimport time\nimport numpy as np\nimport pandas as pd\nimport librosa\nimport cv2\nfrom tqdm.auto import tqdm\nfrom pathlib import Path\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport timm\nfrom timm.models.layers import to_2tuple, trunc_normal_\n\nwarnings.filterwarnings(\"ignore\")\nlogging.basicConfig(level=logging.INFO)\n\n# Configuration class\nclass CFG:\n    # Paths\n    test_soundscapes = '/kaggle/input/birdclef-2025/test_soundscapes'\n    sample_submission = '/kaggle/input/birdclef-2025/sample_submission.csv'\n    taxonomy_csv = '/kaggle/input/birdclef-2025/taxonomy.csv'\n    model_path = '/kaggle/input/ast/pytorch/default/1/final_model.pth'\n    \n    # Audio parameters\n    sample_rate = 32000\n    duration = 5  # seconds\n    \n    # Mel spectrogram parameters - must match training settings\n    n_mels = 128\n    n_fft = 1024\n    hop_length = 512\n    fmin = 50\n    fmax = 14000\n    \n    # Image parameters\n    target_height = 224\n    target_width = 224\n    \n    # AST model parameters\n    fstride = 10\n    tstride = 10\n    patch_size = 16\n    model_size = 'base224'\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 = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\ncfg = CFG()\nprint(f\"Using device: {cfg.device}\")\n\n# Define AST model architecture - must match the training model\nclass PatchEmbed(nn.Module):\n    \"\"\"2D Image to Patch Embedding\"\"\"\n    def __init__(self, img_size=224, patch_size=16, in_chans=3, embed_dim=768):\n        super().__init__()\n        \n        img_size = to_2tuple(img_size)\n        patch_size = to_2tuple(patch_size)\n        num_patches = (img_size[1] // patch_size[1]) * (img_size[0] // patch_size[0])\n        self.img_size = img_size\n        self.patch_size = patch_size\n        self.num_patches = num_patches\n        \n        self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size)\n        \n    def forward(self, x):\n        x = self.proj(x).flatten(2).transpose(1, 2)\n        return x\n\nclass ASTModel(nn.Module):\n    \"\"\"Audio Spectrogram Transformer model\"\"\"\n    def __init__(self, label_dim=527, fstride=10, tstride=10, input_fdim=128, input_tdim=1024, \n                 imagenet_pretrain=True, model_size='base224'):\n        super(ASTModel, self).__init__()\n        \n        # Override timm input shape restriction\n        timm.models.vision_transformer.PatchEmbed = PatchEmbed\n        \n        # Print model configuration for debugging\n        print(f'AST Model: size={model_size}, input_fdim={input_fdim}, input_tdim={input_tdim}')\n        print(f'frequency stride={fstride}, time stride={tstride}')\n        \n        # Load the appropriate ViT model\n        if model_size == 'tiny224':\n            try:\n                self.v = timm.create_model('vit_deit_tiny_distilled_patch16_224', pretrained=imagenet_pretrain)\n            except RuntimeError:\n                print(\"Falling back to vit_tiny_patch16_224\")\n                self.v = timm.create_model('vit_tiny_patch16_224', pretrained=imagenet_pretrain)\n        elif model_size == 'small224':\n            try:\n                self.v = timm.create_model('vit_deit_small_distilled_patch16_224', pretrained=imagenet_pretrain)\n            except RuntimeError:\n                print(\"Falling back to vit_small_patch16_224\")\n                self.v = timm.create_model('vit_small_patch16_224', pretrained=imagenet_pretrain)\n        elif model_size == 'base224':\n            try:\n                self.v = timm.create_model('vit_deit_base_distilled_patch16_224', pretrained=imagenet_pretrain)\n            except RuntimeError:\n                print(\"Falling back to vit_base_patch16_224\")\n                self.v = timm.create_model('vit_base_patch16_224', pretrained=imagenet_pretrain)\n        elif model_size == 'base384':\n            try:\n                self.v = timm.create_model('vit_deit_base_distilled_patch16_384', pretrained=imagenet_pretrain)\n            except RuntimeError:\n                print(\"Falling back to vit_base_patch16_384\")\n                try:\n                    self.v = timm.create_model('vit_base_patch16_384', pretrained=imagenet_pretrain)\n                except RuntimeError:\n                    print(\"Falling back to vit_base_patch16_224\")\n                    self.v = timm.create_model('vit_base_patch16_224', pretrained=imagenet_pretrain)\n        else:\n            raise Exception('Model size must be one of tiny224, small224, base224, base384.')\n            \n        # Check if model has distillation token\n        self.has_dist_token = hasattr(self.v, 'dist_token')\n        print(f\"Model has distillation token: {self.has_dist_token}\")\n        \n        self.original_num_patches = self.v.patch_embed.num_patches\n        self.oringal_hw = int(self.original_num_patches ** 0.5)\n        self.original_embedding_dim = self.v.pos_embed.shape[2]\n        self.mlp_head = nn.Sequential(nn.LayerNorm(self.original_embedding_dim), \n                                     nn.Linear(self.original_embedding_dim, label_dim))\n        \n        # Get shape automatically\n        f_dim, t_dim = self.get_shape(fstride, tstride, input_fdim, input_tdim)\n        num_patches = f_dim * t_dim\n        self.v.patch_embed.num_patches = num_patches\n        \n        print(f'number of patches={num_patches}')\n            \n        # Linear projection\n        new_proj = torch.nn.Conv2d(1, self.original_embedding_dim, kernel_size=(16, 16), stride=(fstride, tstride))\n        if imagenet_pretrain:\n            new_proj.weight = torch.nn.Parameter(torch.sum(self.v.patch_embed.proj.weight, dim=1).unsqueeze(1))\n            new_proj.bias = self.v.patch_embed.proj.bias\n        self.v.patch_embed.proj = new_proj\n        \n        # Positional embedding\n        if imagenet_pretrain:\n            # Get the positional embedding from model\n            if self.has_dist_token:\n                new_pos_embed = self.v.pos_embed[:, 2:, :].detach().reshape(1, self.original_num_patches, self.original_embedding_dim).transpose(1, 2).reshape(1, self.original_embedding_dim, self.oringal_hw, self.oringal_hw)\n            else:\n                new_pos_embed = self.v.pos_embed[:, 1:, :].detach().reshape(1, self.original_num_patches, self.original_embedding_dim).transpose(1, 2).reshape(1, self.original_embedding_dim, self.oringal_hw, self.oringal_hw)\n            \n            # Cut or interpolate position embedding\n            if t_dim <= self.oringal_hw:\n                new_pos_embed = new_pos_embed[:, :, :, int(self.oringal_hw / 2) - int(t_dim / 2): int(self.oringal_hw / 2) - int(t_dim / 2) + t_dim]\n            else:\n                new_pos_embed = torch.nn.functional.interpolate(new_pos_embed, size=(self.oringal_hw, t_dim), mode='bilinear')\n                \n            # Cut or interpolate position embedding\n            if f_dim <= self.oringal_hw:\n                new_pos_embed = new_pos_embed[:, :, int(self.oringal_hw / 2) - int(f_dim / 2): int(self.oringal_hw / 2) - int(f_dim / 2) + f_dim, :]\n            else:\n                new_pos_embed = torch.nn.functional.interpolate(new_pos_embed, size=(f_dim, t_dim), mode='bilinear')\n                \n            # Flatten the position embedding\n            new_pos_embed = new_pos_embed.reshape(1, self.original_embedding_dim, num_patches).transpose(1, 2)\n            \n            # Concatenate with cls token and distillation token\n            if self.has_dist_token:\n                self.v.pos_embed = nn.Parameter(torch.cat([self.v.pos_embed[:, :2, :].detach(), new_pos_embed], dim=1))\n            else:\n                self.v.pos_embed = nn.Parameter(torch.cat([self.v.pos_embed[:, :1, :].detach(), new_pos_embed], dim=1))\n        else:\n            # Random initialization\n            if self.has_dist_token:\n                new_pos_embed = nn.Parameter(torch.zeros(1, self.v.patch_embed.num_patches + 2, self.original_embedding_dim))\n            else:\n                new_pos_embed = nn.Parameter(torch.zeros(1, self.v.patch_embed.num_patches + 1, self.original_embedding_dim))\n            self.v.pos_embed = new_pos_embed\n            trunc_normal_(self.v.pos_embed, std=.02)\n        \n    def get_shape(self, fstride, tstride, input_fdim=128, input_tdim=1024):\n        test_input = torch.randn(1, 1, input_fdim, input_tdim)\n        test_proj = nn.Conv2d(1, self.original_embedding_dim, kernel_size=(16, 16), stride=(fstride, tstride))\n        test_out = test_proj(test_input)\n        f_dim = test_out.shape[2]\n        t_dim = test_out.shape[3]\n        return f_dim, t_dim\n    \n    def forward(self, x):\n        \"\"\"\n        :param x: Input spectrogram, expected shape: (batch_size, time_frame_num, frequency_bins)\n        :return: prediction\n        \"\"\"\n        # Input shape: (batch_size, time_frame_num, frequency_bins)\n        x = x.unsqueeze(1)        # Add channel dimension: (B, 1, T, F)\n        x = x.transpose(2, 3)     # -> (B, 1, F, T)\n        \n        B = x.shape[0]\n        x = self.v.patch_embed(x)\n        \n        # Handle both model types (with and without distillation token)\n        if self.has_dist_token:\n            cls_tokens = self.v.cls_token.expand(B, -1, -1)\n            dist_token = self.v.dist_token.expand(B, -1, -1)\n            x = torch.cat((cls_tokens, dist_token, x), dim=1)\n        else:\n            cls_tokens = self.v.cls_token.expand(B, -1, -1)\n            x = torch.cat((cls_tokens, x), dim=1)\n            \n        x = x + self.v.pos_embed\n        x = self.v.pos_drop(x)\n        \n        for blk in self.v.blocks:\n            x = blk(x)\n            \n        x = self.v.norm(x)\n        \n        # Handle both model types for output\n        if self.has_dist_token:\n            x = (x[:, 0] + x[:, 1]) / 2  # Average of cls and dist token\n        else:\n            x = x[:, 0]  # Just use cls token\n        \n        x = self.mlp_head(x)\n        return x\n\n# Audio processing functions\ndef audio_to_melspec(audio_data, cfg):\n    \"\"\"Convert audio data to mel spectrogram\"\"\"\n    # Handle NaN values\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    # Generate mel spectrogram\n    mel_spec = librosa.feature.melspectrogram(\n        y=audio_data,\n        sr=cfg.sample_rate,\n        n_fft=cfg.n_fft,\n        hop_length=cfg.hop_length,\n        n_mels=cfg.n_mels,\n        fmin=cfg.fmin,\n        fmax=cfg.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 to [0, 1]\n    mel_spec_norm = (mel_spec_db - mel_spec_db.min()) / (mel_spec_db.max() - mel_spec_db.min() + 1e-8)\n    \n    return mel_spec_norm\n\ndef apply_tta(mel_spec, step):\n    \"\"\"Apply test-time augmentation transformations\"\"\"\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.model_path}\")\n    \n    # First load the sample submission to get expected column names\n    sample_sub = pd.read_csv(cfg.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 - handle different checkpoint formats\n    try:\n        checkpoint = torch.load(cfg.model_path, map_location=cfg.device)\n        print(f\"Checkpoint keys: {list(checkpoint.keys())}\")\n        \n        # Extract number of classes and label map if available in checkpoint\n        if 'num_classes' in checkpoint:\n            num_classes = checkpoint['num_classes']\n            print(f\"Using num_classes from checkpoint: {num_classes}\")\n        else:\n            num_classes = num_species\n            print(f\"Using num_classes from sample submission: {num_classes}\")\n        \n        if 'label_map' in checkpoint:\n            label_map = checkpoint['label_map']\n            print(f\"Found label_map in checkpoint with {len(label_map)} entries\")\n        else:\n            # If no label map in checkpoint, use 1:1 mapping\n            label_map = {i: species for i, species in enumerate(species_columns)}\n            print(\"Created 1:1 label mapping\")\n        \n        # Create model with the appropriate number of classes\n        model = ASTModel(\n            label_dim=num_classes,\n            fstride=cfg.fstride,\n            tstride=cfg.tstride,\n            input_fdim=cfg.target_height,\n            input_tdim=cfg.target_width,\n            imagenet_pretrain=False,  # Not using pretrained weights for inference\n            model_size=cfg.model_size\n        )\n        \n        # Load model weights\n        if 'model_state_dict' in checkpoint:\n            model.load_state_dict(checkpoint['model_state_dict'])\n            print(\"Loaded weights from 'model_state_dict'\")\n        elif 'state_dict' in checkpoint:\n            model.load_state_dict(checkpoint['state_dict'])\n            print(\"Loaded weights from 'state_dict'\")\n        else:\n            print(\"WARNING: Could not find model weights in checkpoint\")\n            \n    except Exception as e:\n        print(f\"Error loading model: {e}\")\n        import traceback\n        traceback.print_exc()\n        return None, {}\n    \n    model = model.to(cfg.device)\n    model.eval()\n    \n    # Create output mapping from model outputs to species indices\n    if isinstance(label_map, dict):\n        # If label_map maps from labels to indices, invert it\n        if not all(isinstance(k, int) for k in label_map.keys()):\n            index_to_species = {v: k for k, v in label_map.items()}\n        else:\n            index_to_species = label_map\n    else:\n        index_to_species = {i: species for i, species in enumerate(species_columns)}\n    \n    return model, index_to_species, species_columns\n\ndef predict_on_soundscape(model, audio_path, index_to_species, species_columns):\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.sample_rate,\n            res_type='kaiser_fast'  # Faster resampling\n        )\n        \n        # Calculate total segments\n        segment_samples = cfg.sample_rate * cfg.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 5-second 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.duration\n            row_id = f\"{soundscape_id}_{end_time_sec}\"\n            row_ids.append(row_id)\n            \n            # Initialize prediction dictionary with zeros for all species\n            pred_dict = {species: 0.0 for species in species_columns}\n            \n            if cfg.use_tta:\n                # Apply test-time augmentation\n                segment_preds = []\n                \n                for tta_step in range(cfg.tta_steps):\n                    # Process audio to mel spectrogram\n                    mel_spec = audio_to_melspec(segment_audio, cfg)\n                    mel_spec = apply_tta(mel_spec, tta_step)\n                    \n                    # Resize to model's expected input dimensions\n                    mel_spec = cv2.resize(mel_spec, (cfg.target_width, cfg.target_height))\n                    \n                    # Convert to tensor\n                    mel_spec_tensor = torch.tensor(mel_spec, dtype=torch.float32).unsqueeze(0)\n                    mel_spec_tensor = mel_spec_tensor.to(cfg.device)\n                    \n                    # Get predictions\n                    with torch.no_grad():\n                        outputs = model(mel_spec_tensor)\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)\n                \n                # Resize to model's expected input dimensions\n                mel_spec = cv2.resize(mel_spec, (cfg.target_width, cfg.target_height))\n                \n                # Convert to tensor\n                mel_spec_tensor = torch.tensor(mel_spec, dtype=torch.float32).unsqueeze(0)\n                mel_spec_tensor = mel_spec_tensor.to(cfg.device)\n                \n                # Get predictions\n                with torch.no_grad():\n                    outputs = model(mel_spec_tensor)\n                    final_preds = torch.sigmoid(outputs).cpu().numpy().squeeze()\n            \n            # Map model outputs to species columns\n            if np.isscalar(final_preds):\n                # Handle case where there's only one class\n                if 0 in index_to_species and index_to_species[0] in species_columns:\n                    pred_dict[index_to_species[0]] = float(final_preds)\n            else:\n                # Map each model output to the corresponding species\n                for i, prob in enumerate(final_preds):\n                    if i in index_to_species and index_to_species[i] in species_columns:\n                        pred_dict[index_to_species[i]] = 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, species_columns):\n    \"\"\"Create submission file in the required format\"\"\"\n    print(\"Creating submission file...\")\n    \n    # Initialize submission with row_ids\n    submission_df = pd.DataFrame({'row_id': all_row_ids})\n    \n    # Add each species column\n    for col in species_columns:\n        # Initialize with zeros\n        submission_df[col] = 0.0\n        \n        # Update with actual predictions\n        for i, pred_dict in enumerate(all_predictions):\n            if i < len(submission_df) and col in pred_dict:\n                submission_df.loc[i, col] = pred_dict[col]\n    \n    # Save to CSV\n    submission_df.to_csv('submission.csv', index=False, float_format='%.6f')\n    print(f\"Submission file created with {len(submission_df)} predictions\")\n    \n    return submission_df\n\ndef run_inference():\n    \"\"\"Main inference function\"\"\"\n    start_time = time.time()\n    print(f\"Starting inference using device: {cfg.device}\")\n    \n    # Load model with species mapping\n    model, index_to_species, species_columns = load_model_and_species()\n    \n    if model is None:\n        print(\"Failed to load model. Exiting.\")\n        return\n    \n    # Find test files\n    test_files = list(Path(cfg.test_soundscapes).glob('*.ogg'))\n    \n    if cfg.debug_mode:\n        print(f\"Debug mode: processing 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    # 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), index_to_species, species_columns)\n        all_row_ids.extend(row_ids)\n        all_predictions.extend(predictions)\n    \n    # Create submission file\n    submission_df = create_submission_file(all_row_ids, all_predictions, species_columns)\n    \n    # Verify submission against sample submission\n    try:\n        sample_sub = pd.read_csv(cfg.sample_submission)\n        print(f\"Sample submission shape: {sample_sub.shape}\")\n        print(f\"Created submission shape: {submission_df.shape}\")\n        \n        missing_cols = set(sample_sub.columns) - set(submission_df.columns)\n        extra_cols = set(submission_df.columns) - set(sample_sub.columns)\n        \n        if missing_cols:\n            print(f\"WARNING: Missing columns in submission: {missing_cols}\")\n        if extra_cols:\n            print(f\"WARNING: Extra columns in submission: {extra_cols}\")\n        \n        if set(submission_df.columns) == set(sample_sub.columns):\n            print(\"✓ Submission has the correct columns\")\n        \n    except Exception as e:\n        print(f\"Error verifying submission: {e}\")\n    \n    print(f\"Inference completed in {(time.time() - start_time)/60:.2f} minutes\")\n\nif __name__ == \"__main__\":\n    try:\n        run_inference()\n    except Exception as e:\n        print(f\"Error during inference: {e}\")\n        import traceback\n        traceback.print_exc()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-05-15T07:00:55.617761Z","iopub.execute_input":"2025-05-15T07:00:55.618152Z","iopub.status.idle":"2025-05-15T07:01:14.236362Z","shell.execute_reply.started":"2025-05-15T07:00:55.618122Z","shell.execute_reply":"2025-05-15T07:01:14.233209Z"},"scrolled":true},"outputs":[],"execution_count":null}]}