{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":37077,"databundleVersionId":4333111,"sourceType":"competition"}],"dockerImageVersionId":31192,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport glob\nimport h5py\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.model_selection import StratifiedKFold\nfrom sklearn.metrics import roc_auc_score\nfrom tqdm import tqdm\nimport cv2  # OpenCV pour redimensionner les images rapidement\nimport timm # Bibliothèque de modèles pré-entrainés (ex: EfficientNet)\n\n# Configuration Globale\nCONFIG = {\n    'ROOT_DIR': '/kaggle/input/g2net-detecting-continuous-gravitational-waves',\n    'OUTPUT_DIR': './',\n    'IMG_SIZE': (256, 256), # Taille d'entrée du modèle\n    'BATCH_SIZE': 32,\n    'EPOCHS': 3,\n    'LR': 1e-3,\n    'SEED': 42,\n    'device': torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n}\n\n# Fixer la graine pour la reproductibilité (Déterminisme)\ndef seed_everything(seed):\n    \"\"\"Fixe la graine pour Numpy, PyTorch et Python pour la reproductibilité.\"\"\"\n    np.random.seed(seed)\n    # Graine Python standard\n    import random as rn\n    rn.seed(seed)\n    # Graine PyTorch\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    # Paramètres déterministes spécifiques à CUDA\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False \n\nseed_everything(CONFIG['SEED'])\nprint(f\"Device utilisé : {CONFIG['device']}\")\n\n# --- 1. Chargement et Analyse du Déséquilibre des Labels ---\n\n# 1. Chargement des labels\ntry:\n    train_labels = pd.read_csv(f\"{CONFIG['ROOT_DIR']}/train_labels.csv\")\n    print(\"\\n✅ train_labels.csv chargé.\")\nexcept FileNotFoundError:\n    print(f\"\\n🔴 Erreur: Le fichier des labels n'a pas été trouvé à {CONFIG['ROOT_DIR']}/train_labels.csv\")\n    train_labels = None\n\n# 2. Analyse du déséquilibre (s'il est chargé)\nif train_labels is not None:\n    print(f\"Nombre total d'échantillons : {len(train_labels)}\")\n    \n    # Calculer le nombre d'échantillons par classe\n    class_counts = train_labels['target'].value_counts()\n    \n    # Calculer les pourcentages\n    total_samples = len(train_labels)\n    neg_count = class_counts.get(0, 0)\n    pos_count = class_counts.get(1, 0)\n    \n    neg_percent = (neg_count / total_samples) * 100\n    pos_percent = (pos_count / total_samples) * 100\n    \n    print(\"\\n--- Déséquilibre des classes (Target) ---\")\n    print(f\"Classe 0 (Bruit/Négatif) : {neg_count} ({neg_percent:.2f}%)\")\n    print(f\"Classe 1 (Signal/Positif) : {pos_count} ({pos_percent:.2f}%)\")\n\n    # Mettre en évidence le déséquilibre\n    if pos_percent < 10:\n        print(\"\\n⚠️ Avertissement : Fort déséquilibre des classes. La classe positive est rare.\")\n        print(\"Cela justifie l'utilisation de la métrique AUC et de techniques anti-déséquilibre (Weighted Loss, Over/Under-sampling, Stratified K-Fold).\")\n    elif pos_percent < 30:\n        print(\"\\nNote : Déséquilibre modéré des classes. Des ajustements pourraient être nécessaires.\")\n    else:\n        print(\"\\nNote : Classes relativement bien équilibrées.\")\n    \n    # --- 3. Affichage graphique du déséquilibre ---\n    \n    plt.figure(figsize=(6, 4))\n    \n    bars = plt.bar(\n        class_counts.index.astype(str), \n        class_counts.values,            \n        color=['#1f77b4', '#ff7f0e']    \n    )\n    \n    total = sum(class_counts.values)\n    for bar in bars:\n        height = bar.get_height()\n        percentage = (height / total) * 100\n        plt.text(\n            bar.get_x() + bar.get_width() / 2., \n            height + 500, \n            f'{height}\\n({percentage:.2f}%)',\n            ha='center', \n            va='bottom',\n            fontsize=10\n        )\n\n    plt.title(\"Distribution des Classes G2Net (Target)\")\n    plt.xlabel(\"Classe (0: Bruit, 1: Signal CW)\")\n    plt.ylabel(\"Nombre d'Échantillons\")\n    plt.xticks([0, 1], ['Classe 0 (Bruit)', 'Classe 1 (Signal)'])\n    plt.grid(axis='y', linestyle='--', alpha=0.7)\n    \n    plt.show()\n\nprint(\"\\n--- Initialisation et analyse de la distribution des classes terminées. ---\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Installation des librairies pour la génération de données (GW physics)\n!pip install -q pyfstat lalsuite\n\nimport os\nimport h5py\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport pyfstat\nfrom tqdm.notebook import tqdm\n\n# Configuration de base\nROOT_DIR = '/kaggle/input/g2net-detecting-continuous-gravitational-waves'\nTRAIN_DIR = f\"{ROOT_DIR}/train\"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport h5py\nimport numpy as np\nimport pandas as pd\n# import matplotlib.pyplot as plt # Pas nécessaire pour cette cellule\n\n# Configuration (Assurez-vous que CONFIG est défini dans la première cellule)\n# Si vous exécutez cette cellule seule, vous aurez besoin de définir CONFIG ici:\n# CONFIG = {\n#     'ROOT_DIR': '/kaggle/input/g2net-detecting-continuous-gravitational-waves',\n#     'device': 'cpu'\n# } \n\ndef inspect_hdf5_structure_corrected(file_id, folder='train'):\n    \"\"\"\n    Inspecte la structure HDF5 en tenant compte du ID_FICHIER comme clé racine.\n    CORRIGÉ : Supprime l'accès à 'timestamps' pour éviter la KeyError.\n    \"\"\"\n    path = f\"{CONFIG['ROOT_DIR']}/{folder}/{file_id}.hdf5\"\n    \n    try:\n        with h5py.File(path, 'r') as f:\n            # 1. On trouve la clé racine (qui est l'ID du fichier)\n            root_keys = list(f.keys())\n            if not root_keys:\n                 print(\"🔴 Erreur: Fichier HDF5 vide.\")\n                 return\n                 \n            # Dans ce cas, la clé est l'ID du fichier lui-même\n            root_key = root_keys[0]\n            detector_group = f[root_key]\n\n            print(f\"--- Structure du fichier {file_id} ---\")\n            print(f\"Clé racine trouvée : '{root_key}'\") \n            print(f\"Clés disponibles sous '{root_key}' : {list(detector_group.keys())}\")\n            \n            # 2. On boucle sur les détecteurs H1 et L1 à l'intérieur de ce groupe\n            for detector in ['H1', 'L1']:\n                if detector in detector_group:\n                    print(f\"\\nDétecteur {detector}:\")\n                    \n                    # --- CORRECTION CRITIQUE : Accès UNIQUEMENT aux SFTs ---\n                    if 'SFTs' in detector_group[detector]:\n                        sfts = detector_group[detector]['SFTs'][:]\n                        \n                        print(f\"  Shape SFTs : {sfts.shape} (Fréquences x Temps)\")\n                        print(f\"  Type de données : {sfts.dtype} (Complexe)\") \n                        print(f\"  Dimensions : {sfts.shape[0]} bins de fréquence x {sfts.shape[1]} segments de temps\")\n                    else:\n                        print(f\"  🔴 Clé 'SFTs' manquante sous {detector}!\")\n\n                else:\n                     print(f\"  🔴 Détecteur {detector} manquant dans ce groupe!\")\n                     \n            if 'frequency_Hz' in detector_group:\n                 frequencies = detector_group['frequency_Hz'][:]\n                 print(f\"\\nFréquences : Shape {frequencies.shape}, de {frequencies[0]:.2f} Hz à {frequencies[-1]:.2f} Hz\")\n            \n    except Exception as e:\n        print(f\"🔴 Erreur lors de l'ouverture ou la lecture de {file_id}.hdf5: {e}\")\n\n# Test avec un fichier aléatoire du train set\n# Assurez-vous que CONFIG est bien défini (avec ROOT_DIR)\n# et que train_labels.csv est accessible\ntry:\n    train_labels = pd.read_csv(f\"{CONFIG['ROOT_DIR']}/train_labels.csv\")\n    sample_id_train = train_labels.iloc[0]['id'] \n    \n    # Appel de la fonction corrigée\n    inspect_hdf5_structure_corrected(sample_id_train)\n    \nexcept NameError:\n    print(\"🔴 Erreur: La variable CONFIG n'est pas définie. Veuillez exécuter la cellule de configuration initiale.\")\nexcept FileNotFoundError:\n    print(f\"🔴 Erreur: Fichier de labels non trouvé à {CONFIG['ROOT_DIR']}/train_labels.csv.\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# INSTALLATION / SETUP / DATASET AVEC CANAL 4 OPTIMISÉ\n# ============================================================\n\nimport os\nimport numpy as np\nimport random as rn\nimport torch\nimport torch.backends.cudnn\nfrom torch.utils.data import Dataset\nimport h5py\nimport cv2\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom scipy.ndimage import gaussian_filter1d\n\n# ============================================================\n# CONFIG\n# ============================================================\n\nCONFIG = {\n    'ROOT_DIR': '/kaggle/input/g2net-detecting-continuous-gravitational-waves',\n    'IMG_SIZE': (256, 256),\n    'BATCH_SIZE': 32,\n    'SEED': 42,\n    'device': torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n}\n\n# ============================================================\n# DÉTERMINISME MAXIMAL\n# ============================================================\n\ndef seed_everything_strict(seed):\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    rn.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n    os.environ['TF_DETERMINISTIC_OPS'] = '1'\n\nseed_everything_strict(CONFIG['SEED'])\nprint(f\"✅ Déterminisme maximal activé avec SEED = {CONFIG['SEED']}\")\n\n\n# ============================================================\n# 🔥 DOPPLER PROXY (remplace pyfstat)\n# ============================================================\n\ndef make_doppler_proxy(freq, ra, dec, timestamps):\n    \"\"\"\n    Approximation réaliste du Doppler CW sans pyfstat.\n    Rotation terrestre + orbite.\n    \"\"\"\n    earth_rot = 2 * np.pi / 86164                   # 1 jour sidéral\n    earth_orbit = 2 * np.pi / (365.25 * 86400)      # orbite Terre\n\n    doppler_rot = 1e-6 * np.sin(earth_rot * timestamps + ra)\n    doppler_orb = 5e-4 * np.cos(earth_orbit * timestamps + dec)\n\n    doppler = freq * (doppler_rot + doppler_orb)\n\n    # Normalisation\n    return ((doppler - doppler.mean()) / (doppler.std() + 1e-8)).astype(np.float32)\n\n\n# ============================================================\n# 🔥 VERSION OPTIMISÉE DU MASQUE 2D DOPPLER\n# ============================================================\n\ndef make_doppler_mask_optimized(doppler_curve, freq_bins, smooth=4, thickness=4):\n    \"\"\"\n    Transforme la courbe Doppler 1D en une carte 2D optimisée :\n    - Lissage gaussien\n    - Ligne épaisse\n    - Atténuation gaussienne autour de la crête\n    \"\"\"\n    # 1) Normalisation\n    doppler_norm = (doppler_curve - doppler_curve.min()) / (doppler_curve.ptp() + 1e-8)\n    \n    # 2) Indices de base\n    base_idx = doppler_norm * (freq_bins - 1)\n    \n    # 3) Lissage\n    smooth_idx = gaussian_filter1d(base_idx, sigma=smooth)\n    \n    # 4) Masque final\n    mask = np.zeros((freq_bins, len(doppler_curve)), dtype=np.float32)\n    \n    for t in range(len(doppler_curve)):\n        f0 = smooth_idx[t]\n        for off in range(-thickness, thickness + 1):\n            f = int(np.clip(f0 + off, 0, freq_bins - 1))\n            weight = np.exp(-(off ** 2) / (2 * (thickness / 2) ** 2))\n            mask[f, t] = max(mask[f, t], weight)\n\n    # 5) Re-normalisation\n    return (mask / (mask.max() + 1e-8)).astype(np.float32)\n\n\n# ============================================================\n# 🔥 DATASET PHYSIQUE G2NET (4 CANAUX) - CORRIGÉ\n# ============================================================\n\nclass G2NetDatasetPhysics(Dataset):\n    # ⭐️ CORRECTION ICI : Ajout du paramètre 'folder' ⭐️\n    def __init__(self, df, root_dir, img_size, folder='train', transforms=None):\n        self.df = df\n        self.root_dir = root_dir\n        self.img_size = img_size\n        self.transforms = transforms\n        self.folder = folder # Stockage du paramètre 'folder'\n        \n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        file_id = row['id']\n        target = row['target']\n        \n        # ⭐️ CORRECTION ICI : Utilisation de self.folder ⭐️\n        path = os.path.join(\n            self.root_dir, self.folder, file_id[0], file_id[1], file_id[2], f'{file_id}.hdf5'\n        )\n\n        # Sources CW (fixes ici)\n        F0, ALPHA, DELTA = 100.0, 1.0, 0.5\n\n        # === Chargement HDF5 ===\n        try:\n            with h5py.File(path, 'r') as f:\n                root_key = list(f.keys())[0]\n                sgram_h1 = np.abs(f[f'{root_key}/H1/SFTs'][:])\n                sgram_l1 = np.abs(f[f'{root_key}/L1/SFTs'][:])\n                timestamps = f[f'{root_key}/H1/timestamps_GPS'][:] \n        except:\n            sgram_h1 = np.zeros((360, 4096), np.float32)\n            sgram_l1 = np.zeros((360, 4096), np.float32)\n            timestamps = np.arange(4096)\n\n        # Alignement\n        T = min(sgram_h1.shape[1], sgram_l1.shape[1])\n        sgram_h1 = sgram_h1[:, :T]\n        sgram_l1 = sgram_l1[:, :T]\n        timestamps = timestamps[:T]\n\n        # Normalisation Z-score\n        def z(s): return (s - s.mean()) / (s.std() + 1e-6)\n\n        s_h1 = z(sgram_h1)\n        s_l1 = z(sgram_l1)\n\n        # === CANAL 4 : DOPPLER OPTIMISÉ ===\n        doppler_curve = make_doppler_proxy(F0, ALPHA, DELTA, timestamps)\n        doppler_mask = make_doppler_mask_optimized(\n            doppler_curve,\n            freq_bins=s_h1.shape[0],\n            smooth=4,\n            thickness=4\n        )\n\n        # Stack final 4 canaux\n        combined = np.stack(\n            [s_h1, s_l1, (s_h1 + s_l1) / 2, doppler_mask],\n            axis=-1\n        )\n\n        # Resize final\n        img = cv2.resize(combined, CONFIG['IMG_SIZE'], interpolation=cv2.INTER_LINEAR)\n        img = torch.from_numpy(img).permute(2, 0, 1).float()\n\n        return img, torch.tensor(target, dtype=torch.float)\n\n\n# ============================================================\n# 🔥 TEST\n# ============================================================\n\ntrain_labels = pd.read_csv(f\"{CONFIG['ROOT_DIR']}/train_labels.csv\")\n\n# NOTE: Le test ci-dessous fonctionne sans le paramètre 'folder' car il a la valeur par défaut 'train'\ndataset = G2NetDatasetPhysics(\n    df=train_labels,\n    root_dir=CONFIG['ROOT_DIR'],\n    img_size=CONFIG['IMG_SIZE']\n)\n\nimg, y = dataset[0]\n\nprint(\"Image shape :\", img.shape)\nprint(\"Target :\", y.item())\n\nplt.figure(figsize=(10,3))\nplt.imshow(img[3].cpu(), aspect='auto', origin='lower', cmap='inferno')\nplt.title(\"Canal 4 : Doppler Optimisé\")\nplt.colorbar()\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import DataLoader\nimport timm \nimport pandas as pd\nfrom sklearn.model_selection import StratifiedKFold\nfrom sklearn.metrics import roc_auc_score\nfrom tqdm import tqdm\nimport os\n\n# --- 0. Dépendances et Configuration (Assumer définies précédemment) ---\n\nCONFIG = {\n    'ROOT_DIR': '/kaggle/input/g2net-detecting-continuous-gravitational-waves',\n    'IMG_SIZE': (256, 256),\n    'BATCH_SIZE': 32,\n    'EPOCHS': 3,\n    'SEED': 42,\n    'device': torch.device('cuda' if torch.cuda.is_available() else 'cpu'),\n    'LR': 1e-4, # Taux d'apprentissage\n    'N_SPLITS': 5 # Nombre de plis pour la CV\n}\n\n# Chargement du DataFrame\ntry:\n    train_labels = pd.read_csv(f\"{CONFIG['ROOT_DIR']}/train_labels.csv\")\nexcept NameError:\n    print(\"🔴 Erreur: Assurez-vous que 'CONFIG' est défini et que le fichier 'train_labels.csv' est chargé.\")\n    exit()\n\n# NOTE IMPORTANTE : La classe G2NetDatasetPhysics est désormais la classe utilisée.\n\nclass CWModel(nn.Module):\n    def __init__(self, model_name='tf_efficientnet_b0_ns', pretrained=True, in_chans=4): \n        super().__init__()\n        self.backbone = timm.create_model(model_name, pretrained=pretrained, in_chans=in_chans)\n        \n        if hasattr(self.backbone, 'classifier'):\n            n_features = self.backbone.classifier.in_features\n            self.backbone.classifier = nn.Linear(n_features, 1)\n        elif hasattr(self.backbone, 'fc'):\n            n_features = self.backbone.fc.in_features\n            self.backbone.fc = nn.Linear(n_features, 1)\n\n    def forward(self, x):\n        return self.backbone(x)\n\ndef train_one_epoch(model, loader, optimizer, criterion, device):\n    model.train()\n    total_loss = 0\n    for img, target in tqdm(loader, desc=\"Train\"):\n        img, target = img.to(device), target.to(device)\n        optimizer.zero_grad()\n        output = model(img).squeeze(1)\n        loss = criterion(output, target)\n        loss.backward()\n        optimizer.step()\n        total_loss += loss.item()\n    return total_loss / len(loader)\n\n# =========================================================================\n# 🔥 FONCTION VALIDATE CORRIGÉE ET ULTRA-ROBUSTE\n# =========================================================================\n\ndef validate(model, loader, device):\n    model.eval()\n    preds = []\n    targets = []\n\n    with torch.no_grad():\n        for img, target in tqdm(loader, desc=\"Valid\"):\n            img = img.to(device)\n            output = torch.sigmoid(model(img)).squeeze(1)\n\n            # Conversion directe des targets en entier pour s'assurer du format binaire\n            preds.extend(output.cpu().numpy().astype(float))\n            targets.extend(target.numpy().astype(int)) # ⭐️ CORRECTION : FORCÉ EN INT BINAIRE\n\n    targets = np.array(targets)\n    preds = np.array(preds)\n\n    # ===⭐️ SÉCURITÉ ULTIME 1 : CAS PLI À UNE SEULE CLASSE ⭐️===\n    if len(np.unique(targets)) < 2:\n        # Comme l'exige la compétition (valeur minimale 0.5 si AUC non calculable)\n        # Note: Rappel de votre information sauvegardée: Le minimum de score requis pour ce type de compétition est de 0,5.\n        print(\"⚠️ Avertissement : Une seule classe trouvée dans ce pli. AUC = 0.5\")\n        return 0.5\n\n    # ===⭐️ SÉCURITÉ ULTIME 2 : Gestion des autres erreurs AUC ⭐️===\n    try:\n        return roc_auc_score(targets, preds)\n    except:\n        print(\"⚠️ Erreur AUC imprévue (y_true mal formé ou autre) → retour 0.5\")\n        return 0.5\n\n\n# -------------------------------------------------------------------\n## 🚀 1. Fonction d'Entraînement avec Cross-Validation (CV)\n# -------------------------------------------------------------------\n\ndef run_cross_validation_train(df):\n    \"\"\"\n    Exécute l'entraînement complet en utilisant la Validation Croisée Stratifiée.\n    \"\"\"\n    all_fold_scores = []\n    \n    # 1. Initialisation du Stratified K-Fold\n    # Le UserWarning sur les petites classes est inévitable sur ce dataset.\n    skf = StratifiedKFold(n_splits=CONFIG['N_SPLITS'], shuffle=True, random_state=CONFIG['SEED'])\n    \n    # 2. Itération sur les plis\n    for fold, (train_idx, val_idx) in enumerate(skf.split(df, df['target'])):\n        print(f\"\\n====================== DÉBUT DU PLI {fold+1}/{CONFIG['N_SPLITS']} ======================\")\n        \n        # Création des DataFrames pour le pli actuel\n        train_df = df.iloc[train_idx].reset_index(drop=True)\n        val_df = df.iloc[val_idx].reset_index(drop=True)\n        \n        # 3. Création des Datasets et DataLoaders (Utilisation de la classe 4 canaux)\n        # Assurez-vous que G2NetDatasetPhysics (Cellule 1/3) accepte 'folder' !\n        train_ds = G2NetDatasetPhysics(\n            train_df, \n            root_dir=CONFIG['ROOT_DIR'], \n            img_size=CONFIG['IMG_SIZE'],\n            folder='train' \n        )\n        val_ds = G2NetDatasetPhysics(\n            val_df, \n            root_dir=CONFIG['ROOT_DIR'], \n            img_size=CONFIG['IMG_SIZE'],\n            folder='train'\n        )\n        \n        train_loader = DataLoader(\n            train_ds, batch_size=CONFIG['BATCH_SIZE'], shuffle=True, num_workers=2\n        )\n        val_loader = DataLoader(\n            val_ds, batch_size=CONFIG['BATCH_SIZE'], shuffle=False, num_workers=2\n        )\n        \n        # 4. Initialisation du Modèle/Optimiseur (in_chans=4 CORRECT)\n        model = CWModel(in_chans=4).to(CONFIG['device'])\n        optimizer = torch.optim.Adam(model.parameters(), lr=CONFIG['LR'])\n        criterion = nn.BCEWithLogitsLoss()\n        \n        best_fold_score = 0\n        \n        # 5. Boucle d'Époques\n        for epoch in range(CONFIG['EPOCHS']):\n            loss = train_one_epoch(model, train_loader, optimizer, criterion, CONFIG['device'])\n            score = validate(model, val_loader, CONFIG['device']) # <-- Appel à la fonction ultra-robuste\n            \n            print(f\"Pli {fold+1} | Epoch {epoch+1}/{CONFIG['EPOCHS']} | Loss: {loss:.4f} | AUC: {score:.4f}\")\n            \n            if score > best_fold_score:\n                best_fold_score = score\n                # Sauvegarde du modèle pour ce pli\n                torch.save(model.state_dict(), f'best_model_fold_{fold+1}.pth')\n                print(f\"  >>> Modèle du pli {fold+1} sauvegardé (AUC: {best_fold_score:.4f})\")\n                \n        all_fold_scores.append(best_fold_score)\n        print(f\"====================== FIN DU PLI {fold+1} ======================\")\n\n    # 6. Rapport Final\n    mean_auc = np.mean(all_fold_scores)\n    std_auc = np.std(all_fold_scores)\n    print(\"\\n\\n#####################################################\")\n    print(f\"RÉSULTAT FINAL (Cross-Validation sur {CONFIG['N_SPLITS']} plis):\")\n    print(f\"  Moyenne AUC: {mean_auc:.4f} ± {std_auc:.4f}\")\n    print(\"#####################################################\")\n    return all_fold_scores\n\n# --- EXÉCUTION ---\nif __name__ == '__main__':\n    # Rappel : Assurez-vous que G2NetDatasetPhysics est la classe définie dans la cellule précédente\n    final_scores = run_cross_validation_train(train_labels)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport torch\nfrom torch.utils.data import DataLoader\nimport os\nfrom tqdm import tqdm\n\n# --- 0. Configuration et Dépendances (Assumer définies) ---\n# CONFIG doit contenir: ROOT_DIR, BATCH_SIZE, device, N_SPLITS\n# La classe CWModel et G2NetDatasetPhysics DOIVENT être définies dans les cellules précédentes.\nCONFIG = {\n    'ROOT_DIR': '/kaggle/input/g2net-detecting-continuous-gravitational-waves',\n    'IMG_SIZE': (256, 256),\n    'BATCH_SIZE': 64,  # Généralement plus grand pour l'inférence\n    'device': torch.device('cuda' if torch.cuda.is_available() else 'cpu'),\n    'N_SPLITS': 5 \n}\n\n# Charger le fichier de soumission pour obtenir les IDs du set de test\ntry:\n    submission_df = pd.read_csv(f\"{CONFIG['ROOT_DIR']}/sample_submission.csv\")\n    print(\"\\n✅ Fichier de soumission/test IDs chargé.\")\nexcept FileNotFoundError:\n    print(\"🔴 Erreur: 'sample_submission.csv' non trouvé. Veuillez vérifier le chemin.\")\n    exit()\n\n# -------------------------------------------------------------------\n## 💻 1. Préparation du DataLoader de Test\n# -------------------------------------------------------------------\n\n# ⭐️ CORRECTION CRITIQUE: Utilisation de G2NetDatasetPhysics pour le 4 canaux ⭐️\ntest_ds = G2NetDatasetPhysics(\n    submission_df, \n    root_dir=CONFIG['ROOT_DIR'], \n    img_size=CONFIG['IMG_SIZE'], \n    folder='test' # Assurez-vous que G2NetDatasetPhysics utilise bien 'folder' ou déduit le chemin de test.\n)\ntest_loader = DataLoader(\n    test_ds, \n    batch_size=CONFIG['BATCH_SIZE'], \n    shuffle=False, \n    num_workers=2\n)\nprint(f\"Dataset de test créé : {len(test_ds)} échantillons.\")\n\n# -------------------------------------------------------------------\n## 🚀 2. Inférence (Prédiction) par Ensembling K-Fold\n# -------------------------------------------------------------------\n\ndef predict_by_kfold_ensemble(test_loader, n_splits):\n    \"\"\"\n    Fait des prédictions en chargeant chaque modèle de pli (fold) et en moyennant les résultats.\n    \"\"\"\n    final_predictions = np.zeros((len(test_loader.dataset),))\n    \n    # Itération sur chaque pli sauvegardé\n    for fold in range(1, n_splits + 1):\n        print(f\"\\n--- Prédiction avec le modèle du pli {fold} ---\")\n        \n        # 1. Charger un nouveau modèle\n        # CWModel doit être défini pour in_chans=4\n        model = CWModel(in_chans=4).to(CONFIG['device'])\n        weights_path = f'best_model_fold_{fold}.pth'\n        \n        # 2. Charger les poids spécifiques à ce pli\n        try:\n            model.load_state_dict(torch.load(weights_path, map_location=CONFIG['device']))\n            print(f\"✅ Poids chargés depuis {weights_path}\")\n        except FileNotFoundError:\n            print(f\"🔴 ERREUR: Fichier de poids {weights_path} non trouvé. Passe au pli suivant.\")\n            continue\n            \n        model.eval()\n        fold_predictions = []\n        \n        # 3. Prédiction sur le set de test\n        with torch.no_grad():\n            for img, _ in tqdm(test_loader, desc=f\"Inférence Pli {fold}\"):\n                img = img.to(CONFIG['device'])\n                \n                # Le modèle renvoie des logits, nous appliquons Sigmoid pour la probabilité [0, 1]\n                output = torch.sigmoid(model(img)).squeeze(1) \n                fold_predictions.extend(output.cpu().numpy())\n        \n        # 4. Cumul des prédictions (Moyenne)\n        final_predictions += np.array(fold_predictions)\n\n    # 5. Calcul de la moyenne sur tous les plis\n    final_predictions /= CONFIG['N_SPLITS']\n    return final_predictions\n\n# Exécution de l'inférence\ntest_predictions = predict_by_kfold_ensemble(test_loader, CONFIG['N_SPLITS'])\n\n# -------------------------------------------------------------------\n## 📝 3. Génération du Fichier de Soumission\n# -------------------------------------------------------------------\n\nsubmission_df['target'] = test_predictions\n\n# Vérification du format\nprint(f\"\\n--- Vérification du Fichier de Soumission ---\")\nprint(f\"Nombre de lignes : {len(submission_df)}\")\nprint(\"Aperçu du DataFrame:\")\nprint(submission_df.head())\n\n# Sauvegarde au format requis\nsubmission_file_name = 'submission.csv'\nsubmission_df.to_csv(submission_file_name, index=False)\n\nprint(f\"\\n✅ Fichier de soumission généré avec succès : {submission_file_name}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}