{"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":292979834,"sourceType":"kernelVersion"},{"sourceId":723671,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":550665,"modelId":563292}],"dockerImageVersionId":31259,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# 🎓 Entraînement du Student Model (Données Réelles + Pseudo-Labels)\n\nCe notebook entraîne un modèle \"Student\" en combinant :\n1. **Données réelles** : Fichiers audio avec labels humains\n2. **Pseudo-labels** : Labels générés par le Teacher sur les soundscapes\n\n**Architecture** : MobileNetV3 Large + SED Head (identique au Teacher)","metadata":{}},{"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, ConcatDataset\nfrom tqdm.auto import tqdm\nfrom sklearn.metrics import accuracy_score\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-20T19:50:09.449866Z","iopub.execute_input":"2026-01-20T19:50:09.450396Z","iopub.status.idle":"2026-01-20T19:50:09.456316Z","shell.execute_reply.started":"2026-01-20T19:50:09.450362Z","shell.execute_reply":"2026-01-20T19:50:09.455488Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Config:\n    # --- CHEMINS KAGGLE ---\n    # Données réelles\n    CSV_PATH = \"/kaggle/input/birdclef-2025/train.csv\"       \n    AUDIO_ROOT = \"/kaggle/input/birdclef-2025/train_audio\"\n    \n    # Pseudo-labels\n    PSEUDO_LABELS_PATH = \"/kaggle/input/pseudolabeling/pseudo_labels.parquet\"  # Généré par pseudoLabeling.ipynb\n    SOUNDSCAPES_DIR = \"/kaggle/input/birdclef-2025/train_soundscapes\"\n    \n    # --- AUDIO ---\n    SR = 32000           \n    TARGET_LEN = 32000 * 5  # 5 secondes\n    \n    # --- IMAGE (SPECTROGRAMME) ---\n    N_MELS = 224\n    IMG_SIZE = 224\n    \n    # --- ENTRAINEMENT ---\n    BATCH_SIZE = 16\n    EPOCHS = 10           # Plus d'epochs pour le Student\n    LR = 0.001          # LR plus bas pour fine-tuning\n    MIXUP_ALPHA = 0.9\n    \n    # --- MIXTE (REAL vs PSEUDO) ---\n    REAL_WEIGHT = 1.0     # Poids des données réelles dans la loss\n    PSEUDO_WEIGHT = 0.6   # Poids des pseudo-labels (moins confiance)\n    \n    # --- MATÉRIEL ---\n    DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\nprint(f\"⚙️ Configuration Student chargée. Device: {Config.DEVICE}\")\nprint(f\"   Segments: {Config.TARGET_LEN // Config.SR}s | Mel bins: {Config.N_MELS}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T19:50:09.457757Z","iopub.execute_input":"2026-01-20T19:50:09.457971Z","iopub.status.idle":"2026-01-20T19:50:09.474486Z","shell.execute_reply.started":"2026-01-20T19:50:09.457949Z","shell.execute_reply":"2026-01-20T19:50:09.473894Z"}},"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        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        # Mono\n        if wav.shape[0] > 1:\n            wav = wav.mean(dim=0, keepdim=True)\n            \n        channels, length = wav.shape\n        \n        # Padding ou Coupe\n        if length < Config.TARGET_LEN:\n            pad = torch.zeros((channels, Config.TARGET_LEN - length))\n            wav = torch.cat((wav, pad), dim=1)\n        elif length > Config.TARGET_LEN:\n            start = np.random.randint(0, length - Config.TARGET_LEN)\n            wav = wav[:, start:start + Config.TARGET_LEN]\n            \n        return wav[0]\n\n    except Exception as e:\n        return torch.zeros(Config.TARGET_LEN)\n\nprint(\"🔧 Fonction de chargement audio prête.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T19:50:09.475449Z","iopub.execute_input":"2026-01-20T19:50:09.475829Z","iopub.status.idle":"2026-01-20T19:50:09.48758Z","shell.execute_reply.started":"2026-01-20T19:50:09.475795Z","shell.execute_reply":"2026-01-20T19:50:09.486923Z"}},"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        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\nprint(\"🖼️ Module de transformation GPU prêt.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T19:50:09.488554Z","iopub.execute_input":"2026-01-20T19:50:09.488925Z","iopub.status.idle":"2026-01-20T19:50:09.500524Z","shell.execute_reply.started":"2026-01-20T19:50:09.488893Z","shell.execute_reply":"2026-01-20T19:50:09.499981Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# === DATASETS ===\n\nclass RealDataset(Dataset):\n    \"\"\"Dataset pour les données réelles avec labels humains\"\"\"\n    def __init__(self, df, root_dir, label_list):\n        self.df = df\n        self.root_dir = root_dir\n        self.labels = label_list\n        self.label_to_id = {label: i for i, label in enumerate(self.labels)}\n        self.num_classes = len(self.labels)\n        self.is_pseudo = False  # Flag pour identifier le type\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        wav_tensor = fast_load_audio(path)\n        \n        # Label one-hot\n        label = torch.zeros(self.num_classes)\n        if row['primary_label'] in self.label_to_id:\n            label[self.label_to_id[row['primary_label']]] = 1.0\n        \n        # Labels secondaires\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        # Poids : données réelles ont plus de poids\n        weight = torch.tensor(Config.REAL_WEIGHT)\n        \n        return wav_tensor, label, weight\n\n\nclass PseudoLabelDataset(Dataset):\n    \"\"\"Dataset pour les pseudo-labels générés par le Teacher\"\"\"\n    def __init__(self, pseudo_df, soundscapes_dir, label_list):\n        self.pseudo_df = pseudo_df\n        self.soundscapes_dir = soundscapes_dir\n        self.labels = label_list\n        self.num_classes = len(self.labels)\n        self.is_pseudo = True\n        \n        # Grouper par fichier pour extraire des segments de 10s\n        self.file_groups = pseudo_df.groupby('filename')\n        self.files = list(self.file_groups.groups.keys())\n        \n        # Créer des indices (fichier, position de départ en secondes)\n        self.samples = []\n        segment_duration = Config.TARGET_LEN // Config.SR  # 10 secondes\n        \n        for filename in self.files:\n            file_data = self.file_groups.get_group(filename)\n            max_second = file_data['second'].max()\n            \n            # Créer des segments avec chevauchement\n            for start_sec in range(0, max_second - segment_duration + 1, segment_duration // 2):\n                self.samples.append((filename, start_sec))\n        \n    def __len__(self):\n        return len(self.samples)\n    \n    def __getitem__(self, idx):\n        filename, start_sec = self.samples[idx]\n        \n        # Charger l'audio\n        audio_path = None\n        for ext in ['.ogg', '.wav', '.mp3']:\n            candidate = os.path.join(self.soundscapes_dir, f\"{filename}{ext}\")\n            if os.path.exists(candidate):\n                audio_path = candidate\n                break\n        \n        if audio_path is None:\n            # Fichier non trouvé, retourner du silence\n            wav_tensor = torch.zeros(Config.TARGET_LEN)\n            label = torch.zeros(self.num_classes)\n            weight = torch.tensor(0.0)\n            return wav_tensor, label, weight\n        \n        # Charger et extraire le segment\n        try:\n            wav, sr = torchaudio.load(audio_path)\n            if sr != Config.SR:\n                wav = T.Resample(sr, Config.SR)(wav)\n            if wav.shape[0] > 1:\n                wav = wav.mean(dim=0, keepdim=True)\n            wav = wav[0]\n            \n            start_sample = start_sec * Config.SR\n            end_sample = start_sample + Config.TARGET_LEN\n            \n            if end_sample <= wav.shape[0]:\n                wav_tensor = wav[start_sample:end_sample]\n            else:\n                # Padding si nécessaire\n                segment = wav[start_sample:]\n                pad = torch.zeros(Config.TARGET_LEN - segment.shape[0])\n                wav_tensor = torch.cat([segment, pad])\n                \n        except Exception as e:\n            wav_tensor = torch.zeros(Config.TARGET_LEN)\n            label = torch.zeros(self.num_classes)\n            weight = torch.tensor(0.0)\n            return wav_tensor, label, weight\n        \n        # Récupérer les pseudo-labels pour ce segment (moyenne sur les secondes)\n        file_data = self.file_groups.get_group(filename)\n        segment_duration = Config.TARGET_LEN // Config.SR\n        end_sec = start_sec + segment_duration\n        \n        segment_data = file_data[(file_data['second'] >= start_sec) & (file_data['second'] < end_sec)]\n        \n        # Soft labels : moyenne des probabilités sur le segment\n        label = torch.zeros(self.num_classes)\n        for i, lbl in enumerate(self.labels):\n            if lbl in segment_data.columns:\n                label[i] = segment_data[lbl].mean()\n        \n        # Poids : pseudo-labels ont moins de poids\n        weight = torch.tensor(Config.PSEUDO_WEIGHT)\n        \n        return wav_tensor, label, weight\n\nprint(\"📁 Datasets (Real + Pseudo) définis.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T19:50:09.552898Z","iopub.execute_input":"2026-01-20T19:50:09.553187Z","iopub.status.idle":"2026-01-20T19:50:09.57064Z","shell.execute_reply.started":"2026-01-20T19:50:09.55316Z","shell.execute_reply":"2026-01-20T19:50:09.569937Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# === ARCHITECTURE MODÈLE ===\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 BirdStudentModel(nn.Module):\n    \"\"\"MobileNetV3 Large avec tête SED pour Student\"\"\"\n    def __init__(self, num_classes):\n        super().__init__()\n        self.backbone = timm.create_model(\n            #'mobilenetv3_large_100', \n            'efficientnet_b0',\n            pretrained=True, \n            num_classes=0,\n            global_pool=''\n        )\n        \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 Student Model définie.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T19:50:09.572137Z","iopub.execute_input":"2026-01-20T19:50:09.572404Z","iopub.status.idle":"2026-01-20T19:50:09.586136Z","shell.execute_reply.started":"2026-01-20T19:50:09.572381Z","shell.execute_reply":"2026-01-20T19:50:09.585538Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# === FONCTIONS D'ENTRAÎNEMENT ===\n\ndef mixup_data(wavs, labels, weights, alpha=0.4):\n    \"\"\"MixUp augmentation avec gestion des poids\"\"\"\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    mixed_weights = lam * weights + (1 - lam) * weights[index]\n    \n    return mixed_wavs, mixed_labels, mixed_weights\n\n\ndef weighted_cross_entropy(pred, target, weights):\n    \"\"\"CrossEntropyLoss pondérée par échantillon\"\"\"\n    # Softmax + log\n    log_probs = torch.log_softmax(pred, dim=1)\n    \n    # Cross entropy avec soft labels\n    loss_per_sample = -(target * log_probs).sum(dim=1)\n    \n    # Pondération\n    weighted_loss = (loss_per_sample * weights).mean()\n    \n    return weighted_loss\n\n\ndef evaluate_model(model, loader, transform, device):\n    \"\"\"Évaluation du modèle sur le set de validation\"\"\"\n    model.eval()\n    all_preds = []\n    all_labels = []\n    \n    with torch.no_grad():\n        for batch in loader:\n            wavs, labels = batch[0], batch[1]\n            wavs = wavs.to(device)\n            labels = labels.to(device)\n            \n            images = transform(wavs)\n            clip_pred, _ = model(images)\n            \n            predicted = torch.argmax(clip_pred, dim=1)\n            true_labels = torch.argmax(labels, dim=1)\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 d'entraînement prêtes.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T19:50:09.587007Z","iopub.execute_input":"2026-01-20T19:50:09.587284Z","iopub.status.idle":"2026-01-20T19:50:09.599812Z","shell.execute_reply.started":"2026-01-20T19:50:09.587244Z","shell.execute_reply":"2026-01-20T19:50:09.599085Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def prepare_datasets():\n    \"\"\"\n    Prépare les datasets combinés (réel + pseudo-labels).\n    Retourne les dataloaders et la liste des labels.\n    \"\"\"\n    print(\"📂 Préparation des datasets...\")\n    \n    # 1. Charger les données réelles\n    print(\"   → Chargement données réelles...\")\n    real_df = pd.read_csv(Config.CSV_PATH)\n    all_labels = sorted(real_df['primary_label'].unique())\n    num_classes = len(all_labels)\n    print(f\"   📊 {len(real_df)} fichiers réels | {num_classes} classes\")\n    \n    # 2. Charger les pseudo-labels (si disponibles)\n    pseudo_df = None\n    if os.path.exists(Config.PSEUDO_LABELS_PATH):\n        print(\"   → Chargement pseudo-labels...\")\n        pseudo_df = pd.read_parquet(Config.PSEUDO_LABELS_PATH)\n        print(f\"   📊 {pseudo_df['filename'].nunique()} fichiers pseudo-labellisés\")\n    else:\n        print(f\"   ⚠️ Pseudo-labels non trouvés: {Config.PSEUDO_LABELS_PATH}\")\n        print(\"   → Entraînement uniquement sur données réelles\")\n    \n    # 3. Split des données réelles (80/20)\n    real_train_df = real_df.sample(frac=0.8, random_state=42)\n    real_val_df = real_df.drop(real_train_df.index)\n    \n    print(f\"   ✂️ Split réel: {len(real_train_df)} train / {len(real_val_df)} val\")\n    \n    # 4. Créer les datasets\n    real_train_ds = RealDataset(real_train_df, Config.AUDIO_ROOT, all_labels)\n    real_val_ds = RealDataset(real_val_df, Config.AUDIO_ROOT, all_labels)\n    \n    # 5. Ajouter les pseudo-labels si disponibles\n    if pseudo_df is not None:\n        pseudo_ds = PseudoLabelDataset(pseudo_df, Config.SOUNDSCAPES_DIR, all_labels)\n        print(f\"   📊 {len(pseudo_ds)} segments pseudo-labellisés\")\n        \n        # Combiner real train + pseudo\n        train_ds = ConcatDataset([real_train_ds, pseudo_ds])\n        print(f\"   🔗 Dataset combiné: {len(train_ds)} échantillons\")\n    else:\n        train_ds = real_train_ds\n    \n    # 6. 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\n    )\n    \n    val_loader = DataLoader(\n        real_val_ds,  # Validation uniquement sur données réelles\n        batch_size=Config.BATCH_SIZE,\n        shuffle=False,\n        num_workers=4,\n        pin_memory=True,\n        persistent_workers=True\n    )\n    \n    return train_loader, val_loader, all_labels, num_classes\n\nprint(\"📦 Fonction de préparation des datasets prête.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T19:50:09.600729Z","iopub.execute_input":"2026-01-20T19:50:09.60107Z","iopub.status.idle":"2026-01-20T19:50:09.616166Z","shell.execute_reply.started":"2026-01-20T19:50:09.601038Z","shell.execute_reply":"2026-01-20T19:50:09.61544Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def run_student_training():\n    \"\"\"\n    Entraînement du modèle Student sur données réelles + pseudo-labels.\n    \"\"\"\n    print(\"=\" * 60)\n    print(\"🎓 ENTRAÎNEMENT DU STUDENT MODEL\")\n    print(\"=\" * 60)\n    \n    # 1. Préparation des données\n    train_loader, val_loader, labels, num_classes = prepare_datasets()\n    \n    # 2. Initialisation du modèle\n    model = BirdStudentModel(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 OneCycle\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    print(f\"\\n🚀 Démarrage entraînement sur {Config.DEVICE}...\")\n    print(f\"   Epochs: {Config.EPOCHS} | Batch: {Config.BATCH_SIZE}\")\n    print(f\"   Poids réel: {Config.REAL_WEIGHT} | Poids pseudo: {Config.PSEUDO_WEIGHT}\")\n    \n    best_acc = 0.0\n    \n    # 3. Boucle d'entraînement\n    for epoch in range(Config.EPOCHS):\n        model.train()\n        train_loss = 0.0\n        num_batches = 0\n        \n        pbar = tqdm(train_loader, desc=f\"Epoch {epoch+1}/{Config.EPOCHS}\")\n        \n        for batch in pbar:\n            wavs, labels_batch, weights = batch\n            wavs = wavs.to(Config.DEVICE)\n            labels_batch = labels_batch.to(Config.DEVICE)\n            weights = weights.to(Config.DEVICE)\n            \n            # MixUp augmentation\n            mixed_wavs, mixed_labels, mixed_weights = mixup_data(\n                wavs, labels_batch, weights, Config.MIXUP_ALPHA\n            )\n            \n            # Transformation GPU\n            with torch.no_grad():\n                images = gpu_transform(mixed_wavs)\n            \n            # Forward\n            clip_pred, _ = model(images)\n            \n            # Loss pondérée\n            loss = weighted_cross_entropy(clip_pred, mixed_labels, mixed_weights)\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            num_batches += 1\n            pbar.set_postfix({'loss': f'{loss.item():.4f}'})\n        \n        # Validation (sur données réelles uniquement)\n        avg_loss = train_loss / num_batches\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': labels\n            }, \"student_best.pth\")\n            print(f\"💾 Nouveau meilleur modèle! (Acc: {val_acc:.2f}%)\")\n        \n        # Checkpoint régulier\n        torch.save(model.state_dict(), f\"student_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 Student sauvegardé: student_best.pth\")\n    \n    return model\n\nprint(\"🚀 Pipeline d'entraînement Student prêt.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T19:50:09.617675Z","iopub.execute_input":"2026-01-20T19:50:09.617898Z","iopub.status.idle":"2026-01-20T19:50:09.634569Z","shell.execute_reply.started":"2026-01-20T19:50:09.61787Z","shell.execute_reply":"2026-01-20T19:50:09.633952Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# === EXÉCUTION ===\nif __name__ == \"__main__\":\n    student_model = run_student_training()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T19:50:09.635433Z","iopub.execute_input":"2026-01-20T19:50:09.636158Z","iopub.status.idle":"2026-01-20T19:52:52.256804Z","shell.execute_reply.started":"2026-01-20T19:50:09.636127Z","shell.execute_reply":"2026-01-20T19:52:52.255884Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# === COMPARAISON TEACHER vs STUDENT ===\n\ndef compare_models():\n    \"\"\"Compare les performances du Teacher et du Student\"\"\"\n    print(\"\\n📊 Comparaison Teacher vs Student\")\n    print(\"-\" * 40)\n    \n    results = {}\n    \n    # Charger les checkpoints\n    for model_name, path in [(\"Teacher\", \"teacher_best.pth\"), (\"Student\", \"student_best.pth\")]:\n        if os.path.exists(path):\n            checkpoint = torch.load(path, map_location='cpu')\n            acc = checkpoint.get('val_acc', 'N/A')\n            epoch = checkpoint.get('epoch', 'N/A')\n            results[model_name] = {'acc': acc, 'epoch': epoch}\n            print(f\"   {model_name}: Val Acc = {acc:.2f}% (epoch {epoch})\")\n        else:\n            print(f\"   {model_name}: Non trouvé ({path})\")\n    \n    if len(results) == 2:\n        diff = results['Student']['acc'] - results['Teacher']['acc']\n        if diff > 0:\n            print(f\"\\n   ✅ Le Student est meilleur de +{diff:.2f}%!\")\n        elif diff < 0:\n            print(f\"\\n   ⚠️ Le Teacher reste meilleur de {-diff:.2f}%\")\n        else:\n            print(f\"\\n   🔄 Performances identiques\")\n\n# Exécuter la comparaison si les modèles existent\ntry:\n    compare_models()\nexcept:\n    print(\"💡 Entraînez d'abord les modèles pour comparer.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T19:52:52.257335Z","iopub.status.idle":"2026-01-20T19:52:52.257571Z","shell.execute_reply.started":"2026-01-20T19:52:52.257458Z","shell.execute_reply":"2026-01-20T19:52:52.257474Z"}},"outputs":[],"execution_count":null}]}