{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"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":"nvidiaTeslaT4","dataSources":[{"sourceId":91844,"databundleVersionId":11361821,"sourceType":"competition"},{"sourceId":292976366,"sourceType":"kernelVersion"},{"sourceId":723199,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":550287,"modelId":562926},{"sourceId":725845,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":552497,"modelId":565058}],"dockerImageVersionId":31259,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# 🏷️ Génération de Pseudo-Labels avec le Teacher Model\n\nCe notebook utilise le modèle Teacher entraîné pour générer des pseudo-labels sur les données non-étiquetées (soundscapes).\n\n**Méthode :**\n1. **Fenêtre glissante** : Blocs de 10s avec chevauchement de 5s\n2. **Moyenne des prédictions** : Chaque seconde est prédite par plusieurs fenêtres\n3. **Raffinement** : Sigmoid → Transformée de puissance → Seuillage","metadata":{}},{"cell_type":"code","source":"import os\nimport glob\nimport torch\nimport torch.nn as nn\nimport torchaudio\nimport torchaudio.transforms as T\nimport torchvision.transforms as VT\nimport pandas as pd\nimport numpy as np\nimport timm\nfrom tqdm.auto import tqdm\nfrom pathlib import Path\n\n# Fix audio backend\ntry:\n    torchaudio.set_audio_backend(\"soundfile\")\nexcept:\n    pass\n\nprint(\"📚 Bibliothèques chargées.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T13:44:35.247244Z","iopub.execute_input":"2026-01-20T13:44:35.247663Z","iopub.status.idle":"2026-01-20T13:44:35.255101Z","shell.execute_reply.started":"2026-01-20T13:44:35.247632Z","shell.execute_reply":"2026-01-20T13:44:35.253972Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Config:\n    # --- CHEMINS KAGGLE ---\n    TEACHER_PATH = \"/kaggle/input/teacher-best-pth/other/default/1/teacher_epoch_10.pth\"  # Modèle Teacher entraîné\n    SOUNDSCAPES_DIR = \"/kaggle/input/birdclef-2025/train_soundscapes\"\n    OUTPUT_PATH = \"pseudo_labels.parquet\"\n    \n    # --- AUDIO (doit correspondre au Teacher) ---\n    SR = 32000\n    WINDOW_LEN = 32000 * 5   # 5 secondes (fenêtre d'analyse)\n    STRIDE = int(32000 * 2.5)        # 2.5 secondes (décalage entre fenêtres)\n    \n    # --- SPECTROGRAMME (doit correspondre au Teacher) ---\n    N_MELS = 224\n    IMG_SIZE = 224\n    \n    # --- RAFFINEMENT DES PSEUDO-LABELS ---\n    POWER_TRANSFORM = 1.5     # Exposant pour nettoyer le bruit\n    THRESHOLD = 0.1           # Seuil minimum de probabilité\n    \n    # --- INFÉRENCE ---\n    BATCH_SIZE = 8            # Batch pour l'inférence\n    DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\nprint(f\"⚙️ Configuration chargée. Device: {Config.DEVICE}\")\nprint(f\"   Fenêtre: {Config.WINDOW_LEN // Config.SR}s | Stride: {Config.STRIDE // Config.SR}s\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T13:44:35.257057Z","iopub.execute_input":"2026-01-20T13:44:35.257370Z","iopub.status.idle":"2026-01-20T13:44:35.279207Z","shell.execute_reply.started":"2026-01-20T13:44:35.257342Z","shell.execute_reply":"2026-01-20T13:44:35.277921Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# === ARCHITECTURE DU MODÈLE (identique au Teacher) ===\n\nclass GPUTransform(nn.Module):\n    \"\"\"Transformation Audio -> Spectrogramme sur GPU\"\"\"\n    def __init__(self):\n        super().__init__()\n        self.mel_spec = T.MelSpectrogram(\n            sample_rate=Config.SR,\n            n_mels=Config.N_MELS,\n            n_fft=2048,\n            hop_length=320,\n            f_min=20,\n            f_max=16000\n        )\n        self.to_db = T.AmplitudeToDB()\n        self.resize = VT.Resize((Config.N_MELS, Config.IMG_SIZE))\n\n    def forward(self, wav_batch):\n        spec = self.mel_spec(wav_batch) \n        spec = self.to_db(spec)\n        spec = spec.unsqueeze(1) \n        spec = self.resize(spec)\n        \n        # Normalisation par échantillon\n        spec_min = spec.amin(dim=(2, 3), keepdim=True)\n        spec_max = spec.amax(dim=(2, 3), keepdim=True)\n        spec = (spec - spec_min) / (spec_max - spec_min + 1e-6)\n        \n        spec = spec.repeat(1, 3, 1, 1) \n        return spec\n\n\nclass SEDHead(nn.Module):\n    \"\"\"Tête SED avec attention temporelle\"\"\"\n    def __init__(self, in_features, num_classes):\n        super().__init__()\n        self.fc = nn.Linear(in_features, num_classes)\n        self.attention = nn.Sequential(\n            nn.Linear(in_features, 128),\n            nn.Tanh(),\n            nn.Linear(128, 1)\n        )\n        \n    def forward(self, x):\n        x = x.transpose(1, 2)\n        att_weights = torch.softmax(self.attention(x), dim=1)\n        frame_pred = self.fc(x)\n        clip_pred = (frame_pred * att_weights).sum(dim=1)\n        return clip_pred, frame_pred\n\n\nclass BirdTeacherModel(nn.Module):\n    \"\"\"MobileNetV3 Large avec tête SED\"\"\"\n    def __init__(self, num_classes):\n        super().__init__()\n        self.backbone = timm.create_model(\n            'mobilenetv3_large_100', \n            pretrained=False,  # On charge les poids manuellement\n            num_classes=0,\n            global_pool=''\n        )\n        \n        # Récupère le nombre de features\n        with torch.no_grad():\n            dummy = torch.zeros(1, 3, Config.N_MELS, Config.IMG_SIZE)\n            feat = self.backbone(dummy)\n            self.num_features = feat.shape[1]\n        \n        self.sed_head = SEDHead(self.num_features, num_classes)\n        \n    def forward(self, x):\n        features = self.backbone(x)\n        features = features.mean(dim=2)\n        clip_pred, frame_pred = self.sed_head(features)\n        return clip_pred, frame_pred\n\nprint(\"🏗️ Architecture du modèle définie.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T13:44:35.280407Z","iopub.execute_input":"2026-01-20T13:44:35.280783Z","iopub.status.idle":"2026-01-20T13:44:35.300547Z","shell.execute_reply.started":"2026-01-20T13:44:35.280744Z","shell.execute_reply":"2026-01-20T13:44:35.299493Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_teacher_model():\n    \"\"\"\n    Charge le modèle Teacher pré-entraîné.\n    Gère différents formats de checkpoint.\n    \"\"\"\n    print(f\"📂 Chargement du modèle: {Config.TEACHER_PATH}\")\n    \n    checkpoint = torch.load(Config.TEACHER_PATH, map_location=Config.DEVICE)\n    \n    # Détection du format du checkpoint\n    if isinstance(checkpoint, dict):\n        # Format avec métadonnées (sauvegardé avec dict complet)\n        if 'num_classes' in checkpoint:\n            num_classes = checkpoint['num_classes']\n            labels = checkpoint['labels']\n            state_dict = checkpoint['model_state_dict']\n            val_acc = checkpoint.get('val_acc', 'N/A')\n            print(f\"   Classes: {num_classes} | Val Acc: {val_acc}\")\n        \n        # Format state_dict direct (torch.save(model.state_dict(), ...))\n        elif 'sed_head.fc.weight' in checkpoint:\n            # Récupère num_classes depuis la dernière couche\n            num_classes = checkpoint['sed_head.fc.weight'].shape[0]\n            state_dict = checkpoint\n            # Labels doivent être chargés séparément\n            labels = None\n            print(f\"   Classes: {num_classes} (state_dict direct)\")\n        \n        # Format model_state_dict sans autres métadonnées\n        elif 'model_state_dict' in checkpoint:\n            state_dict = checkpoint['model_state_dict']\n            if 'sed_head.fc.weight' in state_dict:\n                num_classes = state_dict['sed_head.fc.weight'].shape[0]\n            else:\n                raise ValueError(\"Impossible de détecter num_classes\")\n            labels = checkpoint.get('labels', None)\n            print(f\"   Classes: {num_classes}\")\n        \n        else:\n            raise ValueError(f\"Format de checkpoint non reconnu. Clés: {list(checkpoint.keys())[:10]}\")\n    else:\n        raise ValueError(\"Le checkpoint doit être un dictionnaire\")\n    \n    # Si labels non trouvés, charger depuis le CSV\n    if labels is None:\n        print(\"   ⚠️ Labels non trouvés dans le checkpoint, chargement depuis CSV...\")\n        try:\n            train_csv = Config.TEACHER_PATH.replace('teacher_best.pth', '') + '/kaggle/input/birdclef-2025/train.csv'\n            # Essayer plusieurs chemins\n            possible_paths = [\n                \"/kaggle/input/birdclef-2025/train.csv\",\n                \"train.csv\",\n                \"../birdclef-2025/train.csv\"\n            ]\n            for path in possible_paths:\n                if os.path.exists(path):\n                    df = pd.read_csv(path)\n                    labels = sorted(df['primary_label'].unique())\n                    print(f\"   ✅ Labels chargés depuis {path}\")\n                    break\n            if labels is None:\n                labels = [f\"class_{i}\" for i in range(num_classes)]\n                print(f\"   ⚠️ Labels génériques utilisés (class_0 à class_{num_classes-1})\")\n        except Exception as e:\n            labels = [f\"class_{i}\" for i in range(num_classes)]\n            print(f\"   ⚠️ Labels génériques utilisés: {e}\")\n    \n    # Création et chargement du modèle\n    model = BirdTeacherModel(num_classes).to(Config.DEVICE)\n    model.load_state_dict(state_dict)\n    model.eval()\n    \n    print(f\"   ✅ Modèle chargé avec succès!\")\n    \n    return model, labels\n\nprint(\"🔧 Fonction de chargement prête.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T13:44:35.301934Z","iopub.execute_input":"2026-01-20T13:44:35.302267Z","iopub.status.idle":"2026-01-20T13:44:35.325384Z","shell.execute_reply.started":"2026-01-20T13:44:35.302235Z","shell.execute_reply":"2026-01-20T13:44:35.324254Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def extract_sliding_windows(audio_path):\n    \"\"\"\n    Extrait des fenêtres glissantes d'un fichier audio.\n    Retourne les fenêtres et leurs positions temporelles (en secondes).\n    \"\"\"\n    wav, sr = torchaudio.load(audio_path)\n    \n    # Resample si nécessaire\n    if sr != Config.SR:\n        wav = T.Resample(sr, Config.SR)(wav)\n    \n    # Mono\n    if wav.shape[0] > 1:\n        wav = wav.mean(dim=0, keepdim=True)\n    wav = wav[0]  # (samples,)\n    \n    total_samples = wav.shape[0]\n    total_seconds = total_samples // Config.SR\n    \n    windows = []\n    positions = []  # Position de début de chaque fenêtre (en samples)\n    \n    start = 0\n    while start + Config.WINDOW_LEN <= total_samples:\n        window = wav[start:start + Config.WINDOW_LEN]\n        windows.append(window)\n        positions.append(start)\n        start += Config.STRIDE\n    \n    # Dernière fenêtre (padding si nécessaire)\n    if start < total_samples:\n        remaining = wav[start:]\n        if len(remaining) > Config.SR:  # Au moins 1 seconde\n            pad_len = Config.WINDOW_LEN - len(remaining)\n            padded = torch.cat([remaining, torch.zeros(pad_len)])\n            windows.append(padded)\n            positions.append(start)\n    \n    if len(windows) == 0:\n        # Fichier trop court : padding\n        pad_len = Config.WINDOW_LEN - total_samples\n        padded = torch.cat([wav, torch.zeros(pad_len)])\n        windows.append(padded)\n        positions.append(0)\n    \n    return torch.stack(windows), positions, total_seconds\n\nprint(\"🪟 Fonction d'extraction par fenêtre glissante prête.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T13:44:35.330413Z","iopub.execute_input":"2026-01-20T13:44:35.331709Z","iopub.status.idle":"2026-01-20T13:44:35.350186Z","shell.execute_reply.started":"2026-01-20T13:44:35.331671Z","shell.execute_reply":"2026-01-20T13:44:35.349257Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def refine_predictions(logits):\n    \"\"\"\n    Raffine les prédictions brutes en pseudo-labels de qualité.\n    \n    1. Sigmoid : logits -> probabilités [0, 1]\n    2. Power Transform : P^1.5 pour réduire le bruit\n    3. Threshold : Garde uniquement P > 0.1\n    \"\"\"\n    # 1. Activation Sigmoid (probabilités indépendantes par classe)\n    probs = torch.sigmoid(logits)\n    \n    # 2. Transformée de puissance (nettoie le bruit)\n    probs = torch.pow(probs, Config.POWER_TRANSFORM)\n    \n    # 3. Seuillage (met à 0 les probabilités trop faibles)\n    probs = torch.where(probs > Config.THRESHOLD, probs, torch.zeros_like(probs))\n    \n    return probs\n\n\ndef aggregate_overlapping_predictions(window_preds, positions, total_seconds, num_classes):\n    \"\"\"\n    Agrège les prédictions des fenêtres chevauchantes.\n    Chaque seconde reçoit la moyenne des prédictions des fenêtres qui la couvrent.\n    \n    Retourne un vecteur de probabilités par seconde.\n    \"\"\"\n    # Initialisation : compteur et somme par seconde\n    counts = torch.zeros(total_seconds)\n    predictions = torch.zeros(total_seconds, num_classes)\n    \n    window_duration = Config.WINDOW_LEN // Config.SR\n    \n    for i, (pred, start_sample) in enumerate(zip(window_preds, positions)):\n        start_sec = start_sample // Config.SR\n        end_sec = min(start_sec + window_duration, total_seconds)\n        \n        for sec in range(start_sec, end_sec):\n            predictions[sec] += pred\n            counts[sec] += 1\n    \n    # Moyenne (évite division par 0)\n    counts = counts.clamp(min=1).unsqueeze(1)\n    averaged_predictions = predictions / counts\n    \n    return averaged_predictions\n\nprint(\"🔬 Fonctions de raffinement et agrégation prêtes.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T13:44:35.352153Z","iopub.execute_input":"2026-01-20T13:44:35.352556Z","iopub.status.idle":"2026-01-20T13:44:35.373810Z","shell.execute_reply.started":"2026-01-20T13:44:35.352517Z","shell.execute_reply":"2026-01-20T13:44:35.372660Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@torch.no_grad()\ndef process_single_file(audio_path, model, transform, num_classes):\n    \"\"\"\n    Traite un fichier audio complet avec fenêtre glissante.\n    Retourne les pseudo-labels par seconde.\n    \"\"\"\n    # 1. Extraction des fenêtres\n    windows, positions, total_seconds = extract_sliding_windows(audio_path)\n    \n    if total_seconds == 0:\n        return None\n    \n    # 2. Inférence par batch\n    all_preds = []\n    \n    for i in range(0, len(windows), Config.BATCH_SIZE):\n        batch = windows[i:i + Config.BATCH_SIZE].to(Config.DEVICE)\n        \n        # Transformation spectrogramme\n        images = transform(batch)\n        \n        # Prédiction (clip-level)\n        clip_pred, _ = model(images)\n        \n        # Raffinement\n        refined = refine_predictions(clip_pred)\n        all_preds.append(refined.cpu())\n    \n    all_preds = torch.cat(all_preds, dim=0)\n    \n    # 3. Agrégation des prédictions chevauchantes\n    per_second_preds = aggregate_overlapping_predictions(\n        all_preds, positions, total_seconds, num_classes\n    )\n    \n    return per_second_preds\n\nprint(\"📄 Fonction de traitement de fichier prête.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T13:44:35.375134Z","iopub.execute_input":"2026-01-20T13:44:35.375460Z","iopub.status.idle":"2026-01-20T13:44:35.397104Z","shell.execute_reply.started":"2026-01-20T13:44:35.375425Z","shell.execute_reply":"2026-01-20T13:44:35.396141Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def run_pseudo_labeling():\n    \"\"\"\n    Pipeline complet de génération de pseudo-labels.\n    Traite tous les soundscapes et sauvegarde les résultats.\n    \"\"\"\n    print(\"=\" * 60)\n    print(\"🏷️  GÉNÉRATION DE PSEUDO-LABELS\")\n    print(\"=\" * 60)\n    \n    # 1. Chargement du modèle Teacher\n    model, labels = load_teacher_model()\n    transform = GPUTransform().to(Config.DEVICE)\n    num_classes = len(labels)\n    \n    # 2. Liste des fichiers audio à traiter\n    audio_files = glob.glob(os.path.join(Config.SOUNDSCAPES_DIR, \"*.ogg\"))\n    if len(audio_files) == 0:\n        audio_files = glob.glob(os.path.join(Config.SOUNDSCAPES_DIR, \"*.wav\"))\n    \n    print(f\"\\n📁 {len(audio_files)} fichiers soundscape trouvés.\")\n    \n    if len(audio_files) == 0:\n        print(\"❌ Aucun fichier audio trouvé!\")\n        return\n    \n    # 3. Traitement de chaque fichier\n    results = []\n    \n    for audio_path in tqdm(audio_files, desc=\"🔄 Traitement\"):\n        filename = Path(audio_path).stem\n        \n        try:\n            per_second_preds = process_single_file(\n                audio_path, model, transform, num_classes\n            )\n            \n            if per_second_preds is None:\n                continue\n            \n            # Format: une ligne par seconde\n            for sec_idx, probs in enumerate(per_second_preds):\n                row = {\n                    'filename': filename,\n                    'second': sec_idx,\n                }\n                # Ajoute les probabilités pour chaque espèce\n                for label, prob in zip(labels, probs.numpy()):\n                    row[label] = float(prob)\n                \n                results.append(row)\n                \n        except Exception as e:\n            print(f\"⚠️ Erreur sur {filename}: {e}\")\n            continue\n    \n    # 4. Sauvegarde des résultats\n    df = pd.DataFrame(results)\n    \n    print(f\"\\n📊 Résultats: {len(df)} lignes (secondes)\")\n    print(f\"   Fichiers traités: {df['filename'].nunique()}\")\n    \n    # Statistiques sur les pseudo-labels\n    label_cols = [c for c in df.columns if c not in ['filename', 'second']]\n    non_zero = (df[label_cols] > 0).sum().sum()\n    total = len(df) * len(label_cols)\n    sparsity = (1 - non_zero / total) * 100\n    \n    print(f\"   Sparsité: {sparsity:.1f}% (labels à 0)\")\n    \n    # Sauvegarde en Parquet (efficace) et CSV (lisible)\n    df.to_parquet(Config.OUTPUT_PATH, index=False)\n    df.to_csv(Config.OUTPUT_PATH.replace('.parquet', '.csv'), index=False)\n    \n    print(f\"\\n✅ Pseudo-labels sauvegardés:\")\n    print(f\"   → {Config.OUTPUT_PATH}\")\n    print(f\"   → {Config.OUTPUT_PATH.replace('.parquet', '.csv')}\")\n    \n    return df\n\nprint(\"🚀 Pipeline de pseudo-labeling prêt.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T13:44:35.398599Z","iopub.execute_input":"2026-01-20T13:44:35.398910Z","iopub.status.idle":"2026-01-20T13:44:35.420514Z","shell.execute_reply.started":"2026-01-20T13:44:35.398882Z","shell.execute_reply":"2026-01-20T13:44:35.419586Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# === EXÉCUTION ===\nif __name__ == \"__main__\":\n    pseudo_labels_df = run_pseudo_labeling()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T13:44:35.422536Z","iopub.execute_input":"2026-01-20T13:44:35.422803Z","iopub.status.idle":"2026-01-20T15:54:16.633127Z","shell.execute_reply.started":"2026-01-20T13:44:35.422777Z","shell.execute_reply":"2026-01-20T15:54:16.632129Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# === VISUALISATION DES RÉSULTATS ===\n\ndef visualize_pseudo_labels(df, top_k=10):\n    \"\"\"Affiche un aperçu des pseudo-labels générés.\"\"\"\n    \n    label_cols = [c for c in df.columns if c not in ['filename', 'second']]\n    \n    print(\"📊 Distribution des pseudo-labels:\\n\")\n    \n    # Top espèces détectées\n    species_counts = (df[label_cols] > 0).sum().sort_values(ascending=False)\n    print(f\"🐦 Top {top_k} espèces les plus détectées:\")\n    for species, count in species_counts.head(top_k).items():\n        pct = count / len(df) * 100\n        print(f\"   {species}: {count} secondes ({pct:.1f}%)\")\n    \n    # Moyenne des probabilités non-nulles par espèce\n    print(f\"\\n📈 Confiance moyenne (probs > 0):\")\n    for species in species_counts.head(top_k).index:\n        non_zero = df[df[species] > 0][species]\n        if len(non_zero) > 0:\n            print(f\"   {species}: {non_zero.mean():.3f} ± {non_zero.std():.3f}\")\n    \n    # Exemple d'un fichier\n    sample_file = df['filename'].iloc[0]\n    sample_data = df[df['filename'] == sample_file]\n    print(f\"\\n📝 Exemple: {sample_file}\")\n    print(f\"   Durée: {len(sample_data)} secondes\")\n    \n    active_species = sample_data[label_cols].sum()\n    active_species = active_species[active_species > 0].sort_values(ascending=False)\n    if len(active_species) > 0:\n        print(f\"   Espèces détectées: {list(active_species.head(5).index)}\")\n\n# Visualisation si le dataframe existe\ntry:\n    if pseudo_labels_df is not None:\n        visualize_pseudo_labels(pseudo_labels_df)\nexcept:\n    print(\"💡 Exécutez d'abord la cellule précédente pour générer les pseudo-labels.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T15:54:16.634442Z","iopub.execute_input":"2026-01-20T15:54:16.634797Z","iopub.status.idle":"2026-01-20T15:54:21.041341Z","shell.execute_reply.started":"2026-01-20T15:54:16.634767Z","shell.execute_reply":"2026-01-20T15:54:21.040450Z"}},"outputs":[],"execution_count":null}]}