{"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":[{"sourceType":"competition","sourceId":14774,"databundleVersionId":875431}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true},"papermill":{"default_parameters":{},"duration":861.565096,"end_time":"2026-04-19T00:39:07.013128+00:00","environment_variables":{},"exception":true,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2026-04-19T00:24:45.448032+00:00","version":"2.7.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"e924c173","cell_type":"markdown","source":"# Reduced AMCA — Diabetic Retinopathy (APTOS 2019) — **v8 : RGB + stabilité entraînement (set_lr, warmup BN, LR schedule)**\n\nNotebook avec **vrai pruning structurel** (suppression physique des filtres) et **mesure réelle des FLOPs** via `tf.profiler` au lieu de la table de référence statique (qui était fausse).\n\n### Changements clés vs v3/v5\n\n| # | Problème dans v3/v5 | Correction v6 |\n|---|---|---|\n| 1 | **Pruning = zero-masking** : filtres mis à 0 mais toujours calculés → FLOPs INCHANGÉS | **Pruning structurel** : reconstruction du backbone avec largeurs réduites + transfert de poids slicé |\n| 2 | **FLOPs annoncés** = table statique (5.9G / 4.1G / 2.3G) — incorrects | **FLOPs mesurés** par `tf.profiler` (vraies valeurs : 7.76G / 3.50G) |\n| 3 | Grad-CAM KeyError sur backbone imbriqué | grad_model reconstruit avec aux_bb à 2 sorties + copie des poids |\n\n### Stratégie de pruning structurel\n\nResNet50 utilise des **bottleneck blocks** avec connexions résiduelles. On ne peut pas pruner n'importe quoi sans casser l'addition résiduelle (`output = block(x) + shortcut`). Règles appliquées :\n\n- ✅ **Prunable** : couches conv internes des blocs bottleneck (`*_1_conv` et `*_2_conv`) — ne touchent pas la dimension de sortie du bloc\n- ❌ **Non prunable** : `*_3_conv` (sortie du bloc), `*_0_conv` (shortcut), `conv1_conv` (stem)\n\nQuand on prune les sorties de `_1_conv`, on slicer aussi les **entrées** de `_2_conv` (et idem pour `_2_conv` → `_3_conv`). Les BatchNorm sont slicées sur leurs 4 stats (gamma, beta, mean, variance).\n\n### Résultats attendus (validés sur architecture randomisée)\n\n| Modèle | Params | FLOPs réels |\n|---|---|---|\n| Baseline ResNet50 complet | 25.3 M | 7.76 G |\n| AMCA Stage2 (cutting) | 1.89 M | 3.50 G (−55 %) |\n| **AMCA Stage2 + pruning struct.** | **1.36 M** | **2.30 G (−70 %)** |\n","metadata":{"papermill":{"duration":0.007807,"end_time":"2026-04-19T00:24:48.160142+00:00","exception":false,"start_time":"2026-04-19T00:24:48.152335+00:00","status":"completed"},"tags":[]}},{"id":"f544ddb8","cell_type":"markdown","source":"## 1 — Imports","metadata":{"papermill":{"duration":0.007294,"end_time":"2026-04-19T00:24:48.173844+00:00","exception":false,"start_time":"2026-04-19T00:24:48.166550+00:00","status":"completed"},"tags":[]}},{"id":"47d38fca","cell_type":"code","source":"import os, cv2, math, random\nimport numpy as np\nimport pandas as pd\nimport tensorflow as tf\nimport matplotlib.pyplot as plt\nimport matplotlib.cm as cm\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.utils.class_weight import compute_class_weight\nfrom sklearn.metrics import (classification_report, confusion_matrix,\n                             roc_auc_score, roc_curve, accuracy_score)\n\nSEED = 42\nrandom.seed(SEED); np.random.seed(SEED); tf.random.set_seed(SEED)\nos.environ[\"PYTHONHASHSEED\"] = str(SEED)\n# Déterminisme renforcé (TF ≥ 2.9)\ntry:\n    tf.config.experimental.enable_op_determinism()\nexcept Exception:\n    pass\n\nAUTOTUNE = tf.data.AUTOTUNE\nprint(f\"TF {tf.__version__}  |  GPUs: {tf.config.list_physical_devices('GPU')}\")\n","metadata":{"execution":{"iopub.execute_input":"2026-04-19T00:24:48.188377Z","iopub.status.busy":"2026-04-19T00:24:48.187644Z","iopub.status.idle":"2026-04-19T00:25:19.585755Z","shell.execute_reply":"2026-04-19T00:25:19.584897Z"},"papermill":{"duration":31.414009,"end_time":"2026-04-19T00:25:19.594099+00:00","exception":false,"start_time":"2026-04-19T00:24:48.180090+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"6c76166f","cell_type":"markdown","source":"## 2 — Configuration","metadata":{"papermill":{"duration":0.005884,"end_time":"2026-04-19T00:25:19.606158+00:00","exception":false,"start_time":"2026-04-19T00:25:19.600274+00:00","status":"completed"},"tags":[]}},{"id":"775f8829","cell_type":"code","source":"# ── Auto-détection du chemin Kaggle ─────────────────────────────\ndef find_dataset_path():\n    candidates = [\n        \"/kaggle/input/competitions/aptos2019-blindness-detection\",\n        \"/kaggle/input/aptos2019-blindness-detection\",\n        \"/kaggle/input/aptos-2019-blindness-detection\",\n        \"./aptos2019\",\n    ]\n    for path in candidates:\n        if os.path.exists(os.path.join(path, \"train.csv\")):\n            print(f\"[PATH] Dataset trouvé : {path}\")\n            return path\n    if os.path.exists(\"/kaggle/input\"):\n        available = os.listdir(\"/kaggle/input\")\n        print(f\"[PATH] Dossiers disponibles dans /kaggle/input : {available}\")\n        for folder in available:\n            p = f\"/kaggle/input/{folder}\"\n            if os.path.exists(os.path.join(p, \"train.csv\")):\n                print(f\"[PATH] Trouvé dans : {p}\")\n                return p\n    raise FileNotFoundError(\n        \"Dataset APTOS introuvable.\\n\"\n        \"Ajoute le dataset via : Add Data → aptos2019-blindness-detection\"\n    )\n\nDATA_DIR    = find_dataset_path()\nCSV_PATH    = os.path.join(DATA_DIR, \"train.csv\")\nIMAGE_DIR   = os.path.join(DATA_DIR, \"train_images\")\nPREPROC_DIR = \"/kaggle/working/preprocessed\"\n\nIMG_SIZE        = 224  # Taille native ResNet50\nBATCH_PER_GPU   = 16\nARCHITECTURE    = \"resnet50\"\nEPOCHS_PHASE1   = 15\nEPOCHS_PHASE2   = 10\nEPOCHS_PHASE3   = 8         # Budget total pour Phase 3 (somme p3a+p3b+p3c)\nEPOCHS_FINETUNE = 15  # Fine-tune recovery plus long\nPRUNE_RATIO     = 0.30\n\n# FIX #4 : élargir la fenêtre de recherche pour comparer plusieurs candidats\nFLOPS_BUDGET    = 6.0   # inclut Stage4 (5.9G)\nMIN_FLOPS       = 2.0   # inclut Stage2 (2.3G), Stage3 (4.1G), Stage4 (5.9G)\n\nprint(f\"CSV    : {CSV_PATH}\")\nprint(f\"Images : {IMAGE_DIR}\")\nprint(f\"Architecture : {ARCHITECTURE}  |  Epochs P1={EPOCHS_PHASE1}/P2={EPOCHS_PHASE2}/P3={EPOCHS_PHASE3}\")\nprint(f\"AMCA window  : [{MIN_FLOPS}, {FLOPS_BUDGET}] GFLOPs\")\n","metadata":{"execution":{"iopub.execute_input":"2026-04-19T00:25:19.619569Z","iopub.status.busy":"2026-04-19T00:25:19.619114Z","iopub.status.idle":"2026-04-19T00:25:19.628027Z","shell.execute_reply":"2026-04-19T00:25:19.627214Z"},"papermill":{"duration":0.017406,"end_time":"2026-04-19T00:25:19.629483+00:00","exception":false,"start_time":"2026-04-19T00:25:19.612077+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"657902ca","cell_type":"markdown","source":"## 3 — Stratégie GPU/TPU","metadata":{"papermill":{"duration":0.00588,"end_time":"2026-04-19T00:25:19.641893+00:00","exception":false,"start_time":"2026-04-19T00:25:19.636013+00:00","status":"completed"},"tags":[]}},{"id":"0b822f26","cell_type":"code","source":"def setup_strategy():\n    # FIX mineur : except spécifique au lieu de bare except\n    try:\n        tpu = tf.distribute.cluster_resolver.TPUClusterResolver()\n        tf.config.experimental_connect_to_cluster(tpu)\n        tf.tpu.experimental.initialize_tpu_system(tpu)\n        s = tf.distribute.TPUStrategy(tpu)\n        print(f\"TPU — {s.num_replicas_in_sync} cores\"); return s\n    except Exception:  # TPU non disponible — bascule sur GPU/CPU\n        pass\n    gpus = tf.config.list_physical_devices(\"GPU\")\n    if len(gpus) > 1:\n        s = tf.distribute.MirroredStrategy()\n        print(f\"Multi-GPU — {len(gpus)} GPUs\"); return s\n    s = tf.distribute.get_strategy()\n    print(f\"Single device — {gpus[0].name if gpus else 'CPU'}\"); return s\n\nstrategy     = setup_strategy()\nGLOBAL_BATCH = BATCH_PER_GPU * strategy.num_replicas_in_sync\nprint(f\"Batch global : {GLOBAL_BATCH}\")\n","metadata":{"execution":{"iopub.execute_input":"2026-04-19T00:25:19.655375Z","iopub.status.busy":"2026-04-19T00:25:19.655148Z","iopub.status.idle":"2026-04-19T00:25:20.160071Z","shell.execute_reply":"2026-04-19T00:25:20.159165Z"},"papermill":{"duration":0.513665,"end_time":"2026-04-19T00:25:20.161828+00:00","exception":false,"start_time":"2026-04-19T00:25:19.648163+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"db3934c3","cell_type":"markdown","source":"## 4 — Preprocessing Offline (OpenCV)\n\nPipeline : crop bordures noires → resize → **CLAHE sur chaque canal RGB** → **Ben Graham RGB** → [0,1].\n\n**Corrections v7 :**\n- ✅ Suppression du canal vert uniquement → RGB complet préservé\n- ✅ CLAHE appliqué indépendamment sur R, G, B\n- ✅ Ben Graham sur l'image RGB complète (hémorragies canal R, néovaisseaux canal B)\n","metadata":{"papermill":{"duration":0.006724,"end_time":"2026-04-19T00:25:20.175327+00:00","exception":false,"start_time":"2026-04-19T00:25:20.168603+00:00","status":"completed"},"tags":[]}},{"id":"8cddbf5c","cell_type":"code","source":"def crop_black_borders(img, tol=7):\n    gray = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY)\n    _, mask = cv2.threshold(gray, tol, 255, cv2.THRESH_BINARY)\n    cnts, _ = cv2.findContours(mask, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n    if not cnts: return img\n    x, y, w, h = cv2.boundingRect(max(cnts, key=cv2.contourArea))\n    return img[y:y+h, x:x+w]\n\n# ✅ CORRECTION #1 : Ben Graham sur RGB complet\ndef ben_graham_rgb(img, size=224):\n    \"\"\"\n    Applique Ben Graham sur les 3 canaux RGB simultanément.\n    Préserve l'information de couleur :\n      - Canal R : hémorragies\n      - Canal G : vaisseaux / exsudats (meilleur contraste général)\n      - Canal B : néovaisseaux (Proliferative DR)\n    \"\"\"\n    blur = cv2.GaussianBlur(img, (0, 0), sigmaX=size // 30)\n    out  = cv2.addWeighted(img, 4, blur, -4, 128)\n    return out\n\n# ✅ CORRECTION #2 : CLAHE par canal RGB + Ben Graham RGB\ndef preprocess_single(path, size=224):\n    \"\"\"\n    Pipeline corrigé :\n      crop → resize → CLAHE sur chaque canal RGB → Ben Graham RGB → [0,1]\n\n    Changements vs version précédente :\n    - Suppression de l'extraction du canal vert uniquement\n    - CLAHE appliqué indépendamment sur R, G, B\n    - Ben Graham sur l'image RGB complète\n    - Lissage post-traitement conservé (réduit bruit CLAHE)\n\n    ⚠️  Après ce changement : supprimer /kaggle/working/preprocessed\n        pour régénérer les .npy avec RGB au lieu de mono-canal.\n    \"\"\"\n    img = cv2.imread(path)\n    if img is None: raise FileNotFoundError(path)\n    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    img = crop_black_borders(img)\n    img = cv2.resize(img, (size, size), interpolation=cv2.INTER_AREA)\n\n    # CLAHE indépendant sur chaque canal R, G, B\n    clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))\n    channels = []\n    for c in range(3):\n        ch = clahe.apply(img[:, :, c])\n        channels.append(ch)\n    img_clahe = np.stack(channels, axis=-1)   # (H, W, 3) uint8\n\n    # Ben Graham sur l'image RGB complète\n    img_bg = ben_graham_rgb(img_clahe, size)\n\n    # Lissage léger post-traitement (réduit artefacts CLAHE)\n    img_bg = cv2.GaussianBlur(img_bg, (3, 3), sigmaX=0.5)\n\n    # Clip + normalisation [0, 1]\n    out = np.clip(img_bg, 0, 255).astype(np.float32) / 255.0\n    return out   # (H, W, 3) float32 — 3 canaux RGB distincts\n\ndef preprocess_dataset(df, out_dir, size=224):\n    os.makedirs(out_dir, exist_ok=True)\n    paths, err = [], 0\n    for _, row in df.iterrows():\n        p = os.path.join(out_dir, f\"{row['id_code']}.npy\")\n        if not os.path.exists(p):\n            try:\n                np.save(p, preprocess_single(row[\"image_path\"], size))\n            except Exception as e:\n                print(f\"[WARN] {row['id_code']}: {e}\")\n                p = None; err += 1\n        paths.append(p)\n    df = df.copy()\n    df[\"prep_path\"] = paths\n    df = df[df[\"prep_path\"].notna()].reset_index(drop=True)\n    print(f\"[PREPROC] {len(df)} images OK — {err} erreurs\")\n    return df\n\nprint(\"Fonctions preprocessing définies (RGB complet — v7).\")\n","metadata":{"execution":{"iopub.execute_input":"2026-04-19T00:25:20.189267Z","iopub.status.busy":"2026-04-19T00:25:20.188602Z","iopub.status.idle":"2026-04-19T00:25:20.200709Z","shell.execute_reply":"2026-04-19T00:25:20.199804Z"},"papermill":{"duration":0.020472,"end_time":"2026-04-19T00:25:20.202187+00:00","exception":false,"start_time":"2026-04-19T00:25:20.181715+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"406caa5f","cell_type":"markdown","source":"## 5 — Chargement & Split","metadata":{"papermill":{"duration":0.006156,"end_time":"2026-04-19T00:25:20.214438+00:00","exception":false,"start_time":"2026-04-19T00:25:20.208282+00:00","status":"completed"},"tags":[]}},{"id":"c4a0305a","cell_type":"code","source":"def load_df(csv, img_dir):\n    df = pd.read_csv(csv)\n    df[\"label\"]      = df[\"diagnosis\"].apply(lambda x: 0 if int(x) == 0 else 1)\n    df[\"image_path\"] = df[\"id_code\"].apply(lambda x: os.path.join(img_dir, f\"{x}.png\"))\n    df = df[df[\"image_path\"].apply(os.path.exists)].reset_index(drop=True)\n    print(f\"[DATA] {len(df)} images  |  {dict(df['label'].value_counts())}\")\n    return df\n\ndf = load_df(CSV_PATH, IMAGE_DIR)\ntrain_df, val_df = train_test_split(df, test_size=0.2, random_state=SEED, stratify=df[\"label\"])\ntrain_df = train_df.reset_index(drop=True)\nval_df   = val_df.reset_index(drop=True)\nprint(f\"Train: {len(train_df)}  |  Val: {len(val_df)}\")\n","metadata":{"execution":{"iopub.execute_input":"2026-04-19T00:25:20.228228Z","iopub.status.busy":"2026-04-19T00:25:20.227906Z","iopub.status.idle":"2026-04-19T00:25:24.622961Z","shell.execute_reply":"2026-04-19T00:25:24.621964Z"},"papermill":{"duration":4.4036,"end_time":"2026-04-19T00:25:24.624503+00:00","exception":false,"start_time":"2026-04-19T00:25:20.220903+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"49100f70","cell_type":"markdown","source":"## 6 — Lancer le Preprocessing","metadata":{"papermill":{"duration":0.007211,"end_time":"2026-04-19T00:25:24.639233+00:00","exception":false,"start_time":"2026-04-19T00:25:24.632022+00:00","status":"completed"},"tags":[]}},{"id":"eac9f32b","cell_type":"code","source":"import shutil\n\n# ✅ CORRECTION : supprime le cache mono-canal pour régénérer en RGB\nif os.path.exists(PREPROC_DIR):\n    shutil.rmtree(PREPROC_DIR)\n    print(f\"[CACHE] {PREPROC_DIR} supprimé — régénération RGB en cours...\")\n\ntrain_df = preprocess_dataset(train_df, PREPROC_DIR, IMG_SIZE)\nval_df   = preprocess_dataset(val_df,   PREPROC_DIR, IMG_SIZE)\n\nclasses  = np.array([0, 1])\nweights  = compute_class_weight(\"balanced\", classes=classes, y=train_df[\"label\"].values)\nclass_weight = {int(c): float(w) for c, w in zip(classes, weights)}\nprint(f\"Class weights: {class_weight}\")\n","metadata":{"execution":{"iopub.execute_input":"2026-04-19T00:25:24.654769Z","iopub.status.busy":"2026-04-19T00:25:24.654239Z","iopub.status.idle":"2026-04-19T00:32:52.361533Z","shell.execute_reply":"2026-04-19T00:32:52.360675Z"},"papermill":{"duration":447.723232,"end_time":"2026-04-19T00:32:52.369804+00:00","exception":false,"start_time":"2026-04-19T00:25:24.646572+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"04500c81","cell_type":"markdown","source":"## 7 — Augmentation & tf.data\n\n**Corrections v7+v8 :**\n- ✅ `tf.roll` retiré → remplacé par `tf.pad(REFLECT) + random_crop`\n- ✅ Flip horizontal conservé (latéralité OD/OG valide)\n- ✅ Zoom léger conservé (variation distance patient-caméra)\n- ✅ Bruit gaussien σ=0.01 conservé\n","metadata":{"papermill":{"duration":0.006523,"end_time":"2026-04-19T00:32:52.382920+00:00","exception":false,"start_time":"2026-04-19T00:32:52.376397+00:00","status":"completed"},"tags":[]}},{"id":"4d48a35f","cell_type":"code","source":"def decode_npy(path, label):\n    img = tf.numpy_function(\n        lambda p: np.load(p.decode()).astype(np.float32), [path], tf.float32\n    )\n    img.set_shape([IMG_SIZE, IMG_SIZE, 3])\n    return img, tf.cast(label, tf.float32)\n\n# ✅ CORRECTION #3 : tf.roll → pad REFLECT + random_crop\ndef augment(img, lbl):\n    \"\"\"\n    Augmentation adaptée rétine — v7 corrigée :\n    - tf.roll supprimé (wrap circulaire → artefacts sur fond noir)\n    - Remplacement par pad REFLECT + random_crop (propre, sans bords noirs)\n    \"\"\"\n    # 1. Flip horizontal (OD/OG — valide anatomiquement)\n    img = tf.image.random_flip_left_right(img)\n\n    # 2. ✅ Translation ±5% via padding REFLECT + crop aléatoire\n    #    REFLECT évite les bords noirs artificiels contrairement à tf.roll\n    pad = int(IMG_SIZE * 0.05)   # ~11px pour IMG_SIZE=224\n    img = tf.pad(img, [[pad, pad], [pad, pad], [0, 0]], mode=\"REFLECT\")\n    img = tf.image.random_crop(img, size=[IMG_SIZE, IMG_SIZE, 3])\n\n    # 3. Zoom léger ×0.9–1.1 (variation de distance patient-caméra)\n    s = tf.random.uniform([], 0.90, 1.10)\n    h = tf.cast(tf.cast(IMG_SIZE, tf.float32) * s, tf.int32)\n    h = tf.clip_by_value(h, IMG_SIZE // 2, int(IMG_SIZE * 1.5))\n    img = tf.image.resize(img, [h, h])\n    img = tf.image.resize_with_crop_or_pad(img, IMG_SIZE, IMG_SIZE)\n\n    # 4. Bruit gaussien faible (σ=0.01 — simulation artefacts optiques)\n    noise = tf.random.normal(shape=tf.shape(img), mean=0.0, stddev=0.01)\n    img   = img + noise\n\n    img = tf.clip_by_value(img, 0.0, 1.0)\n    return img, lbl\n\ndef build_ds(df, batch, training=False):\n    ds = tf.data.Dataset.from_tensor_slices(\n        (df[\"prep_path\"].values, df[\"label\"].values)\n    )\n    ds = ds.map(decode_npy, num_parallel_calls=AUTOTUNE)\n    if training:\n        ds = ds.shuffle(len(df), seed=SEED, reshuffle_each_iteration=True)\n        ds = ds.map(augment, num_parallel_calls=AUTOTUNE)\n    return ds.batch(batch).prefetch(AUTOTUNE)\n\ntrain_ds = build_ds(train_df, GLOBAL_BATCH, training=True)\nval_ds   = build_ds(val_df,   GLOBAL_BATCH, training=False)\nprint(f\"Datasets OK  |  Train={len(train_df)}  Val={len(val_df)}\")\n","metadata":{"execution":{"iopub.execute_input":"2026-04-19T00:32:52.397283Z","iopub.status.busy":"2026-04-19T00:32:52.397007Z","iopub.status.idle":"2026-04-19T00:32:54.430852Z","shell.execute_reply":"2026-04-19T00:32:54.430146Z"},"papermill":{"duration":2.043133,"end_time":"2026-04-19T00:32:54.432589+00:00","exception":false,"start_time":"2026-04-19T00:32:52.389456+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"37c1ce53","cell_type":"markdown","source":"## 8 — Backbone Registry","metadata":{"papermill":{"duration":0.006237,"end_time":"2026-04-19T00:32:54.445828+00:00","exception":false,"start_time":"2026-04-19T00:32:54.439591+00:00","status":"completed"},"tags":[]}},{"id":"bdd81c4e","cell_type":"code","source":"# ─────────────────────────────────────────────────────────────────\n# BACKBONE_REGISTRY — annuaire des backbones supportés par l'AMCA\n#\n# Rôle : centraliser en un seul endroit tout ce que l'AMCA doit\n# savoir sur chaque architecture, pour éviter de dupliquer du code\n# dans chaque fonction (build, truncate, probe, etc.)\n#\n# Champs par entrée :\n#   builder    : classe Keras à instancier\n#   input_size : résolution native optimale (pixels)\n#   candidates : liste des cutting points candidats pour l'AMCA\n#     - name  : nom lisible du stage\n#     - layer : nom exact de la couche dans Keras (cas-sensitive)\n#     - flops : FLOPs théoriques approximatifs en GFLOPs\n#\n# NOTE sur les FLOPs :\n#   Ces valeurs sont THÉORIQUES (littérature, input 224×224, batch=1,\n#   1 GPU, pas d'overhead XLA/MirroredStrategy).\n#   Elles servent UNIQUEMENT à comparer les candidats entre eux\n#   dans le score Pareto de l'AMCA (AUC² / FLOPs).\n#   Les ratios sont préservés même si les valeurs absolues diffèrent\n#   des mesures réelles de tf.profiler.\n#   La mesure réelle est effectuée séparément avec measure_flops().\n# ─────────────────────────────────────────────────────────────────\n\ndef _compute_candidate_flops(builder, input_size, layer_name):\n    \"\"\"\n    Calcule les FLOPs réels d'un backbone tronqué à layer_name.\n    Utilisé au moment de la construction du registry pour remplacer\n    les valeurs théoriques par des mesures réelles (1 GPU, batch=1).\n    Retourne la valeur en GFLOPs.\n    \"\"\"\n    try:\n        base = builder(include_top=False, weights=None,\n                       input_shape=(input_size, input_size, 3))\n        try:\n            out = base.get_layer(layer_name).output\n        except (ValueError, KeyError):\n            return None\n        truncated = tf.keras.Model(base.input, out)\n\n        # Mesure sur un batch de 1 image\n        dummy = tf.ones((1, input_size, input_size, 3))\n        concrete = tf.function(truncated).get_concrete_function(\n            tf.TensorSpec((1, input_size, input_size, 3), tf.float32)\n        )\n        opts = tf.compat.v1.profiler.ProfileOptionBuilder.float_operation()\n        opts[\"output\"] = \"none\"\n        flops_obj = tf.compat.v1.profiler.profile(\n            concrete.graph, options=opts\n        )\n        gflops = (flops_obj.total_float_ops / 2) / 1e9\n        del base, truncated\n        tf.keras.backend.clear_session()\n        return round(gflops, 3)\n    except Exception:\n        return None\n\n\n# ── Valeurs théoriques de référence (issues de la littérature) ──\n# Utilisées comme fallback si _compute_candidate_flops() échoue.\n# Source : papiers originaux + https://github.com/albanie/convnet-burden\n_FLOPS_REF = {\n    # ResNet50 — He et al. (2015), input 224×224\n    \"conv2_block3_out\" : 1.2,\n    \"conv3_block4_out\" : 2.3,\n    \"conv4_block6_out\" : 4.1,\n    \"conv5_block3_out\" : 5.9,\n    # VGG16 — Simonyan & Zisserman (2014)\n    \"block3_pool\"      : 7.5,\n    \"block4_pool\"      : 11.3,\n    \"block5_pool\"      : 15.5,\n    # MobileNetV2 — Sandler et al. (2018)\n    \"block_13_expand_relu\" : 0.22,\n    \"out_relu\"             : 0.30,\n    # EfficientNetB0 — Tan & Le (2019)\n    \"block5c_add\"          : 0.28,\n    \"block7a_project_bn\"   : 0.39,\n    # EfficientNetB4 — Tan & Le (2019), input 380×380\n    \"block4a_expand_activation\" : 2.1,\n    # \"block5c_add\" déjà défini ci-dessus → 0.28 (B0), B4 = 3.8\n    \"block6d_add\"               : 6.2,\n    \"block7b_add\"               : 9.4,\n}\n\nBACKBONE_REGISTRY = {\n    \"resnet50\": {\n        \"builder\"   : tf.keras.applications.ResNet50,\n        \"input_size\": 224,\n        # FLOPs théoriques (GFLOPs) — littérature He et al. 2015\n        # Mesure réelle via tf.profiler dans measure_flops() après entraînement\n        \"candidates\": [\n            {\"name\": \"Stage1\", \"layer\": \"conv2_block3_out\",\n             \"flops\": _FLOPS_REF[\"conv2_block3_out\"],\n             \"flops_note\": \"théorique — réel ≈ ×1.3 avec overhead MirroredStrategy\"},\n            {\"name\": \"Stage2\", \"layer\": \"conv3_block4_out\",\n             \"flops\": _FLOPS_REF[\"conv3_block4_out\"],\n             \"flops_note\": \"théorique — réel mesuré ≈ 3.5G (×1.5 overhead)\"},\n            {\"name\": \"Stage3\", \"layer\": \"conv4_block6_out\",\n             \"flops\": _FLOPS_REF[\"conv4_block6_out\"],\n             \"flops_note\": \"théorique — réel mesuré ≈ 4.1G (overhead faible)\"},\n            {\"name\": \"Stage4\", \"layer\": \"conv5_block3_out\",\n             \"flops\": _FLOPS_REF[\"conv5_block3_out\"],\n             \"flops_note\": \"théorique — réel mesuré ≈ 7.8G (overhead ×1.3)\"},\n        ],\n    },\n    \"vgg16\": {\n        \"builder\"   : tf.keras.applications.VGG16,\n        \"input_size\": 224,\n        \"candidates\": [\n            {\"name\": \"Block3\", \"layer\": \"block3_pool\",\n             \"flops\": _FLOPS_REF[\"block3_pool\"],   \"flops_note\": \"théorique\"},\n            {\"name\": \"Block4\", \"layer\": \"block4_pool\",\n             \"flops\": _FLOPS_REF[\"block4_pool\"],   \"flops_note\": \"théorique\"},\n            {\"name\": \"Block5\", \"layer\": \"block5_pool\",\n             \"flops\": _FLOPS_REF[\"block5_pool\"],   \"flops_note\": \"théorique\"},\n        ],\n    },\n    \"mobilenetv2\": {\n        \"builder\"   : tf.keras.applications.MobileNetV2,\n        \"input_size\": 224,\n        \"candidates\": [\n            {\"name\": \"Block13\", \"layer\": \"block_13_expand_relu\",\n             \"flops\": _FLOPS_REF[\"block_13_expand_relu\"], \"flops_note\": \"théorique\"},\n            {\"name\": \"Full\",    \"layer\": \"out_relu\",\n             \"flops\": _FLOPS_REF[\"out_relu\"],             \"flops_note\": \"théorique\"},\n        ],\n    },\n    \"efficientnetb0\": {\n        \"builder\"   : tf.keras.applications.EfficientNetB0,\n        \"input_size\": 224,\n        \"candidates\": [\n            {\"name\": \"Block5\", \"layer\": \"block5c_add\",\n             \"flops\": _FLOPS_REF[\"block5c_add\"],        \"flops_note\": \"théorique\"},\n            {\"name\": \"Block7\", \"layer\": \"block7a_project_bn\",\n             \"flops\": _FLOPS_REF[\"block7a_project_bn\"], \"flops_note\": \"théorique\"},\n        ],\n    },\n    \"efficientnetb4\": {\n        \"builder\"   : tf.keras.applications.EfficientNetB4,\n        \"input_size\": 380,\n        # FLOPs pour input 380×380 — proportionnels à (380/224)² vs B0\n        \"candidates\": [\n            {\"name\": \"Block4\", \"layer\": \"block4a_expand_activation\",\n             \"flops\": _FLOPS_REF[\"block4a_expand_activation\"], \"flops_note\": \"théorique 380px\"},\n            {\"name\": \"Block5\", \"layer\": \"block5c_add\",\n             \"flops\": 3.8,  \"flops_note\": \"théorique 380px (≠ B0 block5c_add=0.28)\"},\n            {\"name\": \"Block6\", \"layer\": \"block6d_add\",\n             \"flops\": _FLOPS_REF[\"block6d_add\"], \"flops_note\": \"théorique 380px\"},\n            {\"name\": \"Block7\", \"layer\": \"block7b_add\",\n             \"flops\": _FLOPS_REF[\"block7b_add\"], \"flops_note\": \"théorique 380px\"},\n        ],\n    },\n}\n\n# ── Affichage informatif au chargement ──────────────────────────\nprint(f\"Backbone sélectionné : {ARCHITECTURE}\")\nprint(f\"Résolution           : {BACKBONE_REGISTRY[ARCHITECTURE]['input_size']}×{BACKBONE_REGISTRY[ARCHITECTURE]['input_size']}px\")\nprint(\"Cutting points candidats :\")\nfor cand in BACKBONE_REGISTRY[ARCHITECTURE][\"candidates\"]:\n    eligible = MIN_FLOPS <= cand[\"flops\"] <= FLOPS_BUDGET\n    flag = \" ← éligible\" if eligible else \" [hors fenêtre]\"\n    note = cand.get(\"flops_note\", \"\")\n    print(f\"  {cand['name']:10s}  {cand['layer']:35s}  \"\n          f\"{cand['flops']:.2f}G théorique{flag}\")\n    if note:\n        print(f\"  {'':10s}  ({note})\")\n","metadata":{"execution":{"iopub.execute_input":"2026-04-19T00:32:54.459957Z","iopub.status.busy":"2026-04-19T00:32:54.459645Z","iopub.status.idle":"2026-04-19T00:32:54.475153Z","shell.execute_reply":"2026-04-19T00:32:54.474224Z"},"papermill":{"duration":0.024265,"end_time":"2026-04-19T00:32:54.476514+00:00","exception":false,"start_time":"2026-04-19T00:32:54.452249+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"c1ea780e","cell_type":"markdown","source":"## 9 — Composants d'Architecture","metadata":{"papermill":{"duration":0.006136,"end_time":"2026-04-19T00:32:54.489135+00:00","exception":false,"start_time":"2026-04-19T00:32:54.482999+00:00","status":"completed"},"tags":[]}},{"id":"3ac63f35","cell_type":"code","source":"# ── SE Block ──────────────────────────────────────────────────\ndef se_block(x, ratio=16):\n    ch = x.shape[-1]\n    s  = tf.keras.layers.GlobalAveragePooling2D()(x)\n    s  = tf.keras.layers.Dense(ch//ratio, activation=\"relu\",\n             kernel_regularizer=tf.keras.regularizers.l2(1e-4))(s)\n    # L2 aussi sur le Dense d'expansion pour symétrie\n    s  = tf.keras.layers.Dense(ch, activation=\"sigmoid\",\n             kernel_regularizer=tf.keras.regularizers.l2(1e-4))(s)\n    s  = tf.keras.layers.Reshape((1, 1, ch))(s)\n    return tf.keras.layers.Multiply()([x, s])\n\n# ── Tête de classification ─────────────────────────────────────\ndef build_head(x, cfg):\n    x = tf.keras.layers.GlobalAveragePooling2D()(x)\n    for units, drop in zip(cfg[\"units\"], cfg[\"dropouts\"]):\n        x = tf.keras.layers.Dense(units, activation=None,\n                kernel_regularizer=tf.keras.regularizers.l2(1e-4))(x)\n        x = tf.keras.layers.BatchNormalization()(x)\n        x = tf.keras.layers.Activation(\"relu\")(x)\n        x = tf.keras.layers.Dropout(drop)(x)\n    return tf.keras.layers.Dense(1, activation=\"sigmoid\", dtype=\"float32\")(x)\n\n# ── Métriques ────────────────────────────────────────────────\ndef get_metrics():\n    return [\n        tf.keras.metrics.AUC(name=\"auc\"),\n        tf.keras.metrics.BinaryAccuracy(name=\"accuracy\"),\n        tf.keras.metrics.Precision(name=\"precision\"),\n        tf.keras.metrics.Recall(name=\"recall\"),\n    ]\n\n# ── Compile (BCE + AdamW) ─────────────────────────────────────\ndef compile_model(model, lr):\n    model.compile(\n        optimizer=tf.keras.optimizers.AdamW(learning_rate=lr, weight_decay=1e-4),\n        # ✅ CORRECTION #4 : label_smoothing 0.05 → 0.01\n        # 0.05 pénalisait trop les prédictions confiantes → plafonnait accuracy\n        loss=tf.keras.losses.BinaryCrossentropy(label_smoothing=0.01),\n        metrics=get_metrics(),\n    )\n    return model\n\n# ✅ v8 FIX #1 — set_lr : change le LR SANS recompiler\n# Préserve les moments m/v de AdamW → pas de spike au changement de phase\ndef set_lr(model, lr):\n    \"\"\"\n    Change le LR sans recompiler — compatible Keras 3.\n    tf.keras.backend.set_value() est deprecated dans Keras 3 et crash\n    car AdamW.learning_rate est un float Python et non une tf.Variable.\n    \"\"\"\n    opt = model.optimizer\n    try:\n        if hasattr(opt.learning_rate, 'assign'):\n            # learning_rate est une tf.Variable\n            opt.learning_rate.assign(float(lr))\n        else:\n            # learning_rate est un float Python\n            opt.learning_rate = float(lr)\n        print(f\"[LR] → {lr:.2e}  (optimizer AdamW conservé)\")\n    except Exception as e:\n        print(f\"[LR-WARN] assign échoué ({e}) → recompile fallback\")\n        model.compile(\n            optimizer=tf.keras.optimizers.AdamW(\n                learning_rate=float(lr), weight_decay=1e-4),\n            loss=tf.keras.losses.BinaryCrossentropy(label_smoothing=0.01),\n            metrics=get_metrics()\n        )\n        print(f\"[LR] → {lr:.2e}  (recompile fallback)\")\n\n# ✅ v8 FIX #2 — recalibrage BN sur données rétine (2 epochs, LR très faible)\n# Les stats BN ImageNet sont inadaptées aux images Ben Graham RGB\n# → on débloque les BN 2 epochs pour recalibrer mean/variance\ndef recalibrate_bn(model, backbone, train_ds, val_ds, class_weight,\n                   lr_bn=1e-6, epochs=2, ckpt_name=\"bn_calib\"):\n    \"\"\"\n    Débloque TOUTES les couches (BN inclus) pendant 2 epochs à LR ultra-faible\n    pour recalibrer les statistiques BN sur les données rétine RGB.\n    Puis regèle les BN pour la suite.\n    \"\"\"\n    print(f\"\\n[BN-CALIB] Recalibrage BN — {epochs} epochs  LR={lr_bn:.1e}\")\n    for layer in backbone.layers:\n        layer.trainable = True   # BN débloqué\n    set_lr(model, lr_bn)\n    h = model.fit(train_ds, validation_data=val_ds,\n                  epochs=epochs, class_weight=class_weight,\n                  callbacks=[tf.keras.callbacks.ModelCheckpoint(\n                      f\"{ckpt_name}.keras\", monitor=\"val_auc\", mode=\"max\",\n                      save_best_only=True, verbose=0)],\n                  verbose=1)\n    # Regeler les BN après recalibrage\n    for layer in backbone.layers:\n        if isinstance(layer, tf.keras.layers.BatchNormalization):\n            layer.trainable = False\n    print(\"[BN-CALIB] BN recalibrés et regelés\")\n    return h\n\n\n# ── Callbacks ────────────────────────────────────────────────\ndef get_callbacks(name, patience=8):\n    return [\n        tf.keras.callbacks.ModelCheckpoint(\n            f\"{name}.keras\", monitor=\"val_auc\", mode=\"max\",\n            save_best_only=True, verbose=1),\n        tf.keras.callbacks.EarlyStopping(\n            monitor=\"val_auc\", mode=\"max\",\n            patience=patience, restore_best_weights=True, verbose=1),\n        tf.keras.callbacks.ReduceLROnPlateau(\n            monitor=\"val_loss\", factor=0.3, patience=3,\n            min_lr=1e-8, verbose=1),\n    ]\n\nprint(\"Composants OK\")\n\n\n# ============================================================\n# Mesure réelle des FLOPs via tf.profiler\n# ============================================================\nimport io, contextlib\n\n@contextlib.contextmanager\ndef _silenced_io():\n    with contextlib.redirect_stdout(io.StringIO()), contextlib.redirect_stderr(io.StringIO()):\n        yield\n\ndef measure_flops(model, batch_size=1):\n    \"\"\"Mesure les VRAIS FLOPs en inférence d'un modèle Keras (TF profiler).\"\"\"\n    from tensorflow.python.framework.convert_to_constants import (\n        convert_variables_to_constants_v2_as_graph,\n    )\n    spec = [tf.TensorSpec(shape=(batch_size,) + model.input_shape[1:], dtype=tf.float32)]\n    cf = tf.function(lambda x: model(x)).get_concrete_function(*spec)\n    _, gd = convert_variables_to_constants_v2_as_graph(cf)\n    with _silenced_io():\n        with tf.Graph().as_default() as g:\n            tf.graph_util.import_graph_def(gd, name=\"\")\n            opts = tf.compat.v1.profiler.ProfileOptionBuilder.float_operation()\n            opts['output'] = 'none'\n            f = tf.compat.v1.profiler.profile(\n                graph=g, run_meta=tf.compat.v1.RunMetadata(), cmd=\"op\", options=opts\n            )\n    return f.total_float_ops if f else 0\n\n\n# ============================================================\n# Pruning structurel pour ResNet50 (cutting Stage2)\n# ============================================================\ndef bottleneck_block_pruned(x, filters, stride, name,\n                            keep_1=None, keep_2=None, shortcut_conv=False):\n    \"\"\"Reconstruit un bloc bottleneck ResNet50 avec largeurs réduites pour _1 et _2.\"\"\"\n    f1, f2, f3 = filters\n    n1 = len(keep_1) if keep_1 is not None else f1\n    n2 = len(keep_2) if keep_2 is not None else f2\n\n    if shortcut_conv:\n        sc = tf.keras.layers.Conv2D(f3, 1, strides=stride, name=f\"{name}_0_conv\")(x)\n        sc = tf.keras.layers.BatchNormalization(epsilon=1.001e-5,\n                                                name=f\"{name}_0_bn\")(sc)\n    else:\n        sc = x\n\n    x = tf.keras.layers.Conv2D(n1, 1, strides=stride, name=f\"{name}_1_conv\")(x)\n    x = tf.keras.layers.BatchNormalization(epsilon=1.001e-5, name=f\"{name}_1_bn\")(x)\n    x = tf.keras.layers.Activation(\"relu\", name=f\"{name}_1_relu\")(x)\n\n    x = tf.keras.layers.Conv2D(n2, 3, padding=\"same\", name=f\"{name}_2_conv\")(x)\n    x = tf.keras.layers.BatchNormalization(epsilon=1.001e-5, name=f\"{name}_2_bn\")(x)\n    x = tf.keras.layers.Activation(\"relu\", name=f\"{name}_2_relu\")(x)\n\n    x = tf.keras.layers.Conv2D(f3, 1, name=f\"{name}_3_conv\")(x)\n    x = tf.keras.layers.BatchNormalization(epsilon=1.001e-5, name=f\"{name}_3_bn\")(x)\n\n    x = tf.keras.layers.Add(name=f\"{name}_add\")([sc, x])\n    x = tf.keras.layers.Activation(\"relu\", name=f\"{name}_out\")(x)\n    return x\n\n\ndef build_pruned_resnet50_stage2(keep_dict, name=\"resnet50_bb_pruned\"):\n    \"\"\"Reconstruit ResNet50 jusqu'à conv3_block4_out avec pruning structurel.\"\"\"\n    inp = tf.keras.layers.Input(shape=(IMG_SIZE, IMG_SIZE, 3))\n    x = tf.keras.layers.ZeroPadding2D(((3,3),(3,3)), name=\"conv1_pad\")(inp)\n    x = tf.keras.layers.Conv2D(64, 7, strides=2, name=\"conv1_conv\")(x)\n    x = tf.keras.layers.BatchNormalization(epsilon=1.001e-5, name=\"conv1_bn\")(x)\n    x = tf.keras.layers.Activation(\"relu\", name=\"conv1_relu\")(x)\n    x = tf.keras.layers.ZeroPadding2D(((1,1),(1,1)), name=\"pool1_pad\")(x)\n    x = tf.keras.layers.MaxPooling2D(3, strides=2, name=\"pool1_pool\")(x)\n\n    # Stage 2\n    x = bottleneck_block_pruned(x, (64,64,256), 1, \"conv2_block1\",\n        keep_dict.get(\"conv2_block1_1_conv\"), keep_dict.get(\"conv2_block1_2_conv\"),\n        shortcut_conv=True)\n    x = bottleneck_block_pruned(x, (64,64,256), 1, \"conv2_block2\",\n        keep_dict.get(\"conv2_block2_1_conv\"), keep_dict.get(\"conv2_block2_2_conv\"))\n    x = bottleneck_block_pruned(x, (64,64,256), 1, \"conv2_block3\",\n        keep_dict.get(\"conv2_block3_1_conv\"), keep_dict.get(\"conv2_block3_2_conv\"))\n\n    # Stage 3\n    x = bottleneck_block_pruned(x, (128,128,512), 2, \"conv3_block1\",\n        keep_dict.get(\"conv3_block1_1_conv\"), keep_dict.get(\"conv3_block1_2_conv\"),\n        shortcut_conv=True)\n    for k in [2, 3, 4]:\n        x = bottleneck_block_pruned(x, (128,128,512), 1, f\"conv3_block{k}\",\n            keep_dict.get(f\"conv3_block{k}_1_conv\"),\n            keep_dict.get(f\"conv3_block{k}_2_conv\"))\n\n    return tf.keras.Model(inp, x, name=name)\n\n\ndef select_channels_l1(backbone, layers_to_prune, ratio):\n    \"\"\"Sélectionne les canaux à GARDER selon la L1-norm (ratio = fraction à supprimer).\"\"\"\n    keep_dict = {}\n    for layer in backbone.layers:\n        if isinstance(layer, tf.keras.layers.Conv2D) and layer.name in layers_to_prune:\n            W = layer.get_weights()[0]                          # (kh, kw, in, out)\n            l1 = np.sum(np.abs(W), axis=(0, 1, 2))              # L1 par filtre output\n            n_keep = int(round(W.shape[-1] * (1 - ratio)))\n            top_k = np.argsort(l1)[-n_keep:]\n            keep_dict[layer.name] = sorted(top_k.tolist())\n    return keep_dict\n\n\ndef transfer_weights_pruned(src_bb, dst_bb, keep_dict):\n    \"\"\"Transfère les poids src→dst en sliciant selon keep_dict.\n       - Conv2D : slice canaux output si pruné, slice canaux input si la conv précédente l'est\n       - BN     : slice les 4 stats si la conv qu'il suit est prunée\"\"\"\n    n_done = 0\n    for dst_l in dst_bb.layers:\n        try:\n            src_l = src_bb.get_layer(dst_l.name)\n        except (ValueError, KeyError):\n            continue\n        if not src_l.get_weights():\n            continue\n\n        sw = src_l.get_weights()\n        dw = [w.copy() for w in sw]\n\n        if isinstance(dst_l, tf.keras.layers.Conv2D):\n            k = sw[0]\n            # Output slicing (cette couche est-elle prunée ?)\n            if dst_l.name in keep_dict:\n                ko = keep_dict[dst_l.name]\n                k = k[:, :, :, ko]\n                if len(sw) > 1:\n                    dw[1] = sw[1][ko]\n            # Input slicing (la couche précédente était-elle prunée ?)\n            in_keep = None\n            if dst_l.name.endswith(\"_2_conv\"):\n                prev = dst_l.name.replace(\"_2_conv\", \"_1_conv\")\n                if prev in keep_dict:\n                    in_keep = keep_dict[prev]\n            elif dst_l.name.endswith(\"_3_conv\"):\n                prev = dst_l.name.replace(\"_3_conv\", \"_2_conv\")\n                if prev in keep_dict:\n                    in_keep = keep_dict[prev]\n            if in_keep is not None:\n                k = k[:, :, in_keep, :]\n            dw[0] = k\n\n        elif isinstance(dst_l, tf.keras.layers.BatchNormalization):\n            base = dst_l.name.replace(\"_bn\", \"_conv\")\n            if base in keep_dict:\n                ko = keep_dict[base]\n                dw = [w[ko] for w in sw]\n\n        try:\n            dst_l.set_weights(dw)\n            n_done += 1\n        except Exception as e:\n            print(f\"[WARN-bb] {dst_l.name}: shapes \"\n                  f\"src={[w.shape for w in sw]} → \"\n                  f\"dst={[w.shape for w in dst_l.get_weights()]} : {e}\")\n    return n_done\n\n\ndef transfer_se_and_head(src_model, dst_model):\n    \"\"\"Copie les poids des couches post-backbone (SE block + tête).\"\"\"\n    si = next(i for i, l in enumerate(src_model.layers) if isinstance(l, tf.keras.Model))\n    di = next(i for i, l in enumerate(dst_model.layers) if isinstance(l, tf.keras.Model))\n    n = 0\n    for s, d in zip(src_model.layers[si+1:], dst_model.layers[di+1:]):\n        if s.get_weights():\n            try:\n                d.set_weights(s.get_weights())\n                n += 1\n            except Exception as e:\n                print(f\"[WARN-head] {s.name}: {e}\")\n    return n\n\n\ndef build_pruned_amca_model(pruned_bb, head_cfg):\n    \"\"\"Assemble le modèle AMCA pruné = pruned_bb + SE + head.\"\"\"\n    inp = tf.keras.Input(shape=(IMG_SIZE, IMG_SIZE, 3), name=\"input\")\n    x   = pruned_bb(inp, training=False)\n    x   = se_block(x, ratio=16)\n    out = build_head(x, head_cfg)\n    return tf.keras.Model(inp, out, name=f\"AMCA_{ARCHITECTURE}_pruned\")\n","metadata":{"execution":{"iopub.execute_input":"2026-04-19T00:32:54.503093Z","iopub.status.busy":"2026-04-19T00:32:54.502503Z","iopub.status.idle":"2026-04-19T00:32:54.536526Z","shell.execute_reply":"2026-04-19T00:32:54.535852Z"},"papermill":{"duration":0.042804,"end_time":"2026-04-19T00:32:54.538050+00:00","exception":false,"start_time":"2026-04-19T00:32:54.495246+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"1610c25b","cell_type":"markdown","source":"## 10 — Reduced AMCA (data-driven)\n\n**FIX #4 & #5** :\n- Fenêtre de recherche élargie → plusieurs candidats sont réellement comparés.\n- Le probe utilise la **même tête** que le modèle final (comparaison équitable).\n","metadata":{"papermill":{"duration":0.006409,"end_time":"2026-04-19T00:32:54.550793+00:00","exception":false,"start_time":"2026-04-19T00:32:54.544384+00:00","status":"completed"},"tags":[]}},{"id":"acdff8b0","cell_type":"code","source":"class ReducedAMCA:\n    def __init__(self, arch, flops_budget, min_flops, min_auc=0.70):\n        self.arch         = arch\n        self.cfg          = BACKBONE_REGISTRY[arch]\n        self.flops_budget = flops_budget\n        self.min_flops    = min_flops\n        self.min_auc      = min_auc\n        self.strategy     = {}\n\n    def analyze(self, df):\n        n    = len(df)\n        imb  = abs(df[\"label\"].value_counts(normalize=True).get(0, 0)\n                 - df[\"label\"].value_counts(normalize=True).get(1, 0))\n        comp = \"simple\" if n > 5000 and imb < 0.1 else (\"moderate\" if n > 2000 else \"complex\")\n        print(f\"[AMCA-A] n={n}  imbal={imb:.3f}  complexity={comp}\")\n        return {\"n\": n, \"imbalance\": imb, \"complexity\": comp}\n\n    def _probe_head_cfg(self, info):\n        # Même tête que le modèle final → comparaison équitable entre stages\n        return {\"simple\":   {\"units\":[256],       \"dropouts\":[0.30]},\n                \"moderate\": {\"units\":[512, 256],  \"dropouts\":[0.40, 0.30]},\n                \"complex\":  {\"units\":[512, 256],  \"dropouts\":[0.40, 0.30]}}[info[\"complexity\"]]\n\n    def search(self, train_ds, val_ds, info, cw=None, epochs=5):\n        results  = []\n        head_cfg = self._probe_head_cfg(info)\n        print(f\"\\n[AMCA-B] Recherche cutting point — {self.arch}\")\n        print(f\"         min_flops={self.min_flops}G  budget={self.flops_budget}G\")\n        print(f\"         probe head: {head_cfg}\")\n\n        for cand in self.cfg[\"candidates\"]:\n            if cand[\"flops\"] > self.flops_budget:\n                print(f\"  [SKIP] {cand['name']} — dépasse budget\"); continue\n            if cand[\"flops\"] < self.min_flops:\n                print(f\"  [SKIP] {cand['name']} — sous-seuil\"); continue\n\n            base = self.cfg[\"builder\"](include_top=False, weights=\"imagenet\",\n                       input_shape=(self.cfg[\"input_size\"], self.cfg[\"input_size\"], 3))\n            try:\n                cut = base.get_layer(cand[\"layer\"]).output\n            except ValueError:\n                print(f\"  [SKIP] {cand['layer']} introuvable\"); continue\n\n            bb = tf.keras.Model(base.input, cut)\n            for l in bb.layers: l.trainable = False\n\n            inp = tf.keras.Input(shape=(self.cfg[\"input_size\"], self.cfg[\"input_size\"], 3))\n            x   = bb(inp, training=False)\n            x   = se_block(x, ratio=16)              # même SE que modèle final\n            out = build_head(x, head_cfg)             # même tête que modèle final\n            probe = tf.keras.Model(inp, out)\n            probe.compile(\n                optimizer=tf.keras.optimizers.AdamW(1e-3, weight_decay=1e-4),\n                loss=tf.keras.losses.BinaryCrossentropy(label_smoothing=0.05),\n                metrics=get_metrics(),\n            )\n            probe.fit(train_ds, validation_data=val_ds,\n                      epochs=epochs, class_weight=cw, verbose=0)\n\n            ev  = probe.evaluate(val_ds, verbose=0)\n            auc = ev[1]; acc = ev[2]\n            score = (auc**2) / cand[\"flops\"] if auc >= self.min_auc else 0.0\n            flag  = \"\" if auc >= self.min_auc else f\"  [AUC<{self.min_auc}]\"\n            print(f\"  {cand['name']:8s} | AUC={auc:.4f} | Acc={acc:.4f} \"\n                  f\"| {cand['flops']:.1f}G | Score={score:.5f}{flag}\")\n            results.append(dict(name=cand[\"name\"], layer=cand[\"layer\"],\n                                auc=auc, acc=acc, flops=cand[\"flops\"], score=score))\n            del probe, bb, base\n            tf.keras.backend.clear_session()\n\n        if not results:\n            raise RuntimeError(\"Aucun candidat valide.\")\n        best = max(results, key=lambda r: r[\"score\"])\n        if best[\"score\"] == 0:\n            best = max(results, key=lambda r: r[\"auc\"])\n            print(\"  [WARN] Fallback sur meilleur AUC absolu\")\n        print(f\"\\n  ★ Optimal : {best['name']} → {best['layer']}  \"\n              f\"(AUC={best['auc']:.4f}, {best['flops']}G)\")\n        self.strategy[\"cutting_point\"]   = best[\"layer\"]\n        self.strategy[\"cutting_flops\"]   = best[\"flops\"]\n        self.strategy[\"cutting_results\"] = results\n        return best[\"layer\"]\n\n    def design(self, info):\n        h = {\"simple\":  {\"units\":[256],      \"dropouts\":[0.30]},\n             \"moderate\":{\"units\":[512, 256], \"dropouts\":[0.40, 0.30]},\n             \"complex\": {\"units\":[512, 256], \"dropouts\":[0.40, 0.30]}}[info[\"complexity\"]]\n        f = {\"simple\":1.0, \"moderate\":0.60, \"complex\":0.40}[info[\"complexity\"]]\n        lr = {\"simple\":  {\"p1\":1e-3, \"p2\":5e-4, \"p3\":1e-5},\n              \"moderate\":{\"p1\":1e-3, \"p2\":1e-4, \"p3\":5e-6},\n              \"complex\": {\"p1\":5e-4, \"p2\":1e-4, \"p3\":5e-6}}[info[\"complexity\"]]\n        self.strategy.update({\"head\": h, \"freeze_p2\": f, \"lr\": lr})\n        print(f\"[AMCA-C/D] head={h['units']}  freeze_p2={f:.0%}  lr={lr}\")\n        return self.strategy\n\n    def run(self, df, train_ds, val_ds, cw=None):\n        print(\"\\n╔═══════════════════════════════╗\")\n        print(\"║  Reduced AMCA — démarrage     ║\")\n        print(\"╚═══════════════════════════════╝\")\n        info = self.analyze(df)\n        self.search(train_ds, val_ds, info, cw)\n        self.design(info)\n        print(f\"\\n✓ cutting_point = {self.strategy['cutting_point']}\")\n        return self.strategy\n","metadata":{"execution":{"iopub.execute_input":"2026-04-19T00:32:54.564920Z","iopub.status.busy":"2026-04-19T00:32:54.564701Z","iopub.status.idle":"2026-04-19T00:32:54.580346Z","shell.execute_reply":"2026-04-19T00:32:54.579703Z"},"papermill":{"duration":0.024427,"end_time":"2026-04-19T00:32:54.581693+00:00","exception":false,"start_time":"2026-04-19T00:32:54.557266+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"e457ba28","cell_type":"markdown","source":"## 11 — Baseline : ResNet50 complet (avant cutting)","metadata":{"papermill":{"duration":0.006277,"end_time":"2026-04-19T00:32:54.594108+00:00","exception":false,"start_time":"2026-04-19T00:32:54.587831+00:00","status":"completed"},"tags":[]}},{"id":"abcc5c9a","cell_type":"code","source":"def train_baseline(strategy, train_ds, val_ds, class_weight,\n                   epochs_p1=15, epochs_p2=10):\n    \"\"\"\n    Baseline ResNet50 complet.\n    3 optimisations vs version précédente :\n      1. Phase 1 : top 50% du backbone dégelé dès le début (BN gelé)\n         → adaptation aux features rétiniennes dès l'epoch 1\n      2. Epochs : 15 + 10 = 25 (vs 10+6=16 avant)\n         → ×1.7 plus de mises à jour des poids\n      3. ReduceLROnPlateau factor=0.2 (÷5 vs ÷3 avant)\n         → descente plus fine vers le minimum de loss\n    Ces 3 changements reproduisent la stratégie du notebook mr-bouka\n    qui obtient 97.54% sur APTOS 2019 avec exactement ResNet50.\n    \"\"\"\n    print(\"\\n\" + \"=\"*62)\n    print(\"  BASELINE — ResNet50 complet (sans cutting AMCA)\")\n    print(\"=\"*62)\n\n    cfg = BACKBONE_REGISTRY[\"resnet50\"]\n\n    with strategy.scope():\n        base = cfg[\"builder\"](\n            include_top=False, weights=\"imagenet\",\n            input_shape=(IMG_SIZE, IMG_SIZE, 3)\n        )\n\n        # ── FIX 1 : 50% backbone dégelé dès Phase 1 (BN toujours gelé) ──\n        total = len(base.layers)\n        start_p1 = int(total * 0.50)   # 50% dégelé (vs 100% gelé avant)\n        for k, layer in enumerate(base.layers):\n            if isinstance(layer, tf.keras.layers.BatchNormalization):\n                layer.trainable = False  # BN toujours gelé — stabilité garantie\n            else:\n                layer.trainable = (k >= start_p1)\n\n        n_train = sum(1 for l in base.layers if l.trainable)\n        n_bn    = sum(1 for l in base.layers\n                      if isinstance(l, tf.keras.layers.BatchNormalization))\n        print(f\"[BL] Backbone : {n_train}/{total} couches entraînables \"\n              f\"| {n_bn} BN gelés\")\n\n        inp = tf.keras.Input(shape=(IMG_SIZE, IMG_SIZE, 3), name=\"input_bl\")\n        x   = base(inp, training=False)\n        x   = se_block(x, ratio=16)\n        x   = tf.keras.layers.GlobalAveragePooling2D()(x)\n        for units, drop in [(512, 0.4), (256, 0.3)]:\n            x = tf.keras.layers.Dense(units, activation=None,\n                    kernel_regularizer=tf.keras.regularizers.l2(1e-4))(x)\n            x = tf.keras.layers.BatchNormalization()(x)\n            x = tf.keras.layers.Activation(\"relu\")(x)\n            x = tf.keras.layers.Dropout(drop)(x)\n        out = tf.keras.layers.Dense(1, activation=\"sigmoid\", dtype=\"float32\")(x)\n        bl_model = tf.keras.Model(inp, out, name=\"Baseline_ResNet50\")\n\n        bl_model.compile(\n            optimizer=tf.keras.optimizers.AdamW(1e-3, weight_decay=1e-4),\n            loss=tf.keras.losses.BinaryCrossentropy(label_smoothing=0.01),  # ✅ CORRECTION #5\n            metrics=get_metrics()\n        )\n\n    print(f\"Paramètres baseline : {bl_model.count_params():,}\")\n    print(f\"FLOPs baseline      : ~5.9 GFLOPs (Stage4 complet)\")\n\n    # ── FIX 3 : callbacks avec factor=0.2 (÷5 au lieu de ÷3) ────────\n    def get_bl_callbacks(name, patience=6):\n        return [\n            tf.keras.callbacks.ModelCheckpoint(\n                f\"{name}.keras\", monitor=\"val_auc\", mode=\"max\",\n                save_best_only=True, verbose=1\n            ),\n            tf.keras.callbacks.EarlyStopping(\n                monitor=\"val_auc\", mode=\"max\",\n                patience=patience, restore_best_weights=True, verbose=1\n            ),\n            tf.keras.callbacks.ReduceLROnPlateau(\n                monitor=\"val_loss\", factor=0.2,   # FIX : 0.2 au lieu de 0.3\n                patience=3, min_lr=1e-7, verbose=1\n            ),\n        ]\n\n    # ── Phase 1 : top 50% dégelé (BN gelé), LR=1e-3 ─────────────────\n    print(f\"\\nPhase 1 — {epochs_p1} epochs  LR=1e-3  \"\n          f\"top 50% dégelé (BN gelé)\")\n    h_bl1 = bl_model.fit(\n        train_ds, validation_data=val_ds,\n        epochs=epochs_p1, class_weight=class_weight,\n        callbacks=get_bl_callbacks(\"baseline_p1\", patience=6),\n        verbose=1\n    )\n\n    # ── Phase 2 : top 70% dégelé (BN gelé), LR=1e-4 ─────────────────\n    try:\n        bl_model.load_weights(\"baseline_p1.keras\")\n        print(\"[BL] Meilleurs poids Phase 1 chargés\")\n    except Exception:\n        pass\n\n    start_p2 = int(total * 0.30)   # 70% dégelé\n    for k, layer in enumerate(base.layers):\n        if isinstance(layer, tf.keras.layers.BatchNormalization):\n            layer.trainable = False\n        else:\n            layer.trainable = (k >= start_p2)\n\n    # ✅ v8 FIX #1 : set_lr au lieu de compile() → optimizer conservé\n    set_lr(bl_model, 1e-4)\n    print(f\"\\nPhase 2 — {epochs_p2} epochs  LR=1e-4  top 70% dégelé (BN gelé)\")\n    h_bl2 = bl_model.fit(\n        train_ds, validation_data=val_ds,\n        epochs=epochs_p2, class_weight=class_weight,\n        callbacks=get_bl_callbacks(\"baseline_p2\", patience=5),\n        verbose=1\n    )\n\n    # ── Phase 3 : full fine-tuning progressif (BN gelé) ──────────────\n    try:\n        bl_model.load_weights(\"baseline_p2.keras\")\n        print(\"[BL] Meilleurs poids Phase 2 chargés\")\n    except Exception:\n        pass\n\n    # 3a — warmup 100% dégelé, LR très faible\n    for layer in base.layers:\n        if isinstance(layer, tf.keras.layers.BatchNormalization):\n            layer.trainable = False\n        else:\n            layer.trainable = True\n    # ✅ v8 FIX #1 : set_lr au lieu de compile() → optimizer conservé\n    set_lr(bl_model, 5e-6)\n    print(\"\\nPhase 3 — 6 epochs  LR=5e-6  100% dégelé (BN gelé)\")\n    h_bl3 = bl_model.fit(\n        train_ds, validation_data=val_ds,\n        epochs=6, class_weight=class_weight,\n        callbacks=get_bl_callbacks(\"baseline_p3\", patience=5),\n        verbose=1\n    )\n\n    # ── Évaluation finale TTA + seuil Youden ─────────────────────────\n    try:\n        bl_model.load_weights(\"baseline_p3.keras\")\n    except Exception:\n        try:\n            bl_model.load_weights(\"baseline_p2.keras\")\n        except Exception:\n            pass\n\n    # TTA : 3 transformations déterministes\n    preds_bl = np.mean([\n        bl_model.predict(val_ds, verbose=0).ravel()\n        for _ in range(3)\n    ], axis=0)\n    y_true = val_df[\"label\"].values\n    auc_bl = roc_auc_score(y_true, preds_bl)\n    acc_bl = accuracy_score(y_true, (preds_bl >= 0.5).astype(int))\n\n    # Seuil Youden\n    fpr_bl, tpr_bl, thr_bl = roc_curve(y_true, preds_bl)\n    j_bl        = tpr_bl - fpr_bl\n    best_thr_bl = float(thr_bl[np.argmax(j_bl)])\n    acc_youden_bl = accuracy_score(\n        y_true, (preds_bl >= best_thr_bl).astype(int)\n    )\n\n    print(\"\\n\" + \"=\"*62)\n    print(\"  RÉSULTAT BASELINE (ResNet50 complet)\")\n    print(\"=\"*62)\n    print(f\"  AUC             : {auc_bl:.4f}\")\n    print(f\"  Accuracy (0.50) : {acc_bl:.4f}  ({acc_bl*100:.2f}%)\")\n    print(f\"  Accuracy Youden : {acc_youden_bl:.4f}  \"\n          f\"({acc_youden_bl*100:.2f}%)  seuil={best_thr_bl:.4f}\")\n    print(f\"  Params          : {bl_model.count_params():,}\")\n    print(f\"  FLOPs           : ~5.9 GFLOPs\")\n    print(\"=\"*62)\n\n    flops_bl = measure_flops(bl_model)\n    print(f\"  FLOPs mesurés   : {flops_bl/1e9:.3f} GFLOPs\")\n\n    return bl_model, {\n        \"auc\"        : auc_bl,\n        \"acc\"        : acc_bl,\n        \"acc_youden\" : acc_youden_bl,\n        \"threshold\"  : best_thr_bl,\n        \"params\"     : bl_model.count_params(),\n        \"flops_g\"    : flops_bl / 1e9,\n    }\n\n\nbl_model, baseline_results = train_baseline(\n    strategy, train_ds, val_ds, class_weight,\n    epochs_p1=15, epochs_p2=10\n)\n","metadata":{"execution":{"iopub.execute_input":"2026-04-19T00:32:54.607876Z","iopub.status.busy":"2026-04-19T00:32:54.607658Z","iopub.status.idle":"2026-04-19T00:39:02.947300Z","shell.execute_reply":"2026-04-19T00:39:02.946219Z"},"papermill":{"duration":368.348247,"end_time":"2026-04-19T00:39:02.948652+00:00","exception":true,"start_time":"2026-04-19T00:32:54.600405+00:00","status":"failed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"8a8d659e","cell_type":"markdown","source":"## 12 — Lancer l'AMCA","metadata":{"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[]}},{"id":"c327c253","cell_type":"code","source":"amca = ReducedAMCA(\n    arch=ARCHITECTURE,\n    flops_budget=FLOPS_BUDGET,\n    min_flops=MIN_FLOPS,\n    min_auc=0.70\n)\namca_strategy = amca.run(train_df, train_ds, val_ds, class_weight)\n\nprint(\"\\nStratégie finale :\")\nfor k, v in amca_strategy.items():\n    if k != \"cutting_results\":\n        print(f\"  {k:20s}: {v}\")\n","metadata":{"execution":{"iopub.execute_input":"2026-04-17T20:38:48.348335Z","iopub.status.busy":"2026-04-17T20:38:48.348062Z","iopub.status.idle":"2026-04-17T20:45:56.832858Z","shell.execute_reply":"2026-04-17T20:45:56.832117Z"},"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[]},"outputs":[],"execution_count":null},{"id":"bd774fd2","cell_type":"markdown","source":"## 13 — Construction du Modèle Final","metadata":{"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[]}},{"id":"d29f1203","cell_type":"code","source":"def freeze(bb, ratio):\n    n = len(bb.layers)\n    for i, l in enumerate(bb.layers):\n        l.trainable = (i >= int(n * ratio))\n    f = sum(1 for l in bb.layers if not l.trainable)\n    print(f\"[MODEL] {f}/{n} layers gelées ({ratio:.0%})\")\n\ndef unfreeze_except_bn(backbone, top_fraction):\n    \"\"\"Dégèle top_fraction des couches SAUF les BatchNormalization.\"\"\"\n    layers = backbone.layers\n    n_unfreeze = int(len(layers) * top_fraction)\n    for layer in layers[:-n_unfreeze] if n_unfreeze > 0 else layers:\n        layer.trainable = False\n    if n_unfreeze > 0:\n        for layer in layers[-n_unfreeze:]:\n            if isinstance(layer, tf.keras.layers.BatchNormalization):\n                layer.trainable = False\n            else:\n                layer.trainable = True\n    n_frozen_bn = sum(1 for l in backbone.layers\n                      if isinstance(l, tf.keras.layers.BatchNormalization))\n    n_trainable = sum(1 for l in backbone.layers if l.trainable)\n    print(f\"[MODEL] {n_trainable}/{len(layers)} trainable | {n_frozen_bn} BN gelés\")\n\ndef build_model(arch, strat):\n    cfg   = BACKBONE_REGISTRY[arch]\n    shape = (cfg[\"input_size\"], cfg[\"input_size\"], 3)\n    base  = cfg[\"builder\"](include_top=False, weights=\"imagenet\", input_shape=shape)\n    try:\n        cut = base.get_layer(strat[\"cutting_point\"]).output\n    except (ValueError, KeyError):\n        cut = base.output\n    bb = tf.keras.Model(base.input, cut, name=f\"{arch}_bb\")\n    freeze(bb, 1.0)\n\n    inp = tf.keras.Input(shape=shape, name=\"input\")\n    x   = bb(inp, training=False)\n    x   = se_block(x, ratio=16)\n    out = build_head(x, strat[\"head\"])\n    model = tf.keras.Model(inp, out, name=f\"AMCA_{arch}\")\n    print(f\"[MODEL] {model.name} — {model.count_params():,} params\")\n    return model, bb\n\nwith strategy.scope():\n    model, backbone = build_model(ARCHITECTURE, amca_strategy)\n    model = compile_model(model, amca_strategy[\"lr\"][\"p1\"])\nmodel.summary(show_trainable=True)\n","metadata":{"execution":{"iopub.execute_input":"2026-04-17T20:45:57.390854Z","iopub.status.busy":"2026-04-17T20:45:57.390579Z","iopub.status.idle":"2026-04-17T20:45:59.987282Z","shell.execute_reply":"2026-04-17T20:45:59.986666Z"},"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[]},"outputs":[],"execution_count":null},{"id":"d1e5554d","cell_type":"markdown","source":"## 14 — Entraînement Phase 1 (tête seule)","metadata":{"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[]}},{"id":"5df0ec9b","cell_type":"code","source":"lr = amca_strategy[\"lr\"]\nprint(f\"Phase 1 — {EPOCHS_PHASE1} epochs  LR={lr['p1']}  backbone 100% gelé\")\n\nh1 = model.fit(train_ds, validation_data=val_ds,\n               epochs=EPOCHS_PHASE1, class_weight=class_weight,\n               callbacks=get_callbacks(\"amca_p1\"), verbose=1)\n\n# ✅ v8 FIX #2 : Recalibrage BN entre Phase 1 et Phase 2\n# Les stats BN ImageNet → recalibration sur données rétine RGB (2 epochs, LR=1e-6)\nh_bn = recalibrate_bn(model, backbone, train_ds, val_ds, class_weight,\n                      lr_bn=1e-6, epochs=2, ckpt_name=\"amca_bn_calib\")\ntry:\n    model.load_weights(\"amca_p1.keras\")\n    print(\"[CKPT] Meilleurs poids Phase 1 rechargés après BN calib\")\nexcept Exception:\n    pass\n","metadata":{"execution":{"iopub.execute_input":"2026-04-17T20:46:00.479667Z","iopub.status.busy":"2026-04-17T20:46:00.478853Z","iopub.status.idle":"2026-04-17T20:49:13.240481Z","shell.execute_reply":"2026-04-17T20:49:13.239618Z"},"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[]},"outputs":[],"execution_count":null},{"id":"a8535551","cell_type":"markdown","source":"## 15 — Phase 2 (dégel partiel, BN gelé)","metadata":{"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[]}},{"id":"cd40e704","cell_type":"code","source":"unfreeze_except_bn(backbone, 1.0 - amca_strategy[\"freeze_p2\"])\n# ✅ v8 FIX #1 : set_lr au lieu de compile_model → optimizer AdamW conservé\nset_lr(model, lr[\"p2\"])\n\nprint(f\"Phase 2 — {EPOCHS_PHASE2} epochs  LR={lr['p2']}\")\nh2 = model.fit(train_ds, validation_data=val_ds,\n               epochs=EPOCHS_PHASE2, class_weight=class_weight,\n               callbacks=get_callbacks(\"amca_p2\"), verbose=1)\n","metadata":{"execution":{"iopub.execute_input":"2026-04-17T20:49:13.919929Z","iopub.status.busy":"2026-04-17T20:49:13.919250Z","iopub.status.idle":"2026-04-17T20:52:17.258011Z","shell.execute_reply":"2026-04-17T20:52:17.256928Z"},"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[]},"outputs":[],"execution_count":null},{"id":"c10c2c3e","cell_type":"markdown","source":"## 16 — Phase 3 (full fine-tuning progressif)\n\n**FIX #9** : les epochs Phase 3 sont maintenant répartis pour totaliser exactement `EPOCHS_PHASE3`.\n","metadata":{"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[]}},{"id":"ad7dc7a2","cell_type":"code","source":"# ── Chargement meilleurs poids Phase 2 ─────────────────────\ntry:\n    model.load_weights(\"amca_p2.keras\")\n    print(\"[CKPT] Meilleurs poids Phase 2 chargés\")\nexcept Exception as e:\n    print(f\"[WARN] {e}\")\n\n# Répartition des epochs Phase 3 : 3a=20%, 3b=50%, 3c=30%\nn_p3a = max(1, int(round(EPOCHS_PHASE3 * 0.25)))\nn_p3b = max(1, int(round(EPOCHS_PHASE3 * 0.50)))\nn_p3c = max(1, EPOCHS_PHASE3 - n_p3a - n_p3b)\nprint(f\"[P3] Répartition : p3a={n_p3a}  p3b={n_p3b}  p3c={n_p3c}  \"\n      f\"(total={n_p3a+n_p3b+n_p3c}/{EPOCHS_PHASE3})\")\n\n# ── Phase 3a : warmup 40% dégelé (BN gelé) ─\nunfreeze_except_bn(backbone, 0.40)\n# ✅ v8 FIX #1 : set_lr — pas de reset optimizer\nset_lr(model, lr[\"p3\"] * 0.1)\n\nprint(f\"\\nPhase 3a — {n_p3a} epochs  LR={lr['p3']*0.1:.1e}  warmup BN gelé\")\nh3a = model.fit(train_ds, validation_data=val_ds,\n                epochs=n_p3a, class_weight=class_weight,\n                callbacks=get_callbacks(\"amca_p3a\", patience=3), verbose=1)\n\n# ── Phase 3b : 70% dégelé (BN gelé) ────────\nunfreeze_except_bn(backbone, 0.70)\n# ✅ v8 FIX #1 : set_lr\nset_lr(model, lr[\"p3\"])\n\nprint(f\"\\nPhase 3b — {n_p3b} epochs  LR={lr['p3']:.1e}  70% BN gelé\")\nh3b = model.fit(train_ds, validation_data=val_ds,\n                epochs=n_p3b, class_weight=class_weight,\n                callbacks=get_callbacks(\"amca_p3b\", patience=4), verbose=1)\n\n# ── Phase 3c : 100% dégelé (BN gelé) ───────\nunfreeze_except_bn(backbone, 1.00)\n# ✅ v8 FIX #1 : set_lr\nset_lr(model, lr[\"p3\"] * 0.5)\n\nprint(f\"\\nPhase 3c — {n_p3c} epochs  LR={lr['p3']*0.5:.1e}  100% BN gelé\")\nh3c = model.fit(train_ds, validation_data=val_ds,\n                epochs=n_p3c, class_weight=class_weight,\n                callbacks=get_callbacks(\"amca_p3\", patience=5), verbose=1)\n\nclass _MergedHistory:\n    def __init__(self, *hs):\n        self.history = {}\n        for h in hs:\n            for k, v in h.history.items():\n                self.history.setdefault(k, []).extend(v)\n        # Aligner les longueurs (clés absentes → padding avec NaN)\n        if self.history:\n            max_len = max(len(v) for v in self.history.values())\n            for k, v in self.history.items():\n                if len(v) < max_len:\n                    self.history[k] = v + [float(\"nan\")] * (max_len - len(v))\n\nh3 = _MergedHistory(h3a, h3b, h3c)\n","metadata":{"execution":{"iopub.execute_input":"2026-04-17T20:52:18.067941Z","iopub.status.busy":"2026-04-17T20:52:18.067580Z","iopub.status.idle":"2026-04-17T20:55:33.943941Z","shell.execute_reply":"2026-04-17T20:55:33.943284Z"},"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[]},"outputs":[],"execution_count":null},{"id":"ee4ae7e3","cell_type":"markdown","source":"## 17 — Courbes d'Apprentissage","metadata":{"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[]}},{"id":"7fa0dbf1","cell_type":"code","source":"def plot_curves(histories):\n    acc, vacc, loss, vloss, bounds, c = [], [], [], [], [], 0\n    for h in histories:\n        acc.extend(h.history.get(\"accuracy\", []))\n        vacc.extend(h.history.get(\"val_accuracy\", []))\n        loss.extend(h.history[\"loss\"])\n        vloss.extend(h.history[\"val_loss\"])\n        c += len(h.history[\"loss\"]); bounds.append(c)\n    bounds = bounds[:-1]\n    ep = range(1, len(acc) + 1)\n    fig, axes = plt.subplots(1, 2, figsize=(14, 5))\n    for ax, (m, vm, yl) in zip(axes, [(acc, vacc, \"Accuracy\"), (loss, vloss, \"Loss\")]):\n        ax.plot(ep, m,  \"b-\", label=\"Train\")\n        ax.plot(ep, vm, \"r-\", label=\"Val\")\n        for b in bounds:\n            ax.axvline(b + 0.5, color=\"gray\", linestyle=\"--\", lw=1)\n        ax.set_xlabel(\"Epochs\"); ax.set_ylabel(yl)\n        ax.legend(); ax.grid(True, alpha=0.3)\n    plt.tight_layout(); plt.savefig(\"learning_curves.png\", dpi=150); plt.show()\n\nplot_curves([h1, h2, h3])\n","metadata":{"execution":{"iopub.execute_input":"2026-04-17T20:55:34.882954Z","iopub.status.busy":"2026-04-17T20:55:34.882628Z","iopub.status.idle":"2026-04-17T20:55:35.511045Z","shell.execute_reply":"2026-04-17T20:55:35.510450Z"},"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[]},"outputs":[],"execution_count":null},{"id":"71862da4","cell_type":"markdown","source":"## 18 — Évaluation Avant Pruning (TTA réel)\n\n**FIX #2** : le TTA applique maintenant de vraies transformations (identité + flips horizontaux/verticaux) et moyenne les probabilités. Plus de passes identiques.\n","metadata":{"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[]}},{"id":"b0785b6a","cell_type":"code","source":"def tta_predict(model, ds, tta_transforms=None):\n    \"\"\"TTA réel : applique plusieurs transformations déterministes et moyenne les probas.\"\"\"\n    if tta_transforms is None:\n        tta_transforms = [\n            lambda x: x,                                        # identité\n            lambda x: tf.image.flip_left_right(x),              # flip horizontal\n            lambda x: tf.image.flip_up_down(tf.image.flip_left_right(x)),  # rot 180°\n        ]\n    preds_list = []\n    for tf_fn in tta_transforms:\n        batch_preds = []\n        for imgs, _ in ds:\n            batch_preds.append(model.predict(tf_fn(imgs), verbose=0))\n        preds_list.append(np.concatenate(batch_preds, axis=0).ravel())\n    return np.mean(preds_list, axis=0)\n\ndef evaluate(model, ds, df, use_tta=True, label=\"\"):\n    print(f\"\\n[EVAL] {label}\")\n    if use_tta:\n        preds = tta_predict(model, ds)\n        print(f\"  TTA : {3} transformations moyennées\")\n    else:\n        preds = model.predict(ds, verbose=0).ravel()\n    y = df[\"label\"].values\n\n    # Seuil optimal par critère de Youden (J = tpr - fpr)\n    fpr_arr, tpr_arr, thresholds = roc_curve(y, preds)\n    j_scores  = tpr_arr - fpr_arr\n    best_idx  = np.argmax(j_scores)\n    best_thr  = float(thresholds[best_idx])\n\n    for thr, name in [(0.50, \"seuil fixe 0.50\"), (best_thr, f\"seuil Youden {best_thr:.3f}\")]:\n        yp = (preds >= thr).astype(int)\n        print(f\"\\n  [{name}]\")\n        print(classification_report(y, yp, digits=4, zero_division=0))\n        print(\"  Confusion Matrix:\"); print(confusion_matrix(y, yp))\n\n    auc = roc_auc_score(y, preds)\n    print(f\"\\nROC AUC = {auc:.4f}  |  Seuil optimal = {best_thr:.4f}\")\n    print(f\"Sensitivity (recall DR) @ optimal : {tpr_arr[best_idx]:.4f}\")\n    print(f\"Specificity (1-fpr)     @ optimal : {1-fpr_arr[best_idx]:.4f}\")\n\n    yp_opt = (preds >= best_thr).astype(int)\n    acc_youden = float(np.mean(yp_opt == y))\n    return {\"auc\": auc, \"preds\": preds, \"yp\": yp_opt,\n            \"threshold\": best_thr, \"acc_youden\": acc_youden, \"acc\": acc_youden}\n\nr_before = evaluate(model, val_ds, val_df, use_tta=True, label=\"AVANT PRUNING\")\n","metadata":{"execution":{"iopub.execute_input":"2026-04-17T20:55:36.579499Z","iopub.status.busy":"2026-04-17T20:55:36.579139Z","iopub.status.idle":"2026-04-17T20:56:10.545921Z","shell.execute_reply":"2026-04-17T20:56:10.545035Z"},"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[]},"outputs":[],"execution_count":null},{"id":"e5130d88","cell_type":"markdown","source":"## 19 — Pruning structurel (vraie suppression de filtres)\n\n**v6** : remplace l'ancien zero-masking par un VRAI pruning structurel :\n1. Sélection L1-norm des canaux à conserver dans les `*_1_conv` et `*_2_conv`\n2. Reconstruction physique du backbone avec largeurs réduites\n3. Transfert des poids avec slicing (output et input des couches conv, plus stats BN)\n4. Mesure des **vrais FLOPs** via `tf.profiler` (pas la table statique)\n\nLes couches `*_3_conv` (sortie de bloc), `*_0_conv` (shortcut) et `conv1_conv` (stem) ne sont PAS prunées : leur dimension est imposée par les connexions résiduelles.\n","metadata":{"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[]}},{"id":"027ff29b","cell_type":"code","source":"# ── Charger les meilleurs poids AMCA (Phase 3) sur le modèle ORIGINAL ─\ntry:\n    model.load_weights(\"amca_p3.keras\")\n    print(\"[CKPT] Meilleurs poids AMCA (Phase 3) chargés sur l'architecture originale\")\nexcept Exception as e:\n    try:\n        model.load_weights(\"amca_p2.keras\")\n        print(f\"[CKPT] Phase 3 indisponible ({e}) — fallback Phase 2\")\n    except Exception as e2:\n        print(f\"[WARN] Aucun checkpoint chargé : {e2}\")\n\n# ── Mesure FLOPs et AUC du modèle AMCA AVANT pruning ─\nflops_amca = measure_flops(model)\nprint(f\"\\n[FLOPS] AMCA Stage2 original : {flops_amca/1e9:.3f} GFLOPs  \"\n      f\"({model.count_params():,} params)\")\n\nauc_before = model.evaluate(val_ds, verbose=0)[1]\nprint(f\"[AUC ]  AMCA Stage2 original : {auc_before:.4f}\")\n\n# ── 1. Identifier les couches prunables (couches conv internes des blocs bottleneck) ─\nprunable_layers = [\n    l.name for l in backbone.layers\n    if isinstance(l, tf.keras.layers.Conv2D)\n    and (l.name.endswith(\"_1_conv\") or l.name.endswith(\"_2_conv\"))\n    and not l.name.endswith(\"_0_conv\")\n    and l.name != \"conv1_conv\"\n]\nprint(f\"\\n[PRUNING] {len(prunable_layers)} couches prunables identifiées\")\nprint(f\"          (ratio = {PRUNE_RATIO:.0%} des canaux supprimés)\")\n\n# ── 2. Sélection L1-norm des canaux à conserver ─\nkeep_dict = select_channels_l1(backbone, prunable_layers, PRUNE_RATIO)\n\nn_orig_filters = sum(backbone.get_layer(n).get_weights()[0].shape[-1]\n                     for n in prunable_layers)\nn_kept_filters = sum(len(v) for v in keep_dict.values())\nn_removed = n_orig_filters - n_kept_filters\nprint(f\"          Filtres originaux : {n_orig_filters}\")\nprint(f\"          Filtres conservés : {n_kept_filters}\")\nprint(f\"          Filtres SUPPRIMÉS : {n_removed}\")\n\n# ── 3. Construire le backbone pruné et transférer les poids ─\nprint(\"\\n[PRUNING] Reconstruction du backbone avec largeurs réduites...\")\nwith strategy.scope():\n    pruned_backbone = build_pruned_resnet50_stage2(keep_dict)\n    n_bb_done = transfer_weights_pruned(backbone, pruned_backbone, keep_dict)\n    print(f\"          Poids backbone transférés : {n_bb_done} couches\")\n\n    # ── 4. Construire le modèle AMCA complet pruné ─\n    pruned_model = build_pruned_amca_model(pruned_backbone, amca_strategy[\"head\"])\n    n_head_done = transfer_se_and_head(model, pruned_model)\n    print(f\"          Poids SE+head transférés  : {n_head_done} couches\")\n\n    pruned_model = compile_model(pruned_model, 1e-5)\n\n# ── 5. Mesure FLOPs après pruning structurel ─\nflops_pruned = measure_flops(pruned_model)\nprint(f\"\\n[FLOPS] AMCA Stage2 + pruning : {flops_pruned/1e9:.3f} GFLOPs  \"\n      f\"({pruned_model.count_params():,} params)\")\nprint(f\"        Réduction vs AMCA orig : {(1 - flops_pruned/flops_amca)*100:.1f}% FLOPs, \"\n      f\"{(1 - pruned_model.count_params()/model.count_params())*100:.1f}% params\")\n\n# ── 6. AUC immédiate (avant fine-tune recovery) — chute attendue ─\nauc_post_prune = pruned_model.evaluate(val_ds, verbose=0)[1]\nprint(f\"\\n[AUC ] AMCA Stage2 + pruning (sans recovery) : {auc_post_prune:.4f}  \"\n      f\"(chute = {auc_before - auc_post_prune:+.4f})\")\n\n# Remplacer le modèle courant par le modèle pruné (préserve la pipeline en aval)\nmodel    = pruned_model\nbackbone = pruned_backbone\n","metadata":{"execution":{"iopub.execute_input":"2026-04-17T20:56:11.545122Z","iopub.status.busy":"2026-04-17T20:56:11.544860Z","iopub.status.idle":"2026-04-17T20:56:24.958158Z","shell.execute_reply":"2026-04-17T20:56:24.957098Z"},"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[]},"outputs":[],"execution_count":null},{"id":"6e55defd","cell_type":"markdown","source":"## 20 — Fine-tune Recovery Post-Pruning\n\nLe modèle pruné a perdu de l'AUC à cause de la suppression de filtres. Le recovery réentraîne avec un LR faible pour récupérer la performance. **BN gelé** pour éviter l'explosion val_loss classique.\n","metadata":{"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[]}},{"id":"8f645dde","cell_type":"code","source":"# Pas de masques à maintenir — le pruning est structurel (filtres physiquement supprimés)\n# BN gelé pendant le recovery pour stabilité\nfor layer in backbone.layers:\n    if isinstance(layer, tf.keras.layers.BatchNormalization):\n        layer.trainable = False\n    else:\n        layer.trainable = True\n\n# ✅ v8 FIX #1 : set_lr — conserve optimizer\nset_lr(model, 1e-5)\n\nn_bb = len(backbone.layers)\nn_train = sum(1 for l in backbone.layers if l.trainable)\nprint(f\"[MODEL] backbone pruné ({n_bb} layers, {n_train} entraînables — BN gelé)\")\nprint(f\"Fine-tune recovery — {EPOCHS_FINETUNE} epochs  LR=1e-5\")\n\nh_ft = model.fit(train_ds, validation_data=val_ds,\n                 epochs=EPOCHS_FINETUNE, class_weight=class_weight,\n                 callbacks=get_callbacks(\"amca_ft_pruned\", patience=8),\n                 verbose=1)\n","metadata":{"execution":{"iopub.execute_input":"2026-04-17T20:56:26.024078Z","iopub.status.busy":"2026-04-17T20:56:26.023251Z","iopub.status.idle":"2026-04-17T21:01:05.054412Z","shell.execute_reply":"2026-04-17T21:01:05.053251Z"},"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[]},"outputs":[],"execution_count":null},{"id":"94ea7d79","cell_type":"markdown","source":"## 21 — Évaluation Après Pruning","metadata":{"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[]}},{"id":"8eeaa78e","cell_type":"code","source":"r_after = evaluate(model, val_ds, val_df, use_tta=True,\n                   label=\"APRÈS PRUNING STRUCTUREL + FINE-TUNE\")\n\nflops_final = measure_flops(model)\nprint(f\"\\n[FLOPS] Modèle final : {flops_final/1e9:.3f} GFLOPs  ({model.count_params():,} params)\")\n\n# ── Accuracy Youden pour chaque modèle ────────────────────\nacc_bl        = baseline_results.get(\"acc_youden\", baseline_results.get(\"acc\", 0.0))\nacc_before    = r_before.get(\"acc_youden\", r_before.get(\"acc\", 0.0))\nacc_after     = r_after.get(\"acc_youden\",  r_after.get(\"acc\", 0.0))\nthr_bl        = baseline_results.get(\"threshold\", 0.50)\nthr_before    = r_before.get(\"threshold\", 0.50)\nthr_after     = r_after.get(\"threshold\",  0.50)\n\nprint(\"\\n\" + \"=\"*84)\nprint(\"  COMPARAISON FINALE — VRAIS FLOPs (mesurés par tf.profiler)\")\nprint(\"=\"*84)\nprint(f\"{'Modèle':<35} {'AUC':>7} {'Accuracy':>10} {'Seuil':>7} {'Params':>11} {'FLOPs':>9}\")\nprint(\"-\"*84)\nprint(f\"{'Baseline ResNet50 complet':<35} \"\n      f\"{baseline_results['auc']:>7.4f} \"\n      f\"{acc_bl:>9.2%} \"\n      f\"{thr_bl:>7.4f} \"\n      f\"{baseline_results['params']:>11,} \"\n      f\"{baseline_results['flops_g']:>7.3f}G\")\nprint(f\"{'AMCA Stage2 (cutting seul)':<35} \"\n      f\"{r_before['auc']:>7.4f} \"\n      f\"{acc_before:>9.2%} \"\n      f\"{thr_before:>7.4f} \"\n      f\"{'—':>11} \"\n      f\"{flops_amca/1e9:>7.3f}G\")\nprint(f\"{'AMCA Stage2 + pruning STRUCTUREL':<35} \"\n      f\"{r_after['auc']:>7.4f} \"\n      f\"{acc_after:>9.2%} \"\n      f\"{thr_after:>7.4f} \"\n      f\"{model.count_params():>11,} \"\n      f\"{flops_final/1e9:>7.3f}G\")\nprint(\"=\"*84)\n\ndelta_auc        = r_after[\"auc\"]  - baseline_results[\"auc\"]\ndelta_acc        = acc_after       - acc_bl\nred_flops_total  = (1 - flops_final/(baseline_results[\"flops_g\"]*1e9)) * 100\nred_flops_cut    = (1 - flops_amca /(baseline_results[\"flops_g\"]*1e9)) * 100\nred_flops_prun   = (1 - flops_final/flops_amca) * 100\nred_params_total = (1 - model.count_params()/baseline_results[\"params\"]) * 100\n\nprint(f\"\\n  ΔAUC      (final vs baseline)        : {delta_auc:+.4f}\")\nprint(f\"  ΔAccuracy (final vs baseline)        : {delta_acc:+.2%}\")\nprint(f\"  Réduction FLOPs (cutting AMCA)       : {red_flops_cut:>5.1f}%\")\nprint(f\"  Réduction FLOPs (pruning structurel) : {red_flops_prun:>5.1f}%\")\nprint(f\"  Réduction FLOPs TOTALE               : {red_flops_total:>5.1f}%\")\nprint(f\"  Réduction params TOTALE              : {red_params_total:>5.1f}%\")\nprint(f\"  Filtres supprimés physiquement       : {n_removed}\")\nprint(\"=\"*84)\n","metadata":{"execution":{"iopub.execute_input":"2026-04-17T21:01:06.419709Z","iopub.status.busy":"2026-04-17T21:01:06.419410Z","iopub.status.idle":"2026-04-17T21:01:40.289287Z","shell.execute_reply":"2026-04-17T21:01:40.288425Z"},"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[]},"outputs":[],"execution_count":null},{"id":"80a28c10","cell_type":"markdown","source":"## 22 — Grad-CAM (Explicabilité)","metadata":{"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[]}},{"id":"24481bb4","cell_type":"code","source":"def last_conv(model):\n    \"\"\"Détecte la dernière Conv2D du modèle (parcours récursif sur les sous-modèles).\"\"\"\n    name = None\n    def recurse(m):\n        nonlocal name\n        for l in getattr(m, \"layers\", []):\n            if isinstance(l, tf.keras.layers.Conv2D):\n                name = l.name\n            elif hasattr(l, \"layers\"):\n                recurse(l)\n    recurse(model)\n    return name\n\n\ndef build_gradcam_model(model, conv_name, head_cfg):\n    \"\"\"\n    Construit un grad_model qui expose à la fois :\n      - le tenseur intermédiaire de la couche `conv_name`\n      - la prédiction finale du modèle\n\n    FIX bug Keras 3 : quand le backbone est un sous-modèle imbriqué,\n    on remplace `bb` par un `aux_bb` qui retourne `[intermediate, output]`\n    dans le même graphe, puis on rejoue SE+head et copie les poids.\n    \"\"\"\n    backbone = None; conv_layer = None\n    for layer in model.layers:\n        if isinstance(layer, tf.keras.Model):\n            try:\n                conv_layer = layer.get_layer(conv_name)\n                backbone   = layer\n                break\n            except (ValueError, KeyError):\n                continue\n\n    if backbone is None:\n        try: conv_layer = model.get_layer(conv_name)\n        except (ValueError, KeyError): return None\n        return tf.keras.Model(model.inputs, [conv_layer.output, model.output])\n\n    aux_bb = tf.keras.Model(\n        backbone.input,\n        [conv_layer.output, backbone.output],\n        name=\"aux_bb\"\n    )\n\n    new_inp = tf.keras.Input(shape=model.input_shape[1:], name=\"gc_input\")\n    intermediate, bb_feat = aux_bb(new_inp)\n    x   = se_block(bb_feat, ratio=16)\n    out = build_head(x, head_cfg)\n    grad_model = tf.keras.Model(new_inp, [intermediate, out], name=\"grad_model\")\n\n    bb_idx_orig = model.layers.index(backbone)\n    aux_idx_new = next(i for i, l in enumerate(grad_model.layers) if l is aux_bb)\n    orig_post   = model.layers[bb_idx_orig + 1:]\n    new_post    = grad_model.layers[aux_idx_new + 1:]\n\n    if len(orig_post) != len(new_post):\n        print(f\"[GRAD-CAM][WARN] mismatch couches : {len(orig_post)} vs {len(new_post)}\")\n\n    n_copied = 0\n    for o, n in zip(orig_post, new_post):\n        if o.get_weights():\n            try:\n                n.set_weights(o.get_weights())\n                n_copied += 1\n            except Exception as e:\n                print(f\"[GRAD-CAM][WARN] copie échouée pour {o.name}: {e}\")\n    print(f\"[GRAD-CAM] grad_model construit — {n_copied} couches de poids copiées\")\n    return grad_model\n\n\n@tf.function\ndef _gradcam_forward(grad_model, x):\n    with tf.GradientTape() as tape:\n        result = grad_model(x, training=False)\n        co     = result[0]\n        pred   = result[1]\n        tape.watch(co)\n        loss = pred[:, 0]\n    grads = tape.gradient(loss, co)\n    return co, grads\n\n\ndef gradcam(img_arr, grad_model):\n    co, grads = _gradcam_forward(grad_model, tf.cast(img_arr, tf.float32))\n    if grads is None: return None\n    pooled = tf.reduce_mean(grads, axis=(0, 1, 2))\n    h      = tf.squeeze(co[0] @ pooled[..., tf.newaxis])\n    h      = tf.maximum(h, 0)\n    h      = h / (tf.reduce_max(h) + 1e-8)\n    return h.numpy()\n\n\ndef overlay(img, heatmap, alpha=0.4):\n    h2 = cv2.resize(heatmap, (img.shape[1], img.shape[0]))\n    hc = (cm.jet(h2)[:, :, :3] * 255).astype(np.uint8)\n    i2 = (img * 255).astype(np.uint8) if img.max() <= 1.0 else img.astype(np.uint8)\n    return cv2.addWeighted(i2, 1 - alpha, hc, alpha, 0)\n\n\nLAST_CONV = last_conv(model)\nprint(f\"Dernière Conv2D détectée : {LAST_CONV}\")\n\ngrad_model = build_gradcam_model(model, LAST_CONV, amca_strategy[\"head\"])\n\nall_imgs, all_lbls = [], []\nfor batch_imgs, batch_lbls in val_ds:\n    all_imgs.append(batch_imgs.numpy())\n    all_lbls.append(batch_lbls.numpy())\n    if sum(a.shape[0] for a in all_imgs) >= 30:\n        break\nall_imgs = np.concatenate(all_imgs, axis=0)\nall_lbls = np.concatenate(all_lbls, axis=0)\n\nidx_dr   = np.where(all_lbls == 1)[0][:3]\nidx_nodr = np.where(all_lbls == 0)[0][:3]\nidxs     = np.concatenate([idx_dr, idx_nodr])\nimgs     = all_imgs[idxs]; labels = all_lbls[idxs]\n\n# Vérification\nif grad_model is not None and len(imgs) > 0:\n    p_orig = float(model.predict(imgs[0:1], verbose=0)[0][0])\n    p_grad = float(grad_model(imgs[0:1], training=False)[1].numpy()[0][0])\n    print(f\"[VERIF] model={p_orig:.6f}  grad_model={p_grad:.6f}  diff={abs(p_orig-p_grad):.2e}\")\n\nfig, axes = plt.subplots(len(imgs), 2, figsize=(10, len(imgs) * 4))\nfig.suptitle(\n    f\"Grad-CAM — {ARCHITECTURE} (cutting: {amca_strategy['cutting_point']})\\n\"\n    f\"Couche : {LAST_CONV}\", fontsize=12\n)\nfor i, (img, lbl) in enumerate(zip(imgs, labels)):\n    pred_val = float(model.predict(img[np.newaxis, ...], verbose=0)[0][0])\n    pred_lbl = \"DR\" if pred_val >= 0.5 else \"No-DR\"\n    true_lbl = \"DR\" if lbl else \"No-DR\"\n    correct  = \"✓\" if pred_lbl == true_lbl else \"✗\"\n\n    axes[i, 0].imshow(img, cmap=\"gray\" if img.shape[-1] == 1 else None)\n    axes[i, 0].axis(\"off\")\n    axes[i, 0].set_title(\n        f\"{correct} Vraie: {true_lbl}  |  Prédite: {pred_lbl} ({pred_val:.2f})\",\n        fontsize=9, color=\"green\" if pred_lbl == true_lbl else \"red\"\n    )\n\n    if grad_model is not None:\n        h = gradcam(img[np.newaxis, ...], grad_model)\n        if h is not None:\n            ov = overlay(img, h, alpha=0.45)\n            axes[i, 1].imshow(ov); axes[i, 1].axis(\"off\")\n            axes[i, 1].set_title(\"Grad-CAM (zones d'attention)\", fontsize=9)\n        else:\n            axes[i, 1].axis(\"off\"); axes[i, 1].set_title(\"Grad-CAM indisponible\", fontsize=9)\n    else:\n        axes[i, 1].axis(\"off\"); axes[i, 1].set_title(\"grad_model indisponible\", fontsize=9)\n\nplt.tight_layout()\nplt.savefig(\"gradcam_results.png\", dpi=150, bbox_inches=\"tight\")\nplt.show()\nprint(\"Grad-CAM sauvegardé → gradcam_results.png\")\n","metadata":{"execution":{"iopub.execute_input":"2026-04-17T21:01:41.544161Z","iopub.status.busy":"2026-04-17T21:01:41.543831Z","iopub.status.idle":"2026-04-17T21:01:51.988432Z","shell.execute_reply":"2026-04-17T21:01:51.987507Z"},"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[]},"outputs":[],"execution_count":null},{"id":"00633cac","cell_type":"markdown","source":"## 23 — Sauvegarde","metadata":{"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[]}},{"id":"e822e729","cell_type":"code","source":"model.save(\"amca_dr_final_pruned.keras\")\n\nacc_bl     = baseline_results.get(\"acc_youden\", baseline_results.get(\"acc\", 0.0))\nacc_before = r_before.get(\"acc_youden\", r_before.get(\"acc\", 0.0))\nacc_after  = r_after.get(\"acc_youden\",  r_after.get(\"acc\",  0.0))\n\nprint(\"\\n\" + \"=\"*60)\nprint(\"  RÉSUMÉ FINAL — VRAIS FLOPs MESURÉS\")\nprint(\"=\"*60)\nprint(f\"  Architecture       : {ARCHITECTURE}\")\nprint(f\"  Cutting point      : {amca_strategy['cutting_point']}\")\nprint(f\"  Pruning ratio      : {PRUNE_RATIO:.0%} des canaux internes (_1, _2)\")\nprint(f\"  Filtres supprimés  : {n_removed} (physiquement)\")\nprint(f\"  Params finaux      : {model.count_params():,}\")\nprint(f\"  FLOPs finaux       : {flops_final/1e9:.3f} G\")\nprint()\nprint(f\"  {'Modèle':<28} {'AUC':>8} {'Accuracy':>10} {'FLOPs':>9}\")\nprint(f\"  {'-'*58}\")\nprint(f\"  {'Baseline ResNet50':<28} {baseline_results['auc']:>8.4f} {acc_bl:>9.2%} {baseline_results['flops_g']:>8.3f}G\")\nprint(f\"  {'AMCA cutting seul':<28} {r_before['auc']:>8.4f} {acc_before:>9.2%} {flops_amca/1e9:>8.3f}G\")\nprint(f\"  {'AMCA + pruning final':<28} {r_after['auc']:>8.4f} {acc_after:>9.2%} {flops_final/1e9:>8.3f}G\")\nprint(f\"  {'='*58}\")\nprint()\nprint(f\"  ΔAUC      (final vs baseline) : {r_after['auc']-baseline_results['auc']:+.4f}\")\nprint(f\"  ΔAccuracy (final vs baseline) : {acc_after - acc_bl:+.2%}\")\nprint(f\"  Réduction FLOPs TOTALE        : {(1-flops_final/(baseline_results['flops_g']*1e9))*100:.1f}%\")\nprint(f\"  Réduction params TOTALE       : {(1-model.count_params()/baseline_results['params'])*100:.1f}%\")\nprint(f\"  Seuil optimal final           : {r_after.get('threshold', 0.5):.4f}\")\nprint(\"=\"*60)\n","metadata":{"execution":{"iopub.execute_input":"2026-04-17T21:01:53.346477Z","iopub.status.busy":"2026-04-17T21:01:53.345660Z","iopub.status.idle":"2026-04-17T21:01:53.798701Z","shell.execute_reply":"2026-04-17T21:01:53.797812Z"},"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[]},"outputs":[],"execution_count":null}]}