{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.12.12"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":6538,"databundleVersionId":44214,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":18613,"sourceType":"datasetVersion","datasetId":5839},{"sourceId":3270629,"sourceType":"datasetVersion","datasetId":1981237},{"sourceId":746564,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":570024,"modelId":582304},{"sourceId":746566,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":570026,"modelId":582306}],"dockerImageVersionId":31260,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Pl@ntNet-300K — Classificateur en Cascade\n\n**Pipeline complet** : chargement → prétraitement/cache → analyse → entraînement → évaluation\n\n> Garcin et al., \"Pl@ntNet-300K: a plant image dataset with high label ambiguity and a long-tailed distribution\", NeurIPS 2021","metadata":{}},{"cell_type":"markdown","source":"## 1. Imports & Configuration","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5"}},{"cell_type":"code","source":"# ── Imports ──────────────────────────────────────────────────\nimport json, os, math, copy, random, time\nfrom datetime import datetime\nfrom collections import Counter, defaultdict\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.nn.functional as F\nfrom torch.utils.data import DataLoader, Subset, TensorDataset\nfrom torch.cuda.amp import autocast, GradScaler\n\nimport torchvision.datasets as datasets\nimport torchvision.transforms as transforms\nimport torchvision.models as models\nimport timm\n\nfrom sklearn.metrics import (\n    classification_report, confusion_matrix, ConfusionMatrixDisplay,\n    f1_score, precision_score, recall_score, roc_auc_score\n)\nfrom tqdm.auto import tqdm\n\n# ── Config GPU ──────────────────────────────────────────────\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nif DEVICE == \"cuda\":\n    torch.backends.cudnn.benchmark = True\n    torch.backends.cuda.matmul.allow_tf32 = True\n    torch.backends.cudnn.allow_tf32 = True\n    print(f\"GPU: {torch.cuda.get_device_name(0)} \"\n          f\"({torch.cuda.get_device_properties(0).total_memory / 1e9:.1f} GB)\")\nelse:\n    print(\"Pas de GPU — entraînement CPU (lent)\")\n\n# ── Chemins Kaggle ──────────────────────────────────────────\nROOT_DIR       = \"/kaggle/input/plantnet-300k-images/plantnet_300K\"\nTRAIN_DIR      = f\"{ROOT_DIR}/images_train\"\nVAL_DIR        = f\"{ROOT_DIR}/images_val\"\nJSON_NAMES     = f\"{ROOT_DIR}/plantnet300K_species_names.json\"\nJSON_META      = f\"{ROOT_DIR}/plantnet300K_metadata.json\"\nSAVE_DIR       = \"./cascade_models\"\nos.makedirs(SAVE_DIR, exist_ok=True)\n\n# ── Hyperparamètres ─────────────────────────────────────────\nTOP_N_DOMINANT      = 5       # Nombre de classes dominantes pour la cascade\nSEUIL_TRES_RARE     = 5       # < 5 images → impossible à apprendre\nSEUIL_RARE          = 20      # < 20 images → difficile\nNUM_WORKERS         = min(8, os.cpu_count() or 4)\nIMAGENET_MEAN       = [0.485, 0.456, 0.406]\nIMAGENET_STD        = [0.229, 0.224, 0.225]\n\nprint(f\"Device: {DEVICE} | Workers: {NUM_WORKERS}\")","metadata":{"execution":{"iopub.status.busy":"2026-02-10T19:34:48.113602Z","iopub.execute_input":"2026-02-10T19:34:48.11394Z","iopub.status.idle":"2026-02-10T19:34:48.123463Z","shell.execute_reply.started":"2026-02-10T19:34:48.113913Z","shell.execute_reply":"2026-02-10T19:34:48.122764Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 2. Fonctions utilitaires (métriques, poids, cache)","metadata":{"trusted":true}},{"cell_type":"code","source":"# ── Coefficient de Gini ──────────────────────────────────────\ndef gini_coefficient(array):\n    \"\"\"0 = égalité parfaite, 1 = inégalité totale.\"\"\"\n    a = np.sort(np.asarray(array, dtype=float))\n    n = len(a)\n    idx = np.arange(1, n + 1)\n    return ((2 * idx - n - 1) * a).sum() / (n * a.sum())\n\n\n# ── Poids de classes ────────────────────────────────────────\ndef compute_class_weights(dataset, num_classes, device=DEVICE, smoothing=0.1):\n    \"\"\"Poids inversement proportionnels à la fréquence, bornés [0.1, 10].\"\"\"\n    if hasattr(dataset, \"targets\"):\n        targets = np.array(dataset.targets)\n    elif hasattr(dataset, \"dataset\") and hasattr(dataset, \"indices\"):\n        targets = np.array([dataset.dataset.targets[i] for i in dataset.indices])\n    elif hasattr(dataset, \"dataset\"):\n        targets = np.array([dataset.dataset.targets[i] for i in dataset.indices])\n    else:\n        targets = np.array([lbl for _, lbl in tqdm(dataset, desc=\"Lecture labels\")])\n\n    counts = np.bincount(targets, minlength=num_classes).clip(min=1)\n    w = (len(targets) / (num_classes * counts)) ** smoothing\n    w = np.clip(w / w.mean(), 0.1, 10.0)\n    print(f\"Class weights — min: {w.min():.3f}, max: {w.max():.3f}\")\n    return torch.FloatTensor(w).to(device)\n\n\n# ── Calcul de TOUTES les métriques ──────────────────────────\ndef compute_all_metrics(y_true, y_pred, y_scores=None, prefix=\"\"):\n    \"\"\"\n    Renvoie un dict avec accuracy, precision, recall, F1 (macro),\n    weighted-F1 et AUROC (si y_scores fourni).\n    \"\"\"\n    acc = 100 * (np.asarray(y_true) == np.asarray(y_pred)).mean()\n    prec  = precision_score(y_true, y_pred, average=\"macro\", zero_division=0)\n    rec   = recall_score(y_true, y_pred, average=\"macro\", zero_division=0)\n    f1_m  = f1_score(y_true, y_pred, average=\"macro\", zero_division=0)\n    f1_w  = f1_score(y_true, y_pred, average=\"weighted\", zero_division=0)\n\n    metrics = {\n        f\"{prefix}accuracy\": acc,\n        f\"{prefix}precision_macro\": prec,\n        f\"{prefix}recall_macro\": rec,\n        f\"{prefix}f1_macro\": f1_m,\n        f\"{prefix}f1_weighted\": f1_w,\n    }\n\n    if y_scores is not None:\n        try:\n            auroc = roc_auc_score(\n                y_true, y_scores, multi_class=\"ovr\", average=\"weighted\"\n            )\n            metrics[f\"{prefix}auroc\"] = auroc\n        except Exception:\n            metrics[f\"{prefix}auroc\"] = float(\"nan\")\n\n    return metrics\n\n\ndef print_metrics(m, title=\"Métriques\"):\n    \"\"\"Affiche joliment un dict de métriques.\"\"\"\n    print(f\"\\n{'─'*50}\")\n    print(f\"  {title}\")\n    print(f\"{'─'*50}\")\n    for k, v in m.items():\n        if \"accuracy\" in k:\n            print(f\"  {k:<25s} {v:>8.2f} %\")\n        else:\n            print(f\"  {k:<25s} {v:>8.4f}\")\n    print(f\"{'─'*50}\")\n\n\n# ── Subset stratifié robuste ────────────────────────────────\ndef stratified_subset(dataset, fraction=0.2, min_per_class=1):\n    \"\"\"Sous-échantillonnage stratifié, conserve toutes les classes.\"\"\"\n    if hasattr(dataset, \"targets\"):\n        targets = np.array(dataset.targets)\n    elif hasattr(dataset, \"dataset\") and hasattr(dataset.dataset, \"targets\"):\n        targets = np.array([dataset.dataset.targets[i] for i in dataset.indices])\n    else:\n        targets = np.array([l for _, l in tqdm(dataset, desc=\"Lecture labels\")])\n\n    indices = np.arange(len(dataset))\n    subset_idx = []\n    for cls in np.unique(targets):\n        cls_idx = indices[targets == cls]\n        n = max(min_per_class, int(math.ceil(len(cls_idx) * fraction)))\n        n = min(n, len(cls_idx))\n        subset_idx.extend(np.random.choice(cls_idx, n, replace=False).tolist())\n\n    print(f\"Subset stratifié: {len(subset_idx)}/{len(dataset)} images, \"\n          f\"{len(np.unique(targets[subset_idx]))} classes\")\n    return Subset(dataset, subset_idx)\n\n\nprint(\"Fonctions utilitaires prêtes.\")","metadata":{"execution":{"iopub.status.busy":"2026-02-10T19:34:48.124698Z","iopub.execute_input":"2026-02-10T19:34:48.12496Z","iopub.status.idle":"2026-02-10T19:34:48.14216Z","shell.execute_reply.started":"2026-02-10T19:34:48.124939Z","shell.execute_reply":"2026-02-10T19:34:48.141427Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3. Chargement des données (lecture directe depuis le disque)\n\nLes images sont lues à la volée via `ImageFolder` + `DataLoader` — **aucun cache RAM ni disque**.  \nL'augmentation de données est configurable via `USE_AUGMENTATION`.","metadata":{"trusted":true}},{"cell_type":"code","source":"# ── 3a. Charger les métadonnées ──────────────────────────────\nwith open(JSON_NAMES) as f:\n    species_map = json.load(f)          # str_id → nom commun\nwith open(JSON_META) as f:\n    meta_dict = json.load(f)\n\ndf_meta = pd.DataFrame.from_dict(meta_dict, orient=\"index\")\ndf_train = df_meta[df_meta[\"split\"] == \"train\"]\ndf_val   = df_meta[df_meta[\"split\"] == \"val\"]\n\nprint(f\"Metadata — Train: {len(df_train):,}  Val: {len(df_val):,}  \"\n      f\"Classes: {df_train['species_id'].nunique()}\")\n\n# ── 3b. Transforms (augmentation optionnelle) ───────────────\nUSE_AUGMENTATION = True   # ← Passer à False pour désactiver l'augmentation\n\nif USE_AUGMENTATION:\n    transform_train = transforms.Compose([\n        transforms.Resize((256, 256)),\n        transforms.RandomCrop(224),\n        transforms.RandomHorizontalFlip(),\n        transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.1),\n        transforms.ToTensor(),\n        transforms.Normalize(IMAGENET_MEAN, IMAGENET_STD),\n    ])\n    print(\"Augmentation ACTIVÉE (RandomCrop, Flip, ColorJitter)\")\nelse:\n    transform_train = transforms.Compose([\n        transforms.Resize((224, 224)),\n        transforms.ToTensor(),\n        transforms.Normalize(IMAGENET_MEAN, IMAGENET_STD),\n    ])\n    print(\"Augmentation DÉSACTIVÉE (Resize simple)\")\n\ntransform_val = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor(),\n    transforms.Normalize(IMAGENET_MEAN, IMAGENET_STD),\n])\n\n# ── 3c. ImageFolder (lecture directe depuis disque) ──────────\nprint(\"Indexation en cours (prend quelques minutes)...\")\nfull_train_ds = datasets.ImageFolder(root=TRAIN_DIR, transform=transform_train)\nfull_val_ds   = datasets.ImageFolder(root=VAL_DIR,   transform=transform_val)\n\n# Mapping string → entier & noms lisibles\nclass_to_idx = full_train_ds.class_to_idx          # str_id → int idx\nidx_to_class = {v: k for k, v in class_to_idx.items()}  # int idx → str_id\nclass_names_readable = {\n    str(idx): species_map.get(sid, sid)\n    for sid, idx in class_to_idx.items()\n}\nnum_classes = len(full_train_ds.classes)\n\nprint(f\"Train: {len(full_train_ds):,}  Val: {len(full_val_ds):,}  Classes: {num_classes}\")\n\n# ── 3d. Analyse de distribution & sélection ────────────────\nclass_counts_train = df_train[\"species_id\"].value_counts()\n\n# Classes dominantes (top N)\ntop_class_ids = class_counts_train.head(TOP_N_DOMINANT).index.tolist()\ndominant_targets = [class_to_idx[sid] for sid in top_class_ids if sid in class_to_idx]\n\n# Classes valides pour l'évaluation (≥ SEUIL_RARE images)\nvalid_class_ids   = class_counts_train[class_counts_train >= SEUIL_RARE].index.tolist()\nvalid_class_idxs  = [class_to_idx[s] for s in valid_class_ids if s in class_to_idx]\nexclude_class_ids = class_counts_train[class_counts_train < SEUIL_RARE].index.tolist()\n\nprint(f\"\\nClasses dominantes ({TOP_N_DOMINANT}):\")\nfor sid in top_class_ids:\n    print(f\"  {species_map.get(sid, sid)} — {class_counts_train[sid]} images\")\nprint(f\"Classes valides (≥{SEUIL_RARE} img): {len(valid_class_idxs)}/{num_classes}\")\n\n# ── 3e. Création subset d'entraînement équilibré ────────────\ntargets_arr = np.array(full_train_ds.targets)\ndist = Counter(targets_arr)\n\nidx_dominant = np.concatenate([np.where(targets_arr == c)[0] for c in dominant_targets]).tolist()\n\n# Filtrer les classes très rares de l'entraînement\nvalid_for_train = {c for c, n in dist.items() if n >= SEUIL_TRES_RARE}\nidx_other = [i for i in range(len(full_train_ds))\n             if i not in set(idx_dominant) and targets_arr[i] in valid_for_train]\n\nn_others = min(len(idx_other), len(idx_dominant) * 2)\nsampled_other = random.sample(idx_other, n_others)\ntrain_indices = idx_dominant + sampled_other\n\n# ── 3f. Datasets & DataLoaders finaux (pas de cache) ────────\ntrain_ds = Subset(full_train_ds, train_indices)\nval_ds   = full_val_ds\n\nprint(f\"\\nSubset train: {len(train_ds):,} images \"\n      f\"({len(idx_dominant):,} dominant + {n_others:,} autres)\")\nprint(f\"Val: {len(val_ds):,} images\")\nprint(\"Données prêtes (lecture directe depuis disque, aucun cache).\")","metadata":{"execution":{"iopub.status.busy":"2026-02-10T19:34:48.216181Z","iopub.execute_input":"2026-02-10T19:34:48.216396Z","iopub.status.idle":"2026-02-10T19:40:39.197285Z","shell.execute_reply.started":"2026-02-10T19:34:48.216375Z","shell.execute_reply":"2026-02-10T19:40:39.196569Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 4. Analyse exploratoire de la distribution","metadata":{"execution":{"iopub.execute_input":"2026-02-03T11:08:46.572619Z","iopub.status.busy":"2026-02-03T11:08:46.57232Z","iopub.status.idle":"2026-02-03T11:08:46.600007Z","shell.execute_reply":"2026-02-03T11:08:46.599313Z","shell.execute_reply.started":"2026-02-03T11:08:46.572593Z"},"trusted":true}},{"cell_type":"code","source":"# ── Distribution des classes (train) ─────────────────────────\ncc = class_counts_train.sort_values(ascending=False)\ngini = gini_coefficient(cc.values)\n\nprint(f\"Min: {cc.min()}  Max: {cc.max()}  Moy: {cc.mean():.1f}  \"\n      f\"Gini: {gini:.4f}  Ratio max/min: {cc.max()/cc.min():.0f}\")\n\ntres_rares = (cc < SEUIL_TRES_RARE).sum()\nrares      = ((cc >= SEUIL_TRES_RARE) & (cc < SEUIL_RARE)).sum()\nfaibles    = ((cc >= SEUIL_RARE) & (cc < 50)).sum()\nnormales   = (cc >= 50).sum()\n\nprint(f\"Très rares (<{SEUIL_TRES_RARE}): {tres_rares}  \"\n      f\"Rares ({SEUIL_TRES_RARE}-{SEUIL_RARE-1}): {rares}  \"\n      f\"Faibles (20-49): {faibles}  Normales (≥50): {normales}\")\n\n# ── Visualisation ────────────────────────────────────────────\nfig, axes = plt.subplots(1, 3, figsize=(18, 5))\n\n# 1) Long tail\nax = axes[0]\nx = range(len(cc))\nax.bar(x, cc.values, width=1, color=\"skyblue\", edgecolor=\"none\")\nax.bar(x[:10], cc.values[:10], width=1, color=\"salmon\", label=\"Top 10\")\nax.axhline(cc.mean(), color=\"red\", ls=\"--\", alpha=.6, label=f\"Moy {cc.mean():.0f}\")\nax.set(title=\"Long Tail (Train)\", xlabel=\"Classe (triée)\", ylabel=\"Images\")\nax.legend()\n\n# 2) Cumulative Pareto\nax = axes[1]\ncum = 100 * np.cumsum(cc.values) / cc.sum()\nax.plot(x, cum, \"g-\", lw=2)\nfor pct, col in [(50, \"red\"), (80, \"orange\")]:\n    ax.axhline(pct, color=col, ls=\"--\", alpha=.5, label=f\"{pct}%\")\n    idx_pct = np.argmax(cum >= pct)\n    ax.axvline(idx_pct, color=col, ls=\":\", alpha=.4)\nax.set(title=\"Pareto cumulatif\", xlabel=\"Nb classes incluses\", ylabel=\"% dataset\"); ax.legend()\n\n# 3) Répartition par catégorie\nax = axes[2]\nsizes = [tres_rares, rares, faibles, normales]\nlabels = [f\"<{SEUIL_TRES_RARE}\", f\"{SEUIL_TRES_RARE}-{SEUIL_RARE-1}\", \"20-49\", \"≥50\"]\ncolors = [\"#ff6b6b\", \"#ffa502\", \"#ffd93d\", \"#6bcb77\"]\nax.pie(sizes, labels=labels, colors=colors, autopct=\"%1.1f%%\", startangle=90)\nax.set_title(\"Répartition par catégorie\")\n\nplt.tight_layout(); plt.show()","metadata":{"execution":{"iopub.status.busy":"2026-02-10T19:40:39.198881Z","iopub.execute_input":"2026-02-10T19:40:39.199206Z","iopub.status.idle":"2026-02-10T19:40:41.051993Z","shell.execute_reply.started":"2026-02-10T19:40:39.199182Z","shell.execute_reply":"2026-02-10T19:40:41.05141Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 5. Définition des modèles & du classificateur en cascade","metadata":{}},{"cell_type":"code","source":"\n# ── Loss functions ───────────────────────────────────────────\nclass FocalLoss(nn.Module):\n    \"\"\"Focal Loss (Lin et al. 2017) — focus sur les exemples difficiles.\"\"\"\n    def __init__(self, alpha=1, gamma=2, reduction=\"mean\"):\n        super().__init__()\n        self.alpha, self.gamma, self.reduction = alpha, gamma, reduction\n\n    def forward(self, inputs, targets):\n        ce = F.cross_entropy(inputs, targets, reduction=\"none\")\n        pt = torch.exp(-ce)\n        loss = self.alpha * (1 - pt) ** self.gamma * ce\n        return loss.mean() if self.reduction == \"mean\" else loss.sum()\n\n\n# ── Mini-modèle binaire (RepViT) ────────────────────────────\nclass MiniCNN_RePViT(nn.Module):\n    def __init__(self, num_classes=2, pretrained=True):\n        super().__init__()\n        self.backbone = timm.create_model(\n            \"repvit_m1.dist_in1k\", pretrained=pretrained,\n            num_classes=0, global_pool=\"avg\")\n        feat = self.backbone.num_features\n        self.classifier = nn.Sequential(nn.Dropout(0.2), nn.Linear(feat, num_classes))\n        for p in self.backbone.parameters():\n            p.requires_grad = False\n\n    def forward(self, x):\n        return self.classifier(self.backbone(x))\n\n\n# ── Wrapper ViT HuggingFace ─────────────────────────────────\nclass ViTWrapper(nn.Module):\n    def __init__(self, vit_model):\n        super().__init__()\n        self.vit = vit_model\n    def forward(self, x):\n        return self.vit(x).logits\n\n\n# ── Classificateur en cascade ────────────────────────────────\nclass CascadeClassifier:\n    \"\"\"\n    Cascade optimisée : mini-modèles binaires rapides + Big Model fallback.\n    Affiche precision, recall, F1, weighted-F1, AUROC à chaque étape.\n    \"\"\"\n    def __init__(self, dominant_classes, big_model, device=DEVICE,\n                 class_names=None, valid_classes=None):\n        self.dominant_classes = dominant_classes\n        self.device = device\n        self.class_names = class_names or {}\n        self.valid_classes = valid_classes\n\n        self.big_model = big_model.to(device)\n        self.mini_models = nn.ModuleDict()\n        for c in dominant_classes:\n            self.mini_models[str(c)] = MiniCNN_RePViT(num_classes=2).to(device)\n\n        self.is_multigpu = torch.cuda.device_count() > 1\n        if self.is_multigpu:\n            self.big_model = nn.DataParallel(self.big_model)\n            for k in self.mini_models:\n                self.mini_models[k] = nn.DataParallel(self.mini_models[k])\n        self.scaler = GradScaler()\n        print(f\"Cascade initialisée — {len(dominant_classes)} mini-modèles, \"\n              f\"device={device}, multi-GPU={self.is_multigpu}\")\n\n    # ─────────── TRAINING ───────────\n    def train_mini_models(self, dataset, epochs=5, lr=1e-3, batch_size=256,\n                      accumulation_steps=1):\n        \"\"\"Entraîne les mini-modèles avec pos_weight adaptatif par classe.\"\"\"\n        loader = DataLoader(dataset, batch_size=batch_size, shuffle=True,\n                            num_workers=NUM_WORKERS, pin_memory=True,\n                            persistent_workers=NUM_WORKERS > 0,\n                            prefetch_factor=2, drop_last=True)\n\n        optimizers = {str(c): optim.AdamW(self.mini_models[str(c)].parameters(),\n                      lr=lr, weight_decay=0.01) for c in self.dominant_classes}\n        schedulers = {k: optim.lr_scheduler.CosineAnnealingLR(o, T_max=epochs)\n                      for k, o in optimizers.items()}\n\n        # ── pos_weight adaptatif par classe dominante ──\n        if hasattr(dataset, \"dataset\") and hasattr(dataset, \"indices\"):\n            _targets = np.array([dataset.dataset.targets[i] for i in dataset.indices])\n        elif hasattr(dataset, \"targets\"):\n            _targets = np.array(dataset.targets)\n        else:\n            _targets = np.array([l for _, l in dataset])\n        _total = len(_targets)\n        _counts = Counter(_targets.tolist())\n\n        criteria = {}\n        print(\"  pos_weight adaptatif par classe :\")\n        for c in self.dominant_classes:\n            pos = _counts.get(c, 1)\n            neg = _total - pos\n            pw = float(np.clip(neg / max(1, pos), 1.0, 20.0))\n            criteria[str(c)] = nn.BCEWithLogitsLoss(\n                pos_weight=torch.tensor([pw]).to(self.device))\n            name = self.class_names.get(str(c), str(c))[:15]\n            print(f\"    {name}: pos={pos:,}, neg={neg:,}, pos_weight={pw:.1f}\")\n\n        history = {str(c): {\"loss\": [], \"acc\": [], \"f1\": []}\n                   for c in self.dominant_classes}\n\n        for epoch in range(epochs):\n            self.mini_models.train()\n            ep_loss  = {str(c): 0.0 for c in self.dominant_classes}\n            ep_true  = {str(c): [] for c in self.dominant_classes}\n            ep_pred  = {str(c): [] for c in self.dominant_classes}\n\n            pbar = tqdm(loader, desc=f\"Mini Epoch {epoch+1}/{epochs}\")\n            for batch_idx, (imgs, lbls) in enumerate(pbar):\n                imgs = imgs.to(self.device, non_blocking=True)\n                lbls = lbls.to(self.device, non_blocking=True)\n\n                # ═══ OPTIMISATION : extraire les features UNE SEULE FOIS ═══\n                with torch.no_grad(), autocast():\n                    first_key = str(self.dominant_classes[0])\n                    first_model = self.mini_models[first_key]\n                    base = first_model.module if hasattr(first_model, 'module') else first_model\n                    shared_features = base.backbone(imgs)\n\n                for cs, model in self.mini_models.items():\n                    tc = int(cs)\n                    binary = (lbls == tc).float().unsqueeze(1)\n\n                    m = model.module if hasattr(model, 'module') else model\n                    with autocast():\n                        out = m.classifier(shared_features)\n                        loss = criteria[cs](out[:, 1:2], binary) / accumulation_steps\n\n                    self.scaler.scale(loss).backward()\n\n                    if (batch_idx + 1) % accumulation_steps == 0:\n                        self.scaler.step(optimizers[cs])\n                        self.scaler.update()\n                        optimizers[cs].zero_grad(set_to_none=True)\n\n                    ep_loss[cs] += loss.item() * accumulation_steps\n                    with torch.no_grad():\n                        p = (torch.sigmoid(out[:, 1:2]) > 0.5).float()\n                        ep_true[cs].extend(binary.cpu().numpy().ravel().tolist())\n                        ep_pred[cs].extend(p.cpu().numpy().ravel().tolist())\n\n                pbar.set_postfix(loss=f\"{np.mean([ep_loss[str(c)]/max(1,batch_idx+1) for c in self.dominant_classes]):.4f}\")\n\n            for s in schedulers.values():\n                s.step()\n\n            # Métriques par mini-modèle\n            print(f\"\\n  Epoch {epoch+1}/{epochs}:\")\n            print(f\"  {'Classe':<14} {'Loss':>8} {'Acc':>8} {'Prec':>8} {'Rec':>8} {'F1':>8}\")\n            for cs in [str(c) for c in self.dominant_classes]:\n                avg_l = ep_loss[cs] / len(loader)\n                yt, yp = np.array(ep_true[cs]), np.array(ep_pred[cs])\n                acc = 100 * (yt == yp).mean()\n                pr  = precision_score(yt, yp, zero_division=0)\n                rc  = recall_score(yt, yp, zero_division=0)\n                f1  = f1_score(yt, yp, zero_division=0)\n                history[cs][\"loss\"].append(avg_l)\n                history[cs][\"acc\"].append(acc)\n                history[cs][\"f1\"].append(f1)\n                name = self.class_names.get(cs, cs)[:12]\n                print(f\"  {name:<14} {avg_l:>8.4f} {acc:>7.2f}% {pr:>8.4f} {rc:>8.4f} {f1:>8.4f}\")\n        print(\"Mini-modèles entraînés.\")\n        return history\n\n    # ─────────── EVALUATION ───────────\n    def evaluate(self, dataset, batch_size=128, threshold=0.5, filter_rare=True):\n        \"\"\"Évaluation complète : accuracy, precision, recall, F1, weighted-F1, AUROC.\"\"\"\n        loader = DataLoader(dataset, batch_size=batch_size, shuffle=False,\n                            num_workers=NUM_WORKERS, pin_memory=True)\n        y_true, y_pred, responsible = [], [], []\n        all_big_probs = []\n\n        self.big_model.eval(); self.mini_models.eval()\n        t0 = time.time()\n\n        with torch.no_grad():\n            for imgs, lbls in tqdm(loader, desc=\"Evaluation\"):\n                imgs = imgs.to(self.device, non_blocking=True)\n                bs = imgs.size(0)\n                final = torch.full((bs,), -1, dtype=torch.long, device=self.device)\n                handled = torch.zeros(bs, dtype=torch.bool, device=self.device)\n                resp = [\"\"] * bs\n\n                with autocast():\n                    for ci in self.dominant_classes:\n                        logits = self.mini_models[str(ci)](imgs)\n                        probs = torch.sigmoid(logits[:, 1])\n                        triggered = (probs > threshold) & ~handled\n                        if triggered.any():\n                            final[triggered] = ci\n                            handled[triggered] = True\n                            for idx in triggered.nonzero().squeeze(-1).tolist():\n                                if isinstance(idx, int):\n                                    resp[idx] = f\"Mini-{ci}\"\n\n                    fb = ~handled\n                    if fb.any():\n                        out = self.big_model(imgs[fb])\n                        probs_big = F.softmax(out, dim=1)\n                        _, preds = torch.max(out, 1)\n                        final[fb] = preds\n                        all_big_probs.append(probs_big.cpu())\n                        for idx in fb.nonzero().squeeze(-1).tolist():\n                            if isinstance(idx, int):\n                                resp[idx] = \"Big\"\n\n                y_true.extend(lbls.cpu().numpy())\n                y_pred.extend(final.cpu().numpy())\n                responsible.extend(resp)\n\n        elapsed = time.time() - t0\n        y_true, y_pred = np.array(y_true), np.array(y_pred)\n\n        big_scores = torch.cat(all_big_probs).numpy() if all_big_probs else None\n\n        # ── Métriques globales ──\n        m_all = compute_all_metrics(y_true, y_pred, prefix=\"all_\")\n        print_metrics(m_all, f\"Toutes classes ({len(y_true):,} images, {elapsed:.1f}s)\")\n\n        # ── Métriques filtrées (classes valides) ──\n        if filter_rare and self.valid_classes is not None:\n            mask = np.isin(y_true, self.valid_classes)\n            m_filt = compute_all_metrics(y_true[mask], y_pred[mask], prefix=\"valid_\")\n            print_metrics(m_filt, f\"Classes valides (≥{SEUIL_RARE} img, {mask.sum():,} samples)\")\n        else:\n            m_filt = m_all\n\n        # ── Stats par étage ──\n        stats = defaultdict(lambda: {\"n\": 0, \"ok\": 0})\n        for t, p, r in zip(y_true, y_pred, responsible):\n            stats[r][\"n\"] += 1\n            if t == p:\n                stats[r][\"ok\"] += 1\n        print(f\"\\n  {'Modèle':<14} {'Appels':>8} {'Charge':>8} {'Acc':>8}\")\n        for name, d in sorted(stats.items()):\n            acc = 100 * d[\"ok\"] / max(1, d[\"n\"])\n            print(f\"  {name:<14} {d['n']:>8,} {100*d['n']/len(y_true):>7.1f}% {acc:>7.2f}%\")\n\n        # ── Matrice de confusion (si ≤ 20 classes) ──\n        uniq = sorted(set(y_true) | set(y_pred))\n        if len(uniq) <= 20:\n            names = [self.class_names.get(str(l), str(l))[:15] for l in uniq]\n            cm = confusion_matrix(y_true, y_pred, labels=uniq)\n            fig, ax = plt.subplots(figsize=(10, 8))\n            ConfusionMatrixDisplay(cm, display_labels=names).plot(\n                ax=ax, cmap=\"Blues\", values_format=\"d\", xticks_rotation=\"vertical\")\n            ax.set_title(f\"Confusion (Acc {m_all['all_accuracy']:.1f}%)\")\n            plt.tight_layout(); plt.show()\n\n        return {**m_all, **m_filt,\n                \"y_true\": y_true, \"y_pred\": y_pred,\n                \"stats\": dict(stats)}\n\n\nprint(\"Modèles et cascade définis.\")\n","metadata":{"execution":{"iopub.status.busy":"2026-02-10T19:40:41.05287Z","iopub.execute_input":"2026-02-10T19:40:41.053068Z","iopub.status.idle":"2026-02-10T19:40:41.086507Z","shell.execute_reply.started":"2026-02-10T19:40:41.053045Z","shell.execute_reply":"2026-02-10T19:40:41.085869Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 6. Sauvegarde & chargement de la cascade","metadata":{}},{"cell_type":"code","source":"def save_cascade(cascade, save_dir=SAVE_DIR, suffix=\"\"):\n    ts = datetime.now().strftime(\"%Y%m%d_%H%M%S\")\n    rd = os.path.join(save_dir, f\"run_{ts}{suffix}\"); os.makedirs(rd, exist_ok=True)\n    # Big Model\n    bm = cascade.big_model.module if hasattr(cascade.big_model, \"module\") else cascade.big_model\n    torch.save({\"state_dict\": bm.state_dict()}, os.path.join(rd, \"big_model.pth\"))\n    # Mini Models\n    mm = {str(c): (m.module if hasattr(m, \"module\") else m).state_dict()\n          for c, m in zip(cascade.dominant_classes,\n                          [cascade.mini_models[str(c)] for c in cascade.dominant_classes])}\n    torch.save(mm, os.path.join(rd, \"mini_models.pth\"))\n    # Config\n    cfg = {\"dominant_classes\": cascade.dominant_classes,\n           \"class_names\": cascade.class_names,\n           \"valid_classes\": cascade.valid_classes,\n           \"timestamp\": ts}\n    json.dump(cfg, open(os.path.join(rd, \"config.json\"), \"w\"), indent=2, default=str)\n    print(f\"Cascade sauvegardée : {rd}\")\n    return rd\n\n\ndef load_cascade(load_dir, device=DEVICE, num_cls=None):\n    cfg = json.load(open(os.path.join(load_dir, \"config.json\")))\n    nc = num_cls or num_classes\n    from transformers import ViTForImageClassification\n    vit = ViTForImageClassification.from_pretrained(\n        \"google/vit-base-patch16-224\", num_labels=nc, ignore_mismatched_sizes=True)\n    bm = ViTWrapper(vit).to(device)\n    bm_data = torch.load(os.path.join(load_dir, \"big_model.pth\"), map_location=device)\n    bm.load_state_dict(bm_data[\"state_dict\"])\n\n    cascade = CascadeClassifier(\n        cfg[\"dominant_classes\"], bm, device,\n        cfg.get(\"class_names\", {}), cfg.get(\"valid_classes\"))\n\n    mm_data = torch.load(os.path.join(load_dir, \"mini_models.pth\"), map_location=device)\n    for k, sd in mm_data.items():\n        m = cascade.mini_models[k]\n        (m.module if hasattr(m, \"module\") else m).load_state_dict(sd)\n\n    cascade.big_model.eval(); cascade.mini_models.eval()\n    print(f\"Cascade chargée depuis {load_dir}\")\n    return cascade\n\n\nprint(\"Fonctions save/load prêtes.\")","metadata":{"execution":{"iopub.status.busy":"2026-02-10T19:40:41.087506Z","iopub.execute_input":"2026-02-10T19:40:41.087724Z","iopub.status.idle":"2026-02-10T19:40:41.103402Z","shell.execute_reply.started":"2026-02-10T19:40:41.087696Z","shell.execute_reply":"2026-02-10T19:40:41.102864Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 7. Entraînement du Big Model (ViT) avec métriques complètes","metadata":{"trusted":true}},{"cell_type":"code","source":"from transformers import ViTForImageClassification\n\n# ── Préparer le Big Model ViT ────────────────────────────────\nvit_model = ViTForImageClassification.from_pretrained(\n    \"google/vit-base-patch16-224\",\n    num_labels=num_classes,\n    ignore_mismatched_sizes=True,\n)\n# Geler toutes les couches sauf les 20 derniers paramètres\nparams = list(vit_model.parameters())\nfor p in params[:-20]:\n    p.requires_grad = False\n\nbig_model = ViTWrapper(vit_model).to(DEVICE)\nbm_data = torch.load(os.path.join(\"/kaggle/input/plantnet-vit/pytorch/default/1/\", \"big_model.pth\"), map_location=DEVICE)\nbig_model.load_state_dict(bm_data[\"state_dict\"])\nprint(f\"ViT — {num_classes} classes, {sum(p.requires_grad for p in params)} params entraînables\")\n\n# ── Loss & Optimizer ─────────────────────────────────────────\nclass_weights = compute_class_weights(train_ds, num_classes, DEVICE, smoothing=0.3)\ncriterion_big = FocalLoss(alpha=1, gamma=2)\noptimizer_big = optim.AdamW(\n    filter(lambda p: p.requires_grad, big_model.parameters()),\n    lr=5e-4, weight_decay=0.02)\n\nN_EPOCHS_BIG = 8\nBIG_BATCH    = 64 if DEVICE == \"cuda\" else 16\nscheduler_big = optim.lr_scheduler.CosineAnnealingLR(optimizer_big, T_max=N_EPOCHS_BIG)\nscaler_big = GradScaler()\n\n# DataLoader classique — lecture directe depuis disque\nbig_loader = DataLoader(train_ds, batch_size=BIG_BATCH, shuffle=True,\n                        num_workers=NUM_WORKERS, pin_memory=True,\n                        persistent_workers=NUM_WORKERS > 0,\n                        prefetch_factor=2)\n\n# ── Boucle d'entraînement avec métriques ─────────────────────\nbig_history = {\"loss\": [], \"acc\": [], \"prec\": [], \"rec\": [], \"f1\": [], \"f1w\": []}\n\nfor epoch in range(N_EPOCHS_BIG):\n    big_model.train()\n    ep_loss, all_t, all_p = 0.0, [], []\n\n    pbar = tqdm(big_loader, desc=f\"ViT Epoch {epoch+1}/{N_EPOCHS_BIG}\")\n    for imgs, lbls in pbar:\n        imgs = imgs.to(DEVICE, non_blocking=True)\n        lbls = lbls.to(DEVICE, non_blocking=True)\n        optimizer_big.zero_grad(set_to_none=True)\n        with autocast():\n            out = big_model(imgs)\n            loss = criterion_big(out, lbls)\n        scaler_big.scale(loss).backward()\n        scaler_big.step(optimizer_big); scaler_big.update()\n\n        ep_loss += loss.item()\n        _, pred = torch.max(out, 1)\n        all_t.extend(lbls.cpu().numpy()); all_p.extend(pred.cpu().numpy())\n        pbar.set_postfix(loss=f\"{loss.item():.4f}\",\n                         acc=f\"{100*(np.array(all_t)==np.array(all_p)).mean():.1f}%\")\n    scheduler_big.step()\n\n    # Métriques epoch\n    m = compute_all_metrics(all_t, all_p, prefix=\"\")\n    big_history[\"loss\"].append(ep_loss / len(big_loader))\n    for k in [\"acc\", \"prec\", \"rec\", \"f1\", \"f1w\"]:\n        key_map = {\"acc\": \"accuracy\", \"prec\": \"precision_macro\",\n                   \"rec\": \"recall_macro\", \"f1\": \"f1_macro\", \"f1w\": \"f1_weighted\"}\n        big_history[k].append(m[key_map[k]])\n\n    print(f\"  Epoch {epoch+1}: loss={big_history['loss'][-1]:.4f}  \"\n          f\"acc={m['accuracy']:.2f}%  prec={m['precision_macro']:.4f}  \"\n          f\"rec={m['recall_macro']:.4f}  F1={m['f1_macro']:.4f}  \"\n          f\"wF1={m['f1_weighted']:.4f}\")\n\n# ── Courbes d'entraînement ──────────────────────────────────\nfig, axes = plt.subplots(1, 3, figsize=(16, 4))\naxes[0].plot(big_history[\"loss\"]); axes[0].set(title=\"Loss\", xlabel=\"Epoch\")\naxes[1].plot(big_history[\"acc\"]);  axes[1].set(title=\"Accuracy (%)\", xlabel=\"Epoch\")\naxes[2].plot(big_history[\"f1\"], label=\"macro-F1\")\naxes[2].plot(big_history[\"f1w\"], label=\"weighted-F1\")\naxes[2].set(title=\"F1 Score\", xlabel=\"Epoch\"); axes[2].legend()\nplt.tight_layout(); plt.show()\n\nprint(\"Big Model entraîné.\")","metadata":{"execution":{"iopub.status.busy":"2026-02-10T19:50:04.808639Z","iopub.execute_input":"2026-02-10T19:50:04.80893Z","iopub.status.idle":"2026-02-10T19:54:32.491555Z","shell.execute_reply.started":"2026-02-10T19:50:04.808904Z","shell.execute_reply":"2026-02-10T19:54:32.490048Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"big_model = ViTWrapper(vit_model).to(DEVICE)\nbm_data = torch.load(os.path.join(\"/kaggle/input/plantnet-vit/pytorch/default/1/\", \"big_model.pth\"), map_location=DEVICE)\nbig_model.load_state_dict(bm_data[\"state_dict\"])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-10T19:55:05.201149Z","iopub.execute_input":"2026-02-10T19:55:05.201981Z","iopub.status.idle":"2026-02-10T19:55:05.514908Z","shell.execute_reply.started":"2026-02-10T19:55:05.20192Z","shell.execute_reply":"2026-02-10T19:55:05.514191Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 8. Entraînement de la cascade (mini-modèles)","metadata":{}},{"cell_type":"code","source":"\n# ── Transform aggressif pour mini-modèles ────────────────────\ntransform_mini = transforms.Compose([\n    transforms.Resize((256, 256)),\n    transforms.RandomResizedCrop(224, scale=(0.6, 1.0)),\n    transforms.RandomHorizontalFlip(),\n    transforms.RandomVerticalFlip(p=0.1),\n    transforms.RandomRotation(20),\n    transforms.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.2, hue=0.05),\n    transforms.RandomGrayscale(p=0.1),\n    transforms.ToTensor(),\n    transforms.Normalize(IMAGENET_MEAN, IMAGENET_STD),\n])\n\nfull_train_ds_mini = datasets.ImageFolder(root=TRAIN_DIR, transform=transform_mini)\ntrain_ds_mini = Subset(full_train_ds_mini, train_indices)\nprint(f\"Dataset mini-modèles: {len(train_ds_mini):,} images (augmentation agressive)\")\n\n# ── Cascade ──────────────────────────────────────────────────\ncascade = CascadeClassifier(\n    dominant_classes=dominant_targets,\n    big_model=big_model,\n    device=DEVICE,\n    class_names=class_names_readable,\n    valid_classes=valid_class_idxs,\n)\n\n# cascade = load_cascade(\"/kaggle/working/cascade_models/run_20260209_175558_v1\")\n\nmini_history = cascade.train_mini_models(\n    train_ds_mini, epochs=5, batch_size=128, lr=1e-3, accumulation_steps=2\n)\n\n# Sauvegarde\nsave_path = save_cascade(cascade, suffix=\"_v2\")\nprint(f\"Sauvegardé : {save_path}\")\n","metadata":{"execution":{"iopub.status.busy":"2026-02-10T19:55:07.934255Z","iopub.execute_input":"2026-02-10T19:55:07.935026Z","iopub.status.idle":"2026-02-10T20:39:05.809366Z","shell.execute_reply.started":"2026-02-10T19:55:07.934995Z","shell.execute_reply":"2026-02-10T20:39:05.808484Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 9. Évaluation complète (Precision, Recall, F1, wF1, AUROC, Top-K)","metadata":{"trusted":true}},{"cell_type":"code","source":"\n# ── Subset de validation stratifié ────────────────────────────\nval_subset = stratified_subset(val_ds, fraction=0.6)\n\n# ══════════════════════════════════════════════════════════════\n# PRÉ-CALCUL UNIQUE — 1 seul DataLoader pour toute l'analyse\n# ══════════════════════════════════════════════════════════════\ndef precompute_all(cascade, dataset, batch_size=128):\n    \"\"\"Pré-calcule TOUTES les sorties (mini + big) en UN seul passage.\"\"\"\n    loader = DataLoader(dataset, batch_size=batch_size, shuffle=False,\n                        num_workers=NUM_WORKERS, pin_memory=True)\n    targets_list = []\n    all_big_logits = []\n    mini_probs = {c: [] for c in cascade.dominant_classes}\n\n    cascade.big_model.eval()\n    cascade.mini_models.eval()\n    t0 = time.time()\n\n    with torch.no_grad():\n        for imgs, lbls in tqdm(loader, desc=\"Pré-calcul (passage unique)\"):\n            imgs = imgs.to(DEVICE, non_blocking=True)\n            targets_list.extend(lbls.numpy())\n            with autocast():\n                out_big = cascade.big_model(imgs)\n                all_big_logits.append(out_big.cpu())\n                for c in cascade.dominant_classes:\n                    logits = cascade.mini_models[str(c)](imgs)\n                    mini_probs[c].extend(\n                        torch.sigmoid(logits[:, 1]).cpu().numpy())\n\n    big_logits = torch.cat(all_big_logits)\n    big_preds = big_logits.argmax(dim=1).numpy()\n\n    elapsed = time.time() - t0\n    print(f\"Pré-calcul terminé en {elapsed:.1f}s — {len(targets_list):,} images\")\n\n    return {\n        \"targets\": np.array(targets_list),\n        \"big_preds\": big_preds,\n        \"big_logits\": big_logits,\n        \"mini_probs\": {k: np.array(v) for k, v in mini_probs.items()},\n    }\n\ncache = precompute_all(cascade, val_subset)\n\n# ══════════════════════════════════════════════════════════════\n# ÉVALUATION MULTI-SEUIL (aucun DataLoader supplémentaire)\n# ══════════════════════════════════════════════════════════════\nthresholds = [0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9]\nall_results = []\n\nfor th in thresholds:\n    handled = np.zeros(len(cache[\"targets\"]), dtype=bool)\n    final = np.full(len(cache[\"targets\"]), -1, dtype=int)\n\n    for c in cascade.dominant_classes:\n        triggered = (cache[\"mini_probs\"][c] > th) & ~handled\n        final[triggered] = c\n        handled[triggered] = True\n\n    fb = ~handled\n    final[fb] = cache[\"big_preds\"][fb]\n\n    y_true, y_pred = cache[\"targets\"], final\n    m_all = compute_all_metrics(y_true, y_pred, prefix=\"all_\")\n\n    mask_valid = np.isin(y_true, valid_class_idxs)\n    m_valid = compute_all_metrics(\n        y_true[mask_valid], y_pred[mask_valid], prefix=\"valid_\")\n\n    all_results.append({\n        \"threshold\": th,\n        **m_all, **m_valid,\n        \"cascade_pct\": 100 * handled.mean(),\n        \"big_pct\": 100 * fb.mean(),\n    })\n\n# ── Tableau récapitulatif ────────────────────────────────────\nprint(f\"\\n{'═'*105}\")\nprint(f\"  RÉSULTATS MULTI-SEUIL ({len(cache['targets']):,} images)\")\nprint(f\"{'═'*105}\")\nprint(f\"  {'Seuil':>6} │ {'Acc(all)':>9} │ {'F1m(all)':>9} │ \"\n      f\"{'wF1(all)':>9} │ {'Acc(val)':>9} │ {'F1m(val)':>9} │ \"\n      f\"{'wF1(val)':>9} │ {'%Mini':>6} │ {'%Big':>6}\")\nprint(f\"  {'─'*6}─┼─{'─'*9}─┼─{'─'*9}─┼─{'─'*9}─┼─\"\n      f\"{'─'*9}─┼─{'─'*9}─┼─{'─'*9}─┼─{'─'*6}─┼─{'─'*6}\")\nfor r in all_results:\n    print(f\"  {r['threshold']:>6.2f} │ {r['all_accuracy']:>8.2f}% │ \"\n          f\"{r['all_f1_macro']:>9.4f} │ {r['all_f1_weighted']:>9.4f} │ \"\n          f\"{r['valid_accuracy']:>8.2f}% │ {r['valid_f1_macro']:>9.4f} │ \"\n          f\"{r['valid_f1_weighted']:>9.4f} │ {r['cascade_pct']:>5.1f}% │ \"\n          f\"{r['big_pct']:>5.1f}%\")\nprint(f\"{'═'*105}\")\n\nbest = max(all_results, key=lambda r: r[\"valid_f1_weighted\"])\nprint(f\"\\n★ Meilleur seuil (wF1 valides): {best['threshold']:.2f}  \"\n      f\"— Acc={best['valid_accuracy']:.2f}%, \"\n      f\"wF1={best['valid_f1_weighted']:.4f}, \"\n      f\"F1m={best['valid_f1_macro']:.4f}\")\n\n# ── Visualisation ────────────────────────────────────────────\nfig, axes = plt.subplots(1, 3, figsize=(16, 5))\nts = [r[\"threshold\"] for r in all_results]\n\naxes[0].plot(ts, [r[\"all_accuracy\"] for r in all_results],\n             \"b-o\", ms=5, label=\"Toutes\")\naxes[0].plot(ts, [r[\"valid_accuracy\"] for r in all_results],\n             \"r-s\", ms=5, label=\"Valides\")\naxes[0].set(title=\"Accuracy vs Seuil\", xlabel=\"Seuil\", ylabel=\"Accuracy (%)\")\naxes[0].legend(); axes[0].grid(alpha=.3)\n\naxes[1].plot(ts, [r[\"all_f1_macro\"] for r in all_results],\n             \"b-o\", ms=5, label=\"macro (all)\")\naxes[1].plot(ts, [r[\"all_f1_weighted\"] for r in all_results],\n             \"b--s\", ms=5, label=\"weighted (all)\")\naxes[1].plot(ts, [r[\"valid_f1_macro\"] for r in all_results],\n             \"r-o\", ms=5, label=\"macro (val)\")\naxes[1].plot(ts, [r[\"valid_f1_weighted\"] for r in all_results],\n             \"r--s\", ms=5, label=\"weighted (val)\")\naxes[1].set(title=\"F1 vs Seuil\", xlabel=\"Seuil\", ylabel=\"F1\")\naxes[1].legend(); axes[1].grid(alpha=.3)\n\naxes[2].plot(ts, [r[\"cascade_pct\"] for r in all_results], \"g-o\", ms=5)\naxes[2].set(title=\"% images traitées par mini-modèles\",\n            xlabel=\"Seuil\", ylabel=\"%\")\naxes[2].grid(alpha=.3)\n\nplt.tight_layout(); plt.show()\n\n# ── Top-K Accuracy (depuis le cache, pas de DataLoader) ──────\ndef compute_topk_from_cache(cache, ks=(1, 3, 5, 10), valid_cls=None):\n    logits = cache[\"big_logits\"]\n    true = cache[\"targets\"]\n    if valid_cls is not None:\n        mask = np.isin(true, valid_cls)\n        logits, true = logits[mask], true[mask]\n    res = {}\n    for k in ks:\n        _, topk = logits.topk(k, dim=1)\n        topk_np = topk.numpy()\n        ok = sum(1 for i, t in enumerate(true) if t in topk_np[i])\n        res[f\"top{k}\"] = 100 * ok / len(true)\n        print(f\"  Top-{k}: {res[f'top{k}']:.2f}%\")\n    return res\n\nprint(\"\\nTop-K Accuracy (classes valides, Big Model) :\")\ntopk = compute_topk_from_cache(cache, valid_cls=valid_class_idxs)\n","metadata":{"execution":{"iopub.status.busy":"2026-02-10T21:32:52.670647Z","iopub.execute_input":"2026-02-10T21:32:52.671343Z","execution_failed":"2026-02-10T21:36:25.916Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 10. Analyse coûts computationnels & seuil optimal","metadata":{}},{"cell_type":"code","source":"\n!pip install -q calflops\n\nfrom calflops import calculate_flops\n\ndef get_flops(model, device=DEVICE):\n    m = model.module if hasattr(model, \"module\") else model\n    m = m.to(device).eval()\n    try:\n        f, _, _ = calculate_flops(m, input_shape=(1, 3, 224, 224),\n                                  print_results=False, output_as_string=False)\n        return f\n    except Exception:\n        return sum(p.numel() for p in m.parameters()) * 2\n\n# ── FLOPs par modèle ─────────────────────────────────────────\nmini_flops = {c: get_flops(cascade.mini_models[str(c)])\n              for c in cascade.dominant_classes}\ncost_mini_avg = np.mean(list(mini_flops.values()))\ncost_big = get_flops(cascade.big_model)\nprint(f\"Mini moyen: {cost_mini_avg/1e9:.4f} GFLOPs | Big: {cost_big/1e9:.4f} GFLOPs | \"\n      f\"Ratio: {cost_big/cost_mini_avg:.1f}x\")\n\nfor c in cascade.dominant_classes:\n    name = class_names_readable.get(str(c), str(c))[:15]\n    print(f\"  Mini-{c} ({name}): {mini_flops[c]/1e9:.4f} GFLOPs\")\n\n# ══════════════════════════════════════════════════════════════\n# ANALYSE : Impact du nombre de mini-modèles dans la cascade\n# ══════════════════════════════════════════════════════════════\nbest_th = best[\"threshold\"]\nn_minis = len(cascade.dominant_classes)\nprint(f\"\\nAnalyse de profondeur de cascade (seuil={best_th:.2f}, \"\n      f\"de 0 à {n_minis} mini-modèles)\")\n\ndepth_results = []\n\nfor n in range(n_minis + 1):  # 0, 1, 2, ..., N\n    subset_classes = cascade.dominant_classes[:n]\n\n    handled = np.zeros(len(cache[\"targets\"]), dtype=bool)\n    final = np.full(len(cache[\"targets\"]), -1, dtype=int)\n    flops_per_img = np.zeros(len(cache[\"targets\"]))\n\n    # Appliquer les n premiers mini-modèles\n    for c in subset_classes:\n        rem = ~handled\n        flops_per_img[rem] += mini_flops[c]\n        triggered = (cache[\"mini_probs\"][c] > best_th) & rem\n        final[triggered] = c\n        handled[triggered] = True\n\n    # Fallback Big Model\n    fb = ~handled\n    flops_per_img[fb] += cost_big\n    final[fb] = cache[\"big_preds\"][fb]\n\n    y_true, y_pred = cache[\"targets\"], final\n    m_all = compute_all_metrics(y_true, y_pred, prefix=\"all_\")\n    mask_valid = np.isin(y_true, valid_class_idxs)\n    m_valid = compute_all_metrics(\n        y_true[mask_valid], y_pred[mask_valid], prefix=\"valid_\")\n\n    avg_gflops = flops_per_img.mean() / 1e9\n    cascade_pct = 100 * handled.mean()\n\n    names_used = [class_names_readable.get(str(c), str(c))[:12]\n                  for c in subset_classes]\n\n    depth_results.append({\n        \"n_mini\": n,\n        \"classes\": names_used,\n        **m_all, **m_valid,\n        \"avg_gflops\": avg_gflops,\n        \"cascade_pct\": cascade_pct,\n    })\n\n# ── Tableau de résultats ─────────────────────────────────────\nprint(f\"\\n{'═'*115}\")\nprint(f\"  IMPACT DU NOMBRE DE MINI-MODÈLES (seuil={best_th:.2f})\")\nprint(f\"{'═'*115}\")\nprint(f\"  {'#Mini':>5} │ {'Acc(all)':>9} │ {'F1m(all)':>9} │ \"\n      f\"{'wF1(all)':>9} │ {'Acc(val)':>9} │ {'F1m(val)':>9} │ \"\n      f\"{'wF1(val)':>9} │ {'GFLOPs':>8} │ {'%Mini':>6} │ Mini-modèles\")\nprint(f\"  {'─'*5}─┼─{'─'*9}─┼─{'─'*9}─┼─{'─'*9}─┼─\"\n      f\"{'─'*9}─┼─{'─'*9}─┼─{'─'*9}─┼─{'─'*8}─┼─{'─'*6}─┼─{'─'*20}\")\n\nfor r in depth_results:\n    cls_str = \", \".join(r[\"classes\"]) if r[\"classes\"] else \"(Big seul)\"\n    print(f\"  {r['n_mini']:>5} │ {r['all_accuracy']:>8.2f}% │ \"\n          f\"{r['all_f1_macro']:>9.4f} │ {r['all_f1_weighted']:>9.4f} │ \"\n          f\"{r['valid_accuracy']:>8.2f}% │ {r['valid_f1_macro']:>9.4f} │ \"\n          f\"{r['valid_f1_weighted']:>9.4f} │ {r['avg_gflops']:>8.4f} │ \"\n          f\"{r['cascade_pct']:>5.1f}% │ {cls_str}\")\nprint(f\"{'═'*115}\")\n\n# ── Gain / Perte vs Big Model seul ───────────────────────────\nbaseline = depth_results[0]  # 0 mini = Big Model seul\nsign_f = lambda x: f\"+{x:.4f}\" if x >= 0 else f\"{x:.4f}\"\nsign_p = lambda x: f\"+{x:.2f}%\" if x >= 0 else f\"{x:.2f}%\"\n\nprint(f\"\\n{'═'*95}\")\nprint(f\"  GAIN / PERTE vs BIG MODEL SEUL (0 mini-modèles)\")\nprint(f\"{'═'*95}\")\nprint(f\"  {'#Mini':>5} │ {'ΔAcc(all)':>10} │ {'ΔF1m(all)':>10} │ \"\n      f\"{'ΔwF1(all)':>10} │ {'ΔAcc(val)':>10} │ {'ΔGFLOPs':>10} │ {'Speedup':>8}\")\nprint(f\"  {'─'*5}─┼─{'─'*10}─┼─{'─'*10}─┼─{'─'*10}─┼─\"\n      f\"{'─'*10}─┼─{'─'*10}─┼─{'─'*8}\")\n\nfor r in depth_results:\n    d_acc  = r[\"all_accuracy\"]     - baseline[\"all_accuracy\"]\n    d_f1m  = r[\"all_f1_macro\"]     - baseline[\"all_f1_macro\"]\n    d_f1w  = r[\"all_f1_weighted\"]  - baseline[\"all_f1_weighted\"]\n    d_accv = r[\"valid_accuracy\"]   - baseline[\"valid_accuracy\"]\n    d_gf   = r[\"avg_gflops\"]       - baseline[\"avg_gflops\"]\n    speedup = baseline[\"avg_gflops\"] / max(r[\"avg_gflops\"], 1e-9)\n    print(f\"  {r['n_mini']:>5} │ {sign_p(d_acc):>10} │ {sign_f(d_f1m):>10} │ \"\n          f\"{sign_f(d_f1w):>10} │ {sign_p(d_accv):>10} │ \"\n          f\"{sign_f(d_gf):>10} │ {speedup:>7.2f}x\")\nprint(f\"{'═'*95}\")\n\n# ── Visualisation ────────────────────────────────────────────\nfig, axes = plt.subplots(1, 3, figsize=(16, 5))\nns = [r[\"n_mini\"] for r in depth_results]\n\naxes[0].plot(ns, [r[\"all_accuracy\"] for r in depth_results],\n             \"b-o\", ms=6, label=\"Toutes\")\naxes[0].plot(ns, [r[\"valid_accuracy\"] for r in depth_results],\n             \"r-s\", ms=6, label=\"Valides\")\naxes[0].set(title=\"Accuracy vs Nb mini-modèles\",\n            xlabel=\"Nb mini-modèles\", ylabel=\"Accuracy (%)\")\naxes[0].set_xticks(ns); axes[0].legend(); axes[0].grid(alpha=.3)\n\naxes[1].plot(ns, [r[\"all_f1_weighted\"] for r in depth_results],\n             \"b-o\", ms=6, label=\"wF1 (all)\")\naxes[1].plot(ns, [r[\"valid_f1_weighted\"] for r in depth_results],\n             \"r-s\", ms=6, label=\"wF1 (val)\")\naxes[1].plot(ns, [r[\"all_f1_macro\"] for r in depth_results],\n             \"b--^\", ms=5, label=\"F1m (all)\")\naxes[1].set(title=\"F1 vs Nb mini-modèles\",\n            xlabel=\"Nb mini-modèles\", ylabel=\"F1\")\naxes[1].set_xticks(ns); axes[1].legend(); axes[1].grid(alpha=.3)\n\naxes[2].plot(ns, [r[\"avg_gflops\"] for r in depth_results], \"g-o\", ms=6)\naxes[2].axhline(baseline[\"avg_gflops\"], color=\"gray\", ls=\"--\",\n                alpha=.5, label=\"Big seul\")\naxes[2].set(title=\"Coût (GFLOPs) vs Nb mini-modèles\",\n            xlabel=\"Nb mini-modèles\", ylabel=\"GFLOPs moyens\")\naxes[2].set_xticks(ns); axes[2].legend(); axes[2].grid(alpha=.3)\n\nplt.tight_layout(); plt.show()\n\n# ── Pareto : Accuracy vs Coût ────────────────────────────────\nfig, ax = plt.subplots(figsize=(8, 6))\nfor r in depth_results:\n    ax.scatter(r[\"avg_gflops\"], r[\"valid_accuracy\"], s=120, zorder=5)\n    ax.annotate(f\"  {r['n_mini']} mini\",\n                (r[\"avg_gflops\"], r[\"valid_accuracy\"]),\n                fontsize=10, va=\"center\")\nax.set(title=\"Pareto : Accuracy (valides) vs Coût computationnel\",\n       xlabel=\"GFLOPs moyens par image\", ylabel=\"Accuracy (%)\")\nax.grid(alpha=.3)\nplt.tight_layout(); plt.show()\n","metadata":{"execution":{"iopub.status.busy":"2026-02-10T21:26:20.02507Z","iopub.execute_input":"2026-02-10T21:26:20.02585Z","iopub.status.idle":"2026-02-10T21:26:24.481367Z","shell.execute_reply.started":"2026-02-10T21:26:20.025818Z","shell.execute_reply":"2026-02-10T21:26:24.480632Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport numpy as np\n\n\n# Préparation des données pour le graphique\nx_labels = []\navg_cost_mflops = []\naccuracies = []\n\nfor r in depth_results:\n    if r[\"n_mini\"] == 0:\n        label = \"Big Model Seul\"\n    else:\n        # On prend juste la dernière classe ajoutée pour l'étiquette\n        last_class = r[\"classes\"][-1]\n        label = f\"+ Mini {r['n_mini']-1} ({last_class})\"\n    x_labels.append(label)\n    # Le graphique exemple utilise MFLOPs, convertissons GFLOPs en MFLOPs\n    avg_cost_mflops.append(r[\"avg_gflops\"])\n    accuracies.append(r[\"valid_accuracy\"])\n\nx = np.arange(len(x_labels))  # Positions des barres sur l'axe X\n\n# Création de la figure et du premier axe (pour le coût - Barres)\nfig, ax1 = plt.subplots(figsize=(10, 6))\n\n# Tracer les barres pour le coût moyen (Axe Y gauche)\nbars = ax1.bar(x, avg_cost_mflops, color='skyblue', alpha=0.7, label='Coût Moyen (GFLOPs)', width=0.4)\nax1.set_ylabel('Coût Moyen (GFLOPs)', color='skyblue', fontsize=12, fontweight='bold')\nax1.tick_params(axis='y', labelcolor='skyblue')\n\n# Création du deuxième axe (pour la précision - Ligne) partageant le même axe X\nax2 = ax1.twinx()\n\n# Tracer la ligne pour la précision (Axe Y droit)\nline = ax2.plot(x, accuracies, color='red', marker='o', linewidth=2, markersize=8, label='Précision (%)')\nax2.set_ylabel('Précision (%)', color='red', fontsize=12, fontweight='bold')\nax2.tick_params(axis='y', labelcolor='red')\n\n# Configuration de l'axe X\nax1.set_xticks(x)\nax1.set_xticklabels(x_labels, rotation=0, fontsize=10)\nax1.set_xlabel(\"Étapes de la cascade\", fontsize=12)\n\n# Titre du graphique\nplt.title(\"Impact de l'ajout progressif de modèles dans la cascade (Coût vs Précision)\", fontsize=14, fontweight='bold')\n\n# Ajout d'une grille\nax1.grid(axis='y', linestyle='--', alpha=0.5)\n\n# Création de la légende combinée\n# On récupère les handles (objets graphiques) et labels des deux axes\nhandles1, labels1 = ax1.get_legend_handles_labels()\nhandles2, labels2 = ax2.get_legend_handles_labels()\n# On combine les handles et labels pour une seule légende\nax1.legend(handles1 + handles2, labels1 + labels2, loc='upper center', bbox_to_anchor=(0.5, -0.15), ncol=2, fontsize=10)\n\n# Ajustement de la mise en page pour éviter que les étiquettes ne soient coupées\nplt.tight_layout()\n\n# Affichage du graphique\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}