{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"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":14515803,"sourceType":"datasetVersion","datasetId":9271149}],"dockerImageVersionId":31259,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport gc\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 torch.utils.data import Dataset, DataLoader\nfrom tqdm.auto import tqdm\nfrom sklearn.metrics import accuracy_score\n\n# Petit fix pour éviter les warnings audio sur certaines machines\ntry:\n    torchaudio.set_audio_backend(\"soundfile\")\nexcept:\n    pass\n\nprint(\"Bibliothèques chargées.\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-01-17T23:53:39.366715Z","iopub.execute_input":"2026-01-17T23:53:39.367039Z","iopub.status.idle":"2026-01-17T23:53:52.537517Z","shell.execute_reply.started":"2026-01-17T23:53:39.367004Z","shell.execute_reply":"2026-01-17T23:53:52.536725Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Config:\n    # --- CHEMINS KAGGLE ---\n    CSV_PATH = \"/kaggle/input/birdclef-2025/train.csv\"       \n    AUDIO_ROOT = \"/kaggle/input/birdclef-2025/train_audio\" \n    \n    # --- AUDIO ---\n    SR = 32000           \n    TARGET_LEN = 32000 * 5  # 5 secondes (segments longs pour insectes/amphibiens)\n    \n    # --- IMAGE (SPECTROGRAMME) ---\n    N_MELS = 224         # Résolution fréquentielle augmentée\n    IMG_SIZE = 224       # Taille standard pour MobileNetV3\n    \n    # --- ENTRAINEMENT (OPTIMISÉ T4) ---\n    BATCH_SIZE = 16      # Réduit car segments plus longs\n    EPOCHS = 7          # Plus d'epochs pour le teacher\n    LR = 1e-3            \n    MIXUP_ALPHA = 0.4    # Paramètre MixUp\n    \n    # --- MATÉRIEL ---\n    DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\nprint(f\"Configuration Teacher chargée. Device: {Config.DEVICE}\")\nprint(f\"Segments: {Config.TARGET_LEN // Config.SR}s | Mel bins: {Config.N_MELS}\")\nif Config.DEVICE.type == 'cpu':\n    print(\"⚠️ ATTENTION : TU ES SUR CPU. Active le GPU T4 dans les options Kaggle !\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-17T23:53:52.539153Z","iopub.execute_input":"2026-01-17T23:53:52.539603Z","iopub.status.idle":"2026-01-17T23:53:52.624495Z","shell.execute_reply.started":"2026-01-17T23:53:52.539577Z","shell.execute_reply":"2026-01-17T23:53:52.623799Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def fast_load_audio(filepath):\n    \"\"\"\n    Charge l'audio, resample et coupe/pad.\n    Renvoie un waveform brut (Tensor 1D).\n    \"\"\"\n    try:\n        # Chargement\n        wav, sr = torchaudio.load(filepath)\n        \n        # Resample si nécessaire\n        if sr != Config.SR:\n            wav = T.Resample(sr, Config.SR)(wav)\n            \n        # Gestion de la longueur (Padding ou Coupe)\n        channels, length = wav.shape\n        if length < Config.TARGET_LEN:\n            # Trop court : on ajoute du silence\n            pad = torch.zeros((channels, Config.TARGET_LEN - length))\n            wav = torch.cat((wav, pad), dim=1)\n        elif length > Config.TARGET_LEN:\n            # Trop long : coupe aléatoire\n            start = np.random.randint(0, length - Config.TARGET_LEN)\n            wav = wav[:, start:start + Config.TARGET_LEN]\n            \n        # On renvoie uniquement le premier canal (Mono)\n        return wav[0] \n\n    except Exception as e:\n        # En cas de fichier corrompu, on renvoie du silence pour ne pas crasher\n        return torch.zeros(Config.TARGET_LEN)\n\nprint(\"Fonction de chargement rapide prête.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-17T23:53:52.625550Z","iopub.execute_input":"2026-01-17T23:53:52.626068Z","iopub.status.idle":"2026-01-17T23:53:52.639373Z","shell.execute_reply.started":"2026-01-17T23:53:52.626040Z","shell.execute_reply":"2026-01-17T23:53:52.638717Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class GPUTransform(nn.Module):\n    \"\"\"Transformation Audio -> Spectrogramme sur GPU\"\"\"\n    def __init__(self):\n        super().__init__()\n        # Transformation Audio -> Spectrogramme Mel haute résolution\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,  # Plus de frames temporelles pour SED\n            f_min=20,\n            f_max=16000\n        )\n        self.to_db = T.AmplitudeToDB()\n        # Redimensionnement\n        self.resize = VT.Resize((Config.N_MELS, Config.IMG_SIZE))\n\n    def forward(self, wav_batch):\n        # 1. Création du Spectrogramme Mel\n        spec = self.mel_spec(wav_batch) \n        spec = self.to_db(spec)\n        \n        # 2. Ajout dimension Channel\n        spec = spec.unsqueeze(1) \n        \n        # 3. Resize\n        spec = self.resize(spec)\n        \n        # 4. Normalisation par batch (plus stable)\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        # 5. Mono -> RGB pour backbone pré-entraîné\n        spec = spec.repeat(1, 3, 1, 1) \n        \n        return spec\n\nprint(\"Module de transformation GPU (haute résolution) prêt.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-17T23:53:52.640444Z","iopub.execute_input":"2026-01-17T23:53:52.641078Z","iopub.status.idle":"2026-01-17T23:53:52.651077Z","shell.execute_reply.started":"2026-01-17T23:53:52.641036Z","shell.execute_reply":"2026-01-17T23:53:52.650527Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class BirdCLEFDataset(Dataset):\n    def __init__(self, df, root_dir):\n        self.df = df\n        self.root_dir = root_dir\n        self.labels = sorted(df['primary_label'].unique())\n        self.label_to_id = {label: i for i, label in enumerate(self.labels)}\n        self.num_classes = len(self.labels)\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        path = os.path.join(self.root_dir, row['filename'])\n        \n        # Charge l'audio brut\n        wav_tensor = fast_load_audio(path)\n        \n        # Label principal (one-hot pour compatibilité MixUp)\n        label = torch.zeros(self.num_classes)\n        label[self.label_to_id[row['primary_label']]] = 1.0\n        \n        # Labels secondaires si présents (multi-label)\n        if pd.notna(row.get('secondary_labels', None)) and row['secondary_labels'] != '[]':\n            try:\n                secondary = eval(row['secondary_labels'])\n                for sec_label in secondary:\n                    if sec_label in self.label_to_id:\n                        label[self.label_to_id[sec_label]] = 1.0\n            except:\n                pass\n        \n        return wav_tensor, label\n\n\n# Tête SED (Sound Event Detection)\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: (batch, features, time)\n        x = x.transpose(1, 2)  # (batch, time, features)\n        \n        # Attention weights\n        att_weights = torch.softmax(self.attention(x), dim=1)  # (batch, time, 1)\n        \n        # Frame-level predictions\n        frame_pred = self.fc(x)  # (batch, time, num_classes)\n        \n        # Clip-level prediction via attention pooling\n        clip_pred = (frame_pred * att_weights).sum(dim=1)  # (batch, num_classes)\n        \n        return clip_pred, frame_pred\n\n\n# Modèle Teacher avec SED\nclass BirdTeacherModel(nn.Module):\n    \"\"\"MobileNetV3 Large avec tête SED pour Teacher\"\"\"\n    def __init__(self, num_classes):\n        super().__init__()\n        # Backbone sans la tête de classification\n        self.backbone = timm.create_model(\n            #'mobilenetv3_large_100', \n            'efficientnet_b0',\n            pretrained=True, \n            num_classes=0,  # Pas de FC finale\n            global_pool=''  # Pas de pooling global\n        )\n        \n        # Récupère le nombre de features du backbone\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        # Tête SED\n        self.sed_head = SEDHead(self.num_features, num_classes)\n        \n    def forward(self, x):\n        # Extraction features (batch, channels, h, w)\n        features = self.backbone(x)\n        \n        # Pool sur la dimension fréquentielle, garde la dimension temporelle\n        features = features.mean(dim=2)  # (batch, channels, time)\n        \n        # Prédictions SED\n        clip_pred, frame_pred = self.sed_head(features)\n        \n        return clip_pred, frame_pred\n\nprint(\"Dataset multi-label et Architecture Teacher SED définis.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-17T23:53:52.651839Z","iopub.execute_input":"2026-01-17T23:53:52.652013Z","iopub.status.idle":"2026-01-17T23:53:52.664666Z","shell.execute_reply.started":"2026-01-17T23:53:52.651996Z","shell.execute_reply":"2026-01-17T23:53:52.664078Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def mixup_data(wavs, labels, alpha=0.4):\n    \"\"\"\n    MixUp augmentation sur les waveforms et labels.\n    Mélange aléatoire de deux échantillons avec coefficient lambda.\n    \"\"\"\n    if alpha > 0:\n        lam = np.random.beta(alpha, alpha)\n    else:\n        lam = 1.0\n    \n    batch_size = wavs.size(0)\n    index = torch.randperm(batch_size).to(wavs.device)\n    \n    mixed_wavs = lam * wavs + (1 - lam) * wavs[index]\n    mixed_labels = lam * labels + (1 - lam) * labels[index]\n    \n    return mixed_wavs, mixed_labels\n\n\ndef evaluate_model(model, loader, transform, device):\n    \"\"\"Évaluation du modèle Teacher sur le set de validation\"\"\"\n    model.eval()\n    all_preds = []\n    all_labels = []\n    \n    with torch.no_grad():\n        for wavs, labels in loader:\n            wavs = wavs.to(device)\n            labels = labels.to(device)\n            \n            # Transformation GPU\n            images = transform(wavs)\n            \n            # Prédiction (clip-level uniquement)\n            clip_pred, _ = model(images)\n            \n            # Top-1 prediction\n            predicted = torch.argmax(clip_pred, dim=1)\n            true_labels = torch.argmax(labels, dim=1)  # Labels one-hot -> indices\n            \n            all_preds.extend(predicted.cpu().numpy())\n            all_labels.extend(true_labels.cpu().numpy())\n            \n    return accuracy_score(all_labels, all_preds) * 100\n\nprint(\"Fonctions MixUp et évaluation prêtes.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-17T23:53:52.665640Z","iopub.execute_input":"2026-01-17T23:53:52.665945Z","iopub.status.idle":"2026-01-17T23:53:52.679898Z","shell.execute_reply.started":"2026-01-17T23:53:52.665925Z","shell.execute_reply":"2026-01-17T23:53:52.679408Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def run_teacher_training():\n    \"\"\"\n    Entraînement du modèle Teacher pour pseudo-labeling.\n    - Segments de 10s\n    - Spectrogrammes 224 bins\n    - MobileNetV3 + SED Head\n    - MixUp augmentation\n    - CrossEntropyLoss sans normalisation (multi-label friendly)\n    \"\"\"\n    # 1. Chargement des données\n    print(\"📂 Lecture du CSV...\")\n    df = pd.read_csv(Config.CSV_PATH)\n    \n    # --- FILTRE DE TEST ---\n    # Décommente pour tester rapidement\n    # df = df.sample(frac=1, random_state=42).reset_index(drop=True).head(1000) \n    \n    print(f\"📊 {len(df)} fichiers disponibles pour l'entraînement.\")\n    \n    full_ds = BirdCLEFDataset(df, Config.AUDIO_ROOT)\n    num_classes = full_ds.num_classes\n    print(f\"🐦 {num_classes} classes détectées.\")\n    \n    # Split 80% Train / 20% Val\n    train_size = int(0.8 * len(full_ds))\n    val_size = len(full_ds) - train_size\n    train_ds, val_ds = torch.utils.data.random_split(\n        full_ds, [train_size, val_size],\n        generator=torch.Generator().manual_seed(42)\n    )\n    print(f\"✂️ Split: {train_size} train / {val_size} val\")\n    \n    # 2. Dataloaders\n    train_loader = DataLoader(\n        train_ds, \n        batch_size=Config.BATCH_SIZE, \n        shuffle=True, \n        num_workers=4, \n        pin_memory=True,\n        persistent_workers=True,\n        drop_last=True  # Important pour MixUp\n    )\n    \n    val_loader = DataLoader(\n        val_ds, \n        batch_size=Config.BATCH_SIZE, \n        shuffle=False, \n        num_workers=4, \n        pin_memory=True,\n        persistent_workers=True\n    )\n    \n    # 3. Initialisation Modèle\n    model = BirdTeacherModel(num_classes).to(Config.DEVICE)\n    gpu_transform = GPUTransform().to(Config.DEVICE)\n    \n    optimizer = torch.optim.AdamW(model.parameters(), lr=Config.LR, weight_decay=0.01)\n    \n    # Scheduler cosine avec warmup\n    total_steps = len(train_loader) * Config.EPOCHS\n    scheduler = torch.optim.lr_scheduler.OneCycleLR(\n        optimizer, \n        max_lr=Config.LR,\n        total_steps=total_steps,\n        pct_start=0.1\n    )\n    \n    # CrossEntropyLoss - les labels ne sont PAS normalisés à 1\n    # Cela donne plus de poids aux fichiers multi-espèces\n    criterion = nn.CrossEntropyLoss()\n    \n    print(f\"🚀 Démarrage entraînement Teacher sur {Config.DEVICE}...\")\n    print(f\"   Segments: {Config.TARGET_LEN // Config.SR}s | MixUp α={Config.MIXUP_ALPHA}\")\n    \n    best_acc = 0.0\n    \n    # 4. Boucle d'entraînement\n    for epoch in range(Config.EPOCHS):\n        model.train()\n        train_loss = 0.0\n        \n        pbar = tqdm(train_loader, desc=f\"Epoch {epoch+1}/{Config.EPOCHS}\")\n        \n        for wavs, labels in pbar:\n            wavs = wavs.to(Config.DEVICE)\n            labels = labels.to(Config.DEVICE)\n            \n            # MixUp augmentation\n            mixed_wavs, mixed_labels = mixup_data(wavs, labels, Config.MIXUP_ALPHA)\n            \n            # Transformation GPU\n            with torch.no_grad():\n                images = gpu_transform(mixed_wavs)\n            \n            # Forward\n            clip_pred, frame_pred = model(images)\n            \n            # Loss sur clip-level (labels soft via MixUp)\n            loss = criterion(clip_pred, mixed_labels)\n            \n            # Backward\n            optimizer.zero_grad()\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n            optimizer.step()\n            scheduler.step()\n            \n            train_loss += loss.item()\n            pbar.set_postfix({'loss': f'{loss.item():.4f}'})\n            \n        # Validation\n        avg_loss = train_loss / len(train_loader)\n        val_acc = evaluate_model(model, val_loader, gpu_transform, Config.DEVICE)\n        \n        print(f\"📈 Epoch {epoch+1} | Loss: {avg_loss:.4f} | Val Acc: {val_acc:.2f}%\")\n        \n        # Sauvegarde du meilleur modèle\n        if val_acc > best_acc:\n            best_acc = val_acc\n            torch.save({\n                'epoch': epoch,\n                'model_state_dict': model.state_dict(),\n                'optimizer_state_dict': optimizer.state_dict(),\n                'val_acc': val_acc,\n                'num_classes': num_classes,\n                'labels': full_ds.labels\n            }, \"teacher_best.pth\")\n            print(f\"💾 Nouveau meilleur modèle sauvegardé! (Acc: {val_acc:.2f}%)\")\n        \n        # Sauvegarde checkpoint régulier\n        torch.save(model.state_dict(), f\"teacher_epoch_{epoch+1}.pth\")\n        \n        # Libération mémoire\n        gc.collect()\n        if Config.DEVICE.type == 'cuda':\n            torch.cuda.empty_cache()\n    \n    print(f\"\\n✅ Entraînement terminé! Meilleure accuracy: {best_acc:.2f}%\")\n    print(\"📁 Modèle Teacher sauvegardé: teacher_best.pth\")\n\nif __name__ == \"__main__\":\n    run_teacher_training()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-17T23:53:52.681618Z","iopub.execute_input":"2026-01-17T23:53:52.681819Z","iopub.status.idle":"2026-01-17T23:57:38.916853Z","shell.execute_reply.started":"2026-01-17T23:53:52.681801Z","shell.execute_reply":"2026-01-17T23:57:38.914744Z"}},"outputs":[],"execution_count":null}]}