{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":91844,"databundleVersionId":11361821,"sourceType":"competition"},{"sourceId":11702705,"sourceType":"datasetVersion","datasetId":7339703},{"sourceId":406521,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":332182,"modelId":353110}],"dockerImageVersionId":31040,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport pandas as pd\nimport numpy as np\nfrom tqdm.auto import tqdm\n\n# PyTorch imports\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.optim import lr_scheduler\nfrom torch.utils.data import DataLoader, Dataset\n\n# Machine learning\nfrom sklearn.model_selection import StratifiedKFold\nfrom sklearn.metrics import roc_auc_score\n\n# Deep learning models\nimport timm\n\n# Device configuration\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Using device: {device}\")\n\n# Paths\nDATA_ROOT = '/kaggle/input/birdclef-2025'  # Current directory\nTAXONOMY_CSV = os.path.join(DATA_ROOT, 'taxonomy.csv')\nTRAIN_CSV = os.path.join(DATA_ROOT, 'train.csv')\n\n# Image parameters\nTARGET_SHAPE = (256, 256)\n\nclass CFG:\n    # Training parameters\n    epochs = 25\n    batch_size = 32\n    \n    # Model optimization\n    model_name = 'vit_base_patch16_224'\n    img_size = (256, 256)\n    dropout_rate = 0.2\n    \n    # Knowledge distillation\n    use_distillation = True\n    distillation_temp = 3.0\n    distillation_alpha = 0.5\n    \n    # Optimizer settings\n    optimizer = 'AdamW'\n    lr = 5e-4\n    weight_decay = 1e-5\n  \n    # Learning rate scheduler\n    scheduler = 'CosineAnnealingLR'\n    min_lr = 1e-6\n    T_max = epochs\n    \n    # Debug mode settings\n    debug = False\n\nclass BirdCLEFDataset(Dataset):\n    def __init__(self, df, spectrograms=None, mode=\"train\"):\n        self.df = df\n        self.mode = mode\n        self.spectrograms = spectrograms\n\n        # Load taxonomy for label mapping\n        taxonomy_df = pd.read_csv(TAXONOMY_CSV)\n        self.species_ids = taxonomy_df['primary_label'].tolist()\n        self.num_classes = len(self.species_ids)\n        self.label_to_idx = {label: idx for idx, label in enumerate(self.species_ids)}\n\n        # Add sample name column for spectrogram lookup\n        if 'samplename' not in self.df.columns:\n            self.df['samplename'] = self.df.filename.map(lambda x: x.split('/')[0] + '-' + x.split('/')[-1].split('.')[0])\n                    \n        # Report stats\n        if self.spectrograms:\n            sample_names = set(self.df['samplename'])\n            found_samples = sum(1 for name in sample_names if name in self.spectrograms)\n            print(f\"Found {found_samples} spectrograms for {mode} dataset out of {len(self.df)} samples\")\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        samplename = row['samplename']\n        \n        # Get spectrogram or create blank if not found\n        if self.spectrograms and samplename in self.spectrograms:\n            spec = self.spectrograms[samplename]\n        else:\n            spec = np.zeros(TARGET_SHAPE, dtype=np.float32)\n        \n        # Add channel dimension: [1, H, W]\n        spec = torch.tensor(spec, dtype=torch.float32).unsqueeze(0)\n\n        # One-hot encode label\n        target = np.zeros(self.num_classes, dtype=np.float32)\n        label = row['primary_label']\n        if label in self.label_to_idx:\n            target[self.label_to_idx[label]] = 1.0\n\n        # Handle secondary labels if present\n        if 'secondary_labels' in row and row['secondary_labels'] not in [[''], None, np.nan]:\n            if isinstance(row['secondary_labels'], str):\n                secondary_labels = eval(row['secondary_labels'])\n            else:\n                secondary_labels = row['secondary_labels']\n            \n            for label in secondary_labels:\n                if label in self.label_to_idx:\n                    target[self.label_to_idx[label]] = 1.0\n\n        return {\n            'melspec': spec,\n            'target': torch.tensor(target, dtype=torch.float32),\n            'filename': row['filename']\n        }\n\nclass BirdCLEFModel(nn.Module):\n    def __init__(self, num_classes, pretrained=True, cfg=None):\n        super().__init__()\n        \n        # Use config if provided, otherwise use defaults\n        self.cfg = cfg if cfg else CFG()\n        \n        # Create ViT for audio spectrograms\n        self.vit = timm.create_model(\n            self.cfg.model_name,\n            pretrained=pretrained,\n            img_size=self.cfg.img_size,\n            in_chans=1,\n            num_classes=0,  # Get embeddings\n            drop_rate=self.cfg.dropout_rate\n        )\n        \n        # Get the output feature dimension\n        backbone_out = self.vit.embed_dim\n        \n        # Classifier head\n        self.classifier = nn.Sequential(\n            nn.LayerNorm(backbone_out),\n            nn.Dropout(0.2),\n            nn.Linear(backbone_out, num_classes)\n        )\n            \n    def forward(self, x):\n        # Extract features from ViT\n        features = self.vit(x)\n        \n        # Get logits from classifier\n        logits = self.classifier(features)\n            \n        return logits\n\nclass AsymmetricLossMultiLabel(nn.Module):\n    def __init__(\n        self,\n        gamma_neg=4,\n        gamma_pos=1,\n        clip=0.05,\n        eps=1e-8,\n        disable_torch_grad_focal_loss=False,\n        reduction=\"mean\",\n    ):\n        super().__init__()\n\n        self.gamma_neg = gamma_neg\n        self.gamma_pos = gamma_pos\n        self.clip = clip\n        self.disable_torch_grad_focal_loss = disable_torch_grad_focal_loss\n        self.eps = eps\n        self.reduction = reduction\n\n    def forward(self, x, y):\n        # Calculating Probabilities\n        x_sigmoid = torch.sigmoid(x)\n        xs_pos = x_sigmoid\n        xs_neg = 1 - x_sigmoid\n\n        # Asymmetric Clipping\n        if self.clip is not None and self.clip > 0:\n            xs_neg = (xs_neg + self.clip).clamp(max=1)\n\n        # Basic CE calculation\n        los_pos = y * torch.log(xs_pos.clamp(min=self.eps))\n        los_neg = (1 - y) * torch.log(xs_neg.clamp(min=self.eps))\n        loss = los_pos + los_neg\n\n        # Asymmetric Focusing\n        if self.gamma_neg > 0 or self.gamma_pos > 0:\n            if self.disable_torch_grad_focal_loss:\n                torch._C.set_grad_enabled(False)\n            pt0 = xs_pos * y\n            pt1 = xs_neg * (1 - y)  # pt = p if t > 0 else 1-p\n            pt = pt0 + pt1\n            one_sided_gamma = self.gamma_pos * y + self.gamma_neg * (1 - y)\n            one_sided_w = torch.pow(1 - pt, one_sided_gamma)\n            if self.disable_torch_grad_focal_loss:\n                torch._C.set_grad_enabled(True)\n            loss *= one_sided_w\n\n        if self.reduction == \"mean\":\n            return -loss.mean()\n        if self.reduction == \"sum\":\n            return -loss.sum()\n\n        return -loss\n\ndef collate_fn(batch):\n    \"\"\"Collate function for dataloaders\"\"\"\n    batch = [item for item in batch if item is not None]\n    if not batch:\n        return {}\n        \n    # Collect items by key\n    melspecs = []\n    targets = []\n    filenames = []\n    \n    for item in batch:\n        if 'melspec' in item:\n            melspecs.append(item['melspec'])\n        if 'target' in item:\n            targets.append(item['target'])\n        if 'filename' in item:\n            filenames.append(item['filename'])\n    \n    # Stack tensors when possible\n    result = {}\n    if melspecs:\n        if isinstance(melspecs[0], torch.Tensor) and all(m.shape == melspecs[0].shape for m in melspecs):\n            result['melspec'] = torch.stack(melspecs)\n        else:\n            result['melspec'] = melspecs\n    \n    if targets and isinstance(targets[0], torch.Tensor):\n        result['target'] = torch.stack(targets)\n    else:\n        result['target'] = targets\n        \n    if filenames:\n        result['filename'] = filenames\n    \n    return result\n\ndef load_teacher_model(model_path, num_classes):\n    \"\"\"Load the teacher model for distillation\"\"\"\n    try:\n        print(f\"Loading teacher model from {model_path}\")\n        checkpoint = {}\n        try:\n            # Try with weights_only=False (has better compatibility)\n            checkpoint = torch.load(model_path, map_location=device, weights_only=False)\n        except:\n            # Fallback: try with weights_only=True\n            try:\n                checkpoint = torch.load(model_path, map_location=device, weights_only=True)\n            except Exception as e:\n                print(f\"Error loading model with standard approach: {e}\")\n                # Last resort: try to load directly with safe_globals\n                try:\n                    from torch.serialization import safe_globals\n                    import numpy as np\n                    with safe_globals([np.core.multiarray.scalar, np.dtype]):\n                        checkpoint = torch.load(model_path, map_location=device, weights_only=False)\n                except Exception as e2:\n                    print(f\"All loading attempts failed: {e2}\")\n                    return None\n        \n        # Create a model instance\n        cfg = CFG()\n        teacher = BirdCLEFModel(num_classes=num_classes, pretrained=False, cfg=cfg)\n        \n        # Load state dict\n        if 'model_state_dict' in checkpoint:\n            teacher.load_state_dict(checkpoint['model_state_dict'])\n            print(\"Teacher model loaded successfully\")\n        else:\n            print(\"Warning: Could not find model_state_dict in checkpoint\")\n            return None\n            \n        # Set to eval mode\n        teacher.eval()\n        return teacher\n            \n    except Exception as e:\n        print(f\"Error loading teacher model: {e}\")\n        return None\n\ndef distillation_loss(outputs, labels, teacher_outputs, temp, alpha):\n    \"\"\"Knowledge distillation loss function\"\"\"\n    # Hard target loss (standard cross-entropy)\n    hard_loss = nn.BCEWithLogitsLoss()(outputs, labels)\n    \n    # Soft target loss (KL divergence between softened distributions)\n    soft_loss = nn.KLDivLoss(reduction='batchmean')(\n        torch.log_softmax(outputs / temp, dim=1),\n        torch.softmax(teacher_outputs / temp, dim=1)\n    ) * (temp * temp)\n    \n    # Combine losses\n    return alpha * hard_loss + (1 - alpha) * soft_loss\n\ndef train_one_epoch(model, teacher_model, loader, optimizer, criterion, device, cfg):\n    model.train()\n    losses = []\n    \n    with tqdm(loader, desc=\"Training\") as pbar:\n        for batch in pbar:\n            if not batch:  # Skip empty batches\n                continue\n                \n            inputs = batch['melspec'].to(device)\n            targets = batch['target'].to(device)\n            \n            # Zero gradients\n            optimizer.zero_grad()\n            \n            # Forward pass\n            outputs = model(inputs)\n            \n            # Calculate loss\n            if cfg.use_distillation and teacher_model is not None:\n                with torch.no_grad():\n                    teacher_outputs = teacher_model(inputs)\n                loss = distillation_loss(\n                    outputs, \n                    targets, \n                    teacher_outputs, \n                    cfg.distillation_temp, \n                    cfg.distillation_alpha\n                )\n            else:\n                loss = criterion(outputs, targets)\n            \n            # Backward pass and optimize\n            loss.backward()\n            optimizer.step()\n            \n            # Record loss\n            losses.append(loss.item())\n            \n            # Update progress bar\n            pbar.set_postfix({'loss': np.mean(losses[-10:]) if losses else 0})\n    \n    return np.mean(losses)\n\ndef validate(model, loader, criterion, device):\n    model.eval()\n    losses = []\n    all_targets = []\n    all_outputs = []\n    \n    with torch.no_grad():\n        for batch in tqdm(loader, desc=\"Validation\"):\n            if not batch:  # Skip empty batches\n                continue\n                \n            inputs = batch['melspec'].to(device)\n            targets = batch['target'].to(device)\n            \n            # Forward pass\n            outputs = model(inputs)\n            loss = criterion(outputs, targets)\n            \n            # Collect results\n            all_outputs.append(outputs.cpu().numpy())\n            all_targets.append(targets.cpu().numpy())\n            losses.append(loss.item())\n    \n    # Calculate metrics\n    all_outputs = np.concatenate(all_outputs)\n    all_targets = np.concatenate(all_targets)\n    \n    auc = calculate_auc(all_targets, all_outputs)\n    avg_loss = np.mean(losses)\n    \n    return avg_loss, auc\n\ndef calculate_auc(targets, outputs):\n    \"\"\"Calculate AUC for multi-label classification\"\"\"\n    num_classes = targets.shape[1]\n    aucs = []\n    \n    # Convert logits to probabilities\n    probs = 1 / (1 + np.exp(-outputs))\n    \n    for i in range(num_classes):\n        # Only calculate AUC for classes that have positive samples\n        if np.sum(targets[:, i]) > 0:\n            class_auc = roc_auc_score(targets[:, i], probs[:, i])\n            aucs.append(class_auc)\n    \n    return np.mean(aucs) if aucs else 0.0\n\ndef run_training(df, spectrograms, cfg):\n    \"\"\"Run training with knowledge distillation\"\"\"\n    \n    # Load taxonomy data\n    taxonomy_df = pd.read_csv(TAXONOMY_CSV)\n    num_classes = len(taxonomy_df['primary_label'])\n    \n    # Setup k-fold cross-validation\n    skf = StratifiedKFold(n_splits=5, shuffle=True, random_state=42)\n    \n    # Train for fold 0 only (simplified)\n    for fold, (train_idx, val_idx) in enumerate(skf.split(df, df['primary_label'])):\n        if fold > 0:  # Only use fold 0 for simplicity\n            break\n            \n        print(f'\\n{\"=\"*30} Fold {fold} {\"=\"*30}')\n        \n        # Split data\n        train_df = df.iloc[train_idx].reset_index(drop=True)\n        val_df = df.iloc[val_idx].reset_index(drop=True)\n        \n        print(f'Training set: {len(train_df)} samples')\n        print(f'Validation set: {len(val_df)} samples')\n        \n        # Create datasets\n        train_dataset = BirdCLEFDataset(train_df, spectrograms=spectrograms, mode='train')\n        val_dataset = BirdCLEFDataset(val_df, spectrograms=spectrograms, mode='valid')\n        \n        # Create data loaders\n        train_loader = DataLoader(\n            train_dataset, \n            batch_size=cfg.batch_size, \n            shuffle=True, \n            num_workers=2,\n            collate_fn=collate_fn\n        )\n        \n        val_loader = DataLoader(\n            val_dataset, \n            batch_size=cfg.batch_size, \n            shuffle=False, \n            num_workers=2,\n            collate_fn=collate_fn\n        )\n        \n        # Create student model\n        model = BirdCLEFModel(num_classes=num_classes, pretrained=True, cfg=cfg).to(device)\n        \n        # Load teacher model for distillation (if enabled)\n        teacher_model = None\n        if cfg.use_distillation:\n            teacher_model = load_teacher_model('/kaggle/input/distill-vit/pytorch/default/1/model_fold4.pth', num_classes)\n            if teacher_model:\n                teacher_model = teacher_model.to(device)\n        \n        # Setup training components\n        optimizer = optim.AdamW(model.parameters(), lr=cfg.lr, weight_decay=cfg.weight_decay)\n        criterion = AsymmetricLossMultiLabel()\n        scheduler = lr_scheduler.CosineAnnealingLR(optimizer, T_max=cfg.T_max, eta_min=cfg.min_lr)\n        \n        # Training loop\n        best_auc = 0\n        best_epoch = 0\n        \n        for epoch in range(cfg.epochs):\n            print(f\"\\nEpoch {epoch+1}/{cfg.epochs}\")\n            \n            # Train\n            train_loss = train_one_epoch(\n                model, \n                teacher_model, \n                train_loader, \n                optimizer, \n                criterion, \n                device,\n                cfg\n            )\n            \n            # Validate\n            val_loss, val_auc = validate(model, val_loader, criterion, device)\n            \n            # Update learning rate\n            scheduler.step()\n            \n            # Print metrics\n            print(f\"Train Loss: {train_loss:.4f}\")\n            print(f\"Val Loss: {val_loss:.4f}, Val AUC: {val_auc:.4f}\")\n            \n            # Save best model\n            if val_auc > best_auc:\n                best_auc = val_auc\n                best_epoch = epoch + 1\n                torch.save({\n                    'model_state_dict': model.state_dict(),\n                    'optimizer_state_dict': optimizer.state_dict(),\n                    'val_auc': val_auc\n                }, f'distilled_model_fold{fold}.pth')\n                print(f\"New best AUC: {best_auc:.4f} - Model saved\")\n        \n        print(f\"\\nBest AUC: {best_auc:.4f} at epoch {best_epoch}\")\n\ndef main():\n    # Check if required files exist\n    if not os.path.exists(TAXONOMY_CSV):\n        print(f\"Error: {TAXONOMY_CSV} not found\")\n        return\n    \n    if not os.path.exists(TRAIN_CSV):\n        print(f\"Error: {TRAIN_CSV} not found\")\n        return\n    \n    # Load data\n    print(\"Loading data...\")\n    taxonomy_df = pd.read_csv(TAXONOMY_CSV)\n    train_df = pd.read_csv(TRAIN_CSV)\n    \n    # Load spectrograms\n    print(\"Loading spectrograms...\")\n    spectrograms = None\n    try:\n        spectrograms = np.load(\"/kaggle/input/birdclef-melspec-data/bird_mel_spectrograms.npy\", allow_pickle=True).item()\n        print(f\"Loaded {len(spectrograms)} spectrograms\")\n    except Exception as e:\n        print(f\"Error loading spectrograms: {e}\")\n    \n    # Run training\n    cfg = CFG()\n    run_training(train_df, spectrograms, cfg)\n\nif __name__ == \"__main__\":\n    main() ","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-05-21T20:47:12.185062Z","iopub.execute_input":"2025-05-21T20:47:12.185663Z","execution_failed":"2025-05-21T20:49:01.183Z"}},"outputs":[],"execution_count":null}]}