{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[],"dockerImageVersionId":28755,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session\n\n# Use the kagglehub client library to attach Kaggle resources like competitions, datasets, and models to your session\n# Learn more about kagglehub: https://github.com/Kaggle/kagglehub/blob/main/README.md\n\nimport kagglehub\n# kagglehub.dataset_download('<owner>/<dataset-slug>')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import json, os\n\n# Créer nouveau kaggle.json avec nouveau compte\nkaggle_creds = {\n    \"username\": \"bejaouikhouloud\",\n    \"key\": \"bdbf1a361b5663214f5e529ac033134f\"\n}\n\nos.makedirs(\"/root/.config/kaggle\", exist_ok=True)\nwith open(\"/root/.config/kaggle/kaggle.json\", \"w\") as f:\n    json.dump(kaggle_creds, f)\nos.chmod(\"/root/.config/kaggle/kaggle.json\", 0o600)\nprint(\"OK nouveau compte configuré\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import subprocess\n\n# Vérifier accès aux datasets de khouloudbejaoui20\ndatasets = [\n    \"khouloudbejaoui20/best-models\",\n    \"khouloudbejaoui20/panda-patches-run1\",\n    \"khouloudbejaoui20/panda-patches-run2\",\n    \"khouloudbejaoui20/panda-patches-run3\",\n    \"khouloudbejaoui20/panda-patches-run4\",\n    \"khouloudbejaoui20/dataset\",\n]\n\nfor ds in datasets:\n    r = subprocess.run([\n        \"kaggle\", \"datasets\", \"files\", ds\n    ], capture_output=True, text=True)\n    status = \"OK\" if \"pth\" in r.stdout or \"png\" in r.stdout else \"NON ACCESSIBLE\"\n    print(f\"{status} : {ds}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\n=============================================================\nStage 0 — Baseline 5-Fold CV — FROM SCRATCH\nDataset  : patches_400 (399 WSIs, 4,319 patches)\nTarget   : WSI-level QWK baseline\nNo PREV_CKPT — random initialization + ImageNet pretrained\n=============================================================\n\"\"\"\n\nimport os, glob, json, random, shutil, subprocess, datetime\nimport numpy as np\nfrom PIL import Image\nfrom collections import Counter\nfrom sklearn.model_selection import StratifiedGroupKFold\nfrom sklearn.metrics import cohen_kappa_score, accuracy_score\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as T\nimport torchvision.models as models\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\n\n# ─────────────────────────────────────────────\n# 1. CONFIG\n# ─────────────────────────────────────────────\nSHARD_ID    = 0\nSTART_FOLD  = 0\nDONE_KAPPAS = []\nDONE_CKPTS  = []\n\n# Stage 0 = FROM SCRATCH — pas de PREV_CKPT\nPREV_CKPT   = None\n\nN_FOLDS     = 5\nSEED        = 42\nBATCH_SIZE  = 32\nNUM_WORKERS = 2\nEPOCHS      = 25\nPATIENCE    = 5\nLR_CNN      = 5e-5\nLR_TR       = 1e-4\n\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\ntorch.backends.cudnn.deterministic = True\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Device    : {DEVICE}\")\nprint(f\"Stage     : {SHARD_ID} (FROM SCRATCH — no PREV_CKPT)\")\nprint(f\"EPOCHS={EPOCHS} | PATIENCE={PATIENCE} | LR_CNN={LR_CNN} | LR_TR={LR_TR}\")\n\n# ─────────────────────────────────────────────\n# 2. CHEMIN PATCHES\n# ─────────────────────────────────────────────\nBASE = \"/kaggle/input/datasets/khouloudbejaoui20/dataset\"\n\n#SHARD_PATHS = {\n#    0: f\"{BASE}/patches_400-20260630T212405Z-3-001/patches_400\",\n#}\nSHARD_PATHS = {\n    0: \"/kaggle/input/datasets/khouloudbejaoui20/dataset/patches_400-20260630T212405Z-3-001/patches_400\",\n}\n\n\nSAVE_DIR     = \"/kaggle/working\"\nKAGGLE_DS    = \"khouloudbejaoui20/best-models\"\nRESULTS_JSON = f\"{SAVE_DIR}/stage0_cv_results.json\"\n\n# ─────────────────────────────────────────────\n# 3. VERIFICATION\n# ─────────────────────────────────────────────\nprint(f\"\\n── Verification ────────────────────────\")\npath = SHARD_PATHS[0]\nif os.path.exists(path):\n    files = glob.glob(f\"{path}/*.png\")\n    if not files:\n        files = glob.glob(f\"{path}/**/*.png\", recursive=True)\n    print(f\"  OK Shard 0: {len(files)} patches -> {path}\")\nelse:\n    print(f\"  ERREUR: {path} non trouve!\")\n\n# ─────────────────────────────────────────────\n# 4. CHARGEMENT PATCHES\n# ─────────────────────────────────────────────\nprint(f\"\\n── Patches Stage 0 ──────────────────────\")\n\ndef load_patches(folder):\n    files = glob.glob(f\"{folder}/*.png\")\n    if not files:\n        files = glob.glob(f\"{folder}/**/*.png\", recursive=True)\n    records = []\n    for fp in files:\n        fname = os.path.basename(fp)\n        try:\n            parts  = fname.split(\"_slide\")\n            label  = int(parts[0].replace(\"grade\", \"\"))\n            wsi_id = parts[1].split(\"_\")[0]\n            records.append({\"path\": fp, \"wsi_id\": wsi_id, \"label\": label})\n        except:\n            continue\n    return records\n\nall_records = load_patches(SHARD_PATHS[0])\n\n# Dedupliquer\nseen = set()\nunique_records = []\nfor r in all_records:\n    fname = os.path.basename(r[\"path\"])\n    if fname not in seen:\n        seen.add(fname)\n        unique_records.append(r)\n\nall_wsis = list(set(r[\"wsi_id\"] for r in unique_records))\nprint(f\"  Total unique : {len(unique_records)} patches\")\nprint(f\"  Total WSIs  : {len(all_wsis)}\")\ncnt = Counter(r[\"label\"] for r in unique_records)\nfor c in range(6):\n    print(f\"    ISUP {c}: {cnt.get(c, 0)}\")\n\nassert len(unique_records) > 100, \"Trop peu de patches!\"\nprint(f\"\\n  OK {len(all_wsis)} WSIs uniques confirmes\")\n\n# ─────────────────────────────────────────────\n# 5. TRANSFORMS\n# ─────────────────────────────────────────────\nMEAN = [0.485, 0.456, 0.406]\nSTD  = [0.229, 0.224, 0.225]\n\ntrain_tf = T.Compose([\n    T.Resize((224, 224)),\n    T.RandomHorizontalFlip(),\n    T.RandomVerticalFlip(),\n    T.RandomRotation(90),\n    T.ColorJitter(brightness=0.2, contrast=0.2,\n                  saturation=0.1, hue=0.05),\n    T.RandomApply([T.GaussianBlur(3)], p=0.2),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n])\n\nval_tf = T.Compose([\n    T.Resize((224, 224)),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n])\n\nclass PatchDataset(Dataset):\n    def __init__(self, records, transform):\n        self.records   = records\n        self.transform = transform\n    def __len__(self):\n        return len(self.records)\n    def __getitem__(self, idx):\n        r = self.records[idx]\n        try:\n            img = Image.open(r[\"path\"]).convert(\"RGB\")\n        except:\n            img = Image.new(\"RGB\", (224, 224), (255, 255, 255))\n        return self.transform(img), int(r[\"label\"])\n\n# ─────────────────────────────────────────────\n# 6. MODELE\n# ─────────────────────────────────────────────\nclass HybridCNNTransformer(nn.Module):\n    def __init__(self, num_classes=6, d_model=512, nhead=8,\n                 num_layers=2, dim_feedforward=2048, dropout=0.1):\n        super().__init__()\n        backbone = models.resnet50(weights=models.ResNet50_Weights.DEFAULT)\n        self.cnn_backbone = nn.Sequential(*list(backbone.children())[:-2])\n        self.pool         = nn.AdaptiveAvgPool2d((1, 1))\n        self.proj         = nn.Sequential(\n            nn.Identity(), nn.Identity(),\n            nn.Linear(2048, d_model),\n            nn.LayerNorm(d_model)\n        )\n        self.cls_token = nn.Parameter(torch.zeros(1, 1, d_model))\n        self.pos_embed = nn.Parameter(torch.zeros(1, 2, d_model))\n        encoder_layer  = nn.TransformerEncoderLayer(\n            d_model=d_model, nhead=nhead,\n            dim_feedforward=dim_feedforward,\n            dropout=dropout, batch_first=True\n        )\n        self.transformer = nn.TransformerEncoder(\n            encoder_layer, num_layers=num_layers,\n            norm=nn.LayerNorm(d_model)\n        )\n        self.classifier = nn.Sequential(\n            nn.Linear(d_model, 256), nn.ReLU(),\n            nn.Dropout(dropout), nn.Linear(256, num_classes)\n        )\n        nn.init.trunc_normal_(self.cls_token, std=0.02)\n        nn.init.trunc_normal_(self.pos_embed, std=0.02)\n\n    def forward(self, x):\n        f   = self.pool(self.cnn_backbone(x)).flatten(1)\n        f   = self.proj(f).unsqueeze(1)\n        cls = self.cls_token.expand(x.size(0), -1, -1)\n        seq = torch.cat([cls, f], dim=1) + self.pos_embed\n        return self.classifier(self.transformer(seq)[:, 0])\n\ndef build_model():\n    \"\"\"Stage 0 — from scratch, no PREV_CKPT\"\"\"\n    model  = HybridCNNTransformer().to(DEVICE)\n    n_gpus = torch.cuda.device_count()\n    if n_gpus > 1:\n        model = nn.DataParallel(model)\n        print(f\"  OK DataParallel: {n_gpus} GPUs\")\n    n_param = sum(p.numel() for p in model.parameters())\n    assert n_param == 30_998_342, f\"Mismatch: {n_param}\"\n    print(f\"  OK Parameters: {n_param:,}\")\n    print(f\"  OK FROM SCRATCH (ImageNet pretrained ResNet-50)\")\n    return model\n\n# ─────────────────────────────────────────────\n# 7. WSI-LEVEL QWK\n# ─────────────────────────────────────────────\ndef wsi_level_qwk(model, val_data, device):\n    model.eval()\n    wsi_preds  = {}\n    wsi_labels = {}\n    all_preds  = []\n    all_labels = []\n    loader = DataLoader(\n        PatchDataset(val_data, val_tf),\n        batch_size=BATCH_SIZE, shuffle=False,\n        num_workers=NUM_WORKERS\n    )\n    with torch.no_grad():\n        for imgs, labels in loader:\n            imgs  = imgs.to(device)\n            preds = model(imgs).argmax(1).cpu().numpy()\n            all_preds.extend(preds)\n            all_labels.extend(labels.numpy())\n    for i, r in enumerate(val_data):\n        wid = r[\"wsi_id\"]\n        if wid not in wsi_preds:\n            wsi_preds[wid]  = []\n            wsi_labels[wid] = r[\"label\"]\n        wsi_preds[wid].append(all_preds[i])\n    wf = []; wl = []\n    for wid in wsi_preds:\n        votes = wsi_preds[wid]\n        wf.append(max(set(votes), key=votes.count))\n        wl.append(wsi_labels[wid])\n    patch_qwk = cohen_kappa_score(all_labels, all_preds, weights=\"quadratic\")\n    wsi_qwk   = cohen_kappa_score(wl, wf, weights=\"quadratic\")\n    return patch_qwk, wsi_qwk, len(wf)\n\n# ─────────────────────────────────────────────\n# 8. PUSH KAGGLE\n# ─────────────────────────────────────────────\ndef push_to_kaggle(ckpt_path, message):\n    try:\n        tmp = f\"{SAVE_DIR}/_push_tmp\"\n        os.makedirs(tmp, exist_ok=True)\n        for c in glob.glob(f\"{SAVE_DIR}/stage{SHARD_ID}_fold*.pth\"):\n            shutil.copy(c, tmp)\n        shutil.copy(ckpt_path, tmp)\n        if os.path.exists(RESULTS_JSON):\n            shutil.copy(RESULTS_JSON, tmp)\n        json.dump({\n            \"title\": \"Memoire IHC Saved Models\",\n            \"id\": KAGGLE_DS,\n            \"licenses\": [{\"name\": \"CC0-1.0\"}]\n        }, open(f\"{tmp}/dataset-metadata.json\", \"w\"))\n        files_pushed = [os.path.basename(f)\n                        for f in glob.glob(f\"{tmp}/*.pth\")]\n        print(f\"   Pushing: {files_pushed}\")\n        r = subprocess.run(\n            [\"kaggle\", \"datasets\", \"version\",\n             \"-p\", tmp, \"-m\", message,\n             \"--dir-mode\", \"skip\"],\n            capture_output=True, text=True, timeout=300\n        )\n        print(f\"   {'OK Kaggle' if r.returncode==0 else 'ERREUR '+r.stderr[:100]}\")\n    except Exception as e:\n        print(f\"   ERREUR push: {e}\")\n\n# ─────────────────────────────────────────────\n# 9. TRAIN ONE FOLD\n# ─────────────────────────────────────────────\ndef train_one_fold(fold_id, train_data, val_data):\n    print(f\"\\n  Fold {fold_id+1}/{N_FOLDS} ──────────────────────────\")\n    print(f\"     Train: {len(train_data)} | Val: {len(val_data)}\")\n\n    train_loader = DataLoader(\n        PatchDataset(train_data, train_tf),\n        batch_size=BATCH_SIZE, shuffle=True,\n        num_workers=NUM_WORKERS, pin_memory=True\n    )\n\n    model = build_model()  # FROM SCRATCH\n\n    # Class weights\n    cnt_tr = Counter(r[\"label\"] for r in train_data)\n    w = torch.tensor(\n        [1.0 / max(cnt_tr.get(c, 1), 1) for c in range(6)],\n        dtype=torch.float\n    )\n    w = (w / w.sum() * 6).to(DEVICE)\n    criterion = nn.CrossEntropyLoss(weight=w)\n\n    # Differential LR\n    base_model = model.module if hasattr(model, \"module\") else model\n    optimizer  = optim.AdamW([\n        {\"params\": base_model.cnn_backbone.parameters(), \"lr\": LR_CNN},\n        {\"params\": [p for n, p in base_model.named_parameters()\n                    if not n.startswith(\"cnn_backbone\")], \"lr\": LR_TR},\n    ], weight_decay=1e-4)\n\n    scheduler  = CosineAnnealingLR(optimizer, T_max=EPOCHS, eta_min=1e-7)\n\n    best_kappa = 0.0\n    no_improve = 0\n    best_ckpt  = f\"{SAVE_DIR}/stage{SHARD_ID}_fold{fold_id+1}_best.pth\"\n\n    for epoch in range(1, EPOCHS + 1):\n        model.train()\n        t_loss, t_correct, t_total = 0.0, 0, 0\n        for imgs, labels in train_loader:\n            imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n            optimizer.zero_grad()\n            out  = model(imgs)\n            loss = criterion(out, labels)\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n            optimizer.step()\n            t_loss    += loss.item() * imgs.size(0)\n            t_correct += (out.argmax(1) == labels).sum().item()\n            t_total   += imgs.size(0)\n        scheduler.step()\n\n        # WSI-level QWK\n        _, wsi_kappa, n_wsis_val = wsi_level_qwk(model, val_data, DEVICE)\n        is_best = wsi_kappa > best_kappa\n\n        print(f\"     Ep[{epoch:02d}/{EPOCHS}] \"\n              f\"loss={t_loss/t_total:.4f} acc={t_correct/t_total:.3f} | \"\n              f\"WSI_QWK={wsi_kappa:.4f} ({n_wsis_val} WSIs)\"\n              + (\" *\" if is_best else \"\"))\n\n        if is_best:\n            best_kappa = wsi_kappa\n            no_improve = 0\n            torch.save({\n                \"epoch\":            epoch,\n                \"model_state_dict\": model.state_dict() if not hasattr(model, \"module\")\n                                    else model.module.state_dict(),\n                \"kappa\":            wsi_kappa,\n                \"shard_id\":         SHARD_ID,\n                \"fold_id\":          fold_id,\n                \"n_wsis\":           len(all_wsis),\n                \"hyperparams\": {\n                    \"lr_cnn\":    LR_CNN,\n                    \"lr_tr\":     LR_TR,\n                    \"epochs\":    EPOCHS,\n                    \"patience\":  PATIENCE,\n                    \"from_scratch\": True\n                }\n            }, best_ckpt)\n        else:\n            no_improve += 1\n            if no_improve >= PATIENCE:\n                print(f\"     Early stopping epoch {epoch}\")\n                break\n\n    print(f\"  OK Fold {fold_id+1} best WSI_QWK: {best_kappa:.4f}\")\n    push_to_kaggle(best_ckpt,\n        f\"Stage{SHARD_ID} Fold{fold_id+1} WSI_QWK={best_kappa:.4f}\")\n    return best_kappa, best_ckpt\n\n# ─────────────────────────────────────────────\n# 10. 5-FOLD CV\n# ─────────────────────────────────────────────\nprint(f\"\\n{'='*60}\")\nprint(f\"Stage {SHARD_ID} -- {len(all_wsis)} WSIs -- 5-Fold CV FROM SCRATCH\")\nprint(f\"{'='*60}\")\n\nwsi_ids    = [r[\"wsi_id\"] for r in unique_records]\nlabels_arr = [r[\"label\"]  for r in unique_records]\nindices    = list(range(len(unique_records)))\n\nsgkf = StratifiedGroupKFold(n_splits=N_FOLDS, shuffle=True, random_state=SEED)\n\nfold_kappas       = DONE_KAPPAS.copy()\nfold_ckpts        = DONE_CKPTS.copy()\nbest_overall      = max(DONE_KAPPAS) if DONE_KAPPAS else -1.0\nbest_ckpt_overall = None\n\nfor fold_id, (train_idx, val_idx) in enumerate(\n        sgkf.split(indices, labels_arr, groups=wsi_ids)):\n\n    if fold_id < START_FOLD:\n        continue\n\n    train_data = [unique_records[i] for i in train_idx]\n    val_data   = [unique_records[i] for i in val_idx]\n\n    # Anti-leakage assertion\n    tw = set(unique_records[i][\"wsi_id\"] for i in train_idx)\n    vw = set(unique_records[i][\"wsi_id\"] for i in val_idx)\n    assert len(tw & vw) == 0, f\"DATA LEAKAGE fold {fold_id+1}!\"\n\n    kappa, ckpt_path = train_one_fold(fold_id, train_data, val_data)\n    fold_kappas.append(kappa)\n    fold_ckpts.append(ckpt_path)\n\n    if kappa > best_overall:\n        best_overall      = kappa\n        best_ckpt_overall = ckpt_path\n\n# ─────────────────────────────────────────────\n# 11. RESUME FINAL\n# ─────────────────────────────────────────────\nmean_kappa = np.mean(fold_kappas)\nstd_kappa  = np.std(fold_kappas)\n\nprint(f\"\\n{'='*60}\")\nprint(f\"OK Stage {SHARD_ID} -- 5-Fold CV Results\")\nprint(f\"{'='*60}\")\nfor i, k in enumerate(fold_kappas):\n    print(f\"   Fold {i+1}: QWK = {k:.4f}\")\nprint(f\"{'─'*40}\")\nprint(f\"   Mean QWK : {mean_kappa:.4f} +/- {std_kappa:.4f}\")\nprint(f\"   Best QWK : {best_overall:.4f}\")\nprint(f\"   WSIs     : {len(all_wsis)}\")\nprint(f\"   Expected : ~0.57 (baseline from scratch)\")\nprint(f\"{'='*60}\")\n\n# Save JSON\nresults = {\n    \"stage\":       SHARD_ID,\n    \"n_wsis\":      len(all_wsis),\n    \"n_patches\":   len(unique_records),\n    \"fold_kappas\": fold_kappas,\n    \"mean_kappa\":  mean_kappa,\n    \"std_kappa\":   std_kappa,\n    \"best_kappa\":  best_overall,\n    \"timestamp\":   datetime.datetime.now().isoformat(),\n    \"from_scratch\": True,\n}\nwith open(RESULTS_JSON, \"w\") as f:\n    json.dump(results, f, indent=2)\n\n# Stage best checkpoint\nif best_ckpt_overall:\n    stage_best = f\"{SAVE_DIR}/stage{SHARD_ID}_best_qwk{best_overall:.4f}.pth\"\n    shutil.copy(best_ckpt_overall, stage_best)\n    print(f\"\\nOK Best checkpoint: {stage_best}\")\n    print(f\"\\n-> Pour Stage 1 :\")\n    print(f\"   SHARD_ID  = 1\")\n    print(f\"   PREV_CKPT = '{stage_best}'\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Après Fold 5 — vérifie l'assertion\nfrom sklearn.model_selection import StratifiedGroupKFold\nimport numpy as np\n\nsgkf = StratifiedGroupKFold(n_splits=5, shuffle=True, random_state=42)\nwsi_ids = [r[\"wsi_id\"] for r in unique_records]\nlabels  = [r[\"label\"]  for r in unique_records]\nindices = list(range(len(unique_records)))\n\nfor fold_id, (train_idx, val_idx) in enumerate(\n        sgkf.split(indices, labels, groups=wsi_ids)):\n    tw = set(unique_records[i][\"wsi_id\"] for i in train_idx)\n    vw = set(unique_records[i][\"wsi_id\"] for i in val_idx)\n    overlap = len(tw & vw)\n    print(f\"Fold {fold_id+1}: overlap={overlap} {'OK' if overlap==0 else 'LEAKAGE!'}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, glob, shutil\n\n# Voir tous les checkpoints\nckpts = glob.glob(\"/kaggle/working/stage0_fold*.pth\")\nprint(\"Checkpoints disponibles:\")\nfor c in ckpts:\n    mb = os.path.getsize(c)/(1024*1024)\n    print(f\"  {os.path.basename(c)} ({mb:.1f} MB)\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import shutil, subprocess, json, glob, os\n\ntmp = \"/kaggle/working/_push_stage0\"\nos.makedirs(tmp, exist_ok=True)\n\n# Copier tous les folds stage0\nfor f in glob.glob(\"/kaggle/working/stage0_fold*.pth\"):\n    shutil.copy(f, tmp)\n    print(f\"  OK {os.path.basename(f)}\")\n\njson.dump({\n    \"title\": \"Memoire IHC Saved Models\",\n    \"id\": \"khouloudbejaoui20/best-models\",\n    \"licenses\": [{\"name\": \"CC0-1.0\"}]\n}, open(f\"{tmp}/dataset-metadata.json\", \"w\"))\n\nr = subprocess.run([\n    \"kaggle\", \"datasets\", \"version\",\n    \"-p\", tmp,\n    \"-m\", \"Stage0 all folds\",\n    \"--dir-mode\", \"skip\"\n], capture_output=True, text=True, timeout=300)\n\nprint(\"OK!\" if r.returncode == 0 else r.stderr[:200])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, glob, shutil, subprocess, json\n\n# Vérifier checkpoint\nckpt = \"/kaggle/working/stage0_fold1_best.pth\"\nprint(f\"Checkpoint: {'OK' if os.path.exists(ckpt) else 'NON TROUVE'}\")\n\nif os.path.exists(ckpt):\n    tmp = \"/kaggle/working/_push_stage0\"\n    os.makedirs(tmp, exist_ok=True)\n    shutil.copy(ckpt, tmp)\n    \n    json.dump({\n        \"title\": \"Memoire IHC Saved Models\",\n        \"id\": \"khouloudbejaoui20/best-models\",\n        \"licenses\": [{\"name\": \"CC0-1.0\"}]\n    }, open(f\"{tmp}/dataset-metadata.json\", \"w\"))\n    \n    r = subprocess.run([\n        \"kaggle\", \"datasets\", \"version\",\n        \"-p\", tmp,\n        \"-m\", \"Stage0 Fold1 best checkpoint\",\n        \"--dir-mode\", \"skip\"\n    ], capture_output=True, text=True, timeout=300)\n    \n    print(\"OK!\" if r.returncode == 0 else f\"ERREUR: {r.stderr[:200]}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import glob, os\n\n# Chercher TOUS les .pth sur le système\nfiles = glob.glob(\"/kaggle/**/*.pth\", recursive=True)\nprint(f\"Total .pth trouvés: {len(files)}\")\nfor f in files:\n    size_mb = os.path.getsize(f) / (1024*1024)\n    print(f\"  {f} ({size_mb:.1f} MB)\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nprint(f\"SAVE_DIR existe: {os.path.exists('/kaggle/working')}\")\nprint(f\"Contenu /kaggle/working:\")\nfor f in os.listdir('/kaggle/working'):\n    print(f\"  {f}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\n=============================================================\nStage 1 — Incremental 5-Fold CV\nFine-tune depuis Stage 0 (best QWK=0.7746)\nDonnées cumulatives :\n  Shard 0 : patches_400     (399 WSIs,  4,319 patches)\n  Shard 1 : panda-patches-run1 (1,000 WSIs, 14,842 patches)\nTotal     : ~1,399 WSIs, ~19,161 patches\n=============================================================\n\"\"\"\n\nimport os, glob, json, random, shutil, subprocess, datetime\nimport numpy as np\nfrom PIL import Image\nfrom collections import Counter\nfrom sklearn.model_selection import StratifiedGroupKFold\nfrom sklearn.metrics import cohen_kappa_score, accuracy_score\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as T\nimport torchvision.models as models\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\n\n# ─────────────────────────────────────────────\n# 1. CONFIG\n# ─────────────────────────────────────────────\nSHARD_ID    = 1\nSTART_FOLD  = 0\nDONE_KAPPAS = []\nDONE_CKPTS  = []\n\nPREV_CKPT   = \"/kaggle/working/stage0_best_qwk0.7746.pth\"\n\n\n\n\n\n\n\n\n\n\n\n\n\nN_FOLDS     = 5\nSEED        = 42\nBATCH_SIZE  = 32\nNUM_WORKERS = 2\nEPOCHS      = 20\nPATIENCE    = 5\nLR_CNN      = 1e-5\nLR_TR       = 2e-5\n\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\ntorch.backends.cudnn.deterministic = True\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Device    : {DEVICE}\")\nprint(f\"Stage     : {SHARD_ID} (fine-tune depuis Stage 0 QWK=0.7746)\")\nprint(f\"EPOCHS={EPOCHS} | PATIENCE={PATIENCE} | LR_CNN={LR_CNN} | LR_TR={LR_TR}\")\n\n# ─────────────────────────────────────────────\n# 2. CHEMINS PATCHES\n# ─────────────────────────────────────────────\nBASE = \"/kaggle/input/datasets/khouloudbejaoui20\"\n\nSHARD_PATHS = {\n    0: f\"{BASE}/dataset/patches_400-20260630T212405Z-3-001/patches_400\",\n    1: f\"{BASE}/panda-patches-run1/kaggle/working/patches_panda_new\",\n}\n\nSAVE_DIR     = \"/kaggle/working\"\nKAGGLE_DS    = \"khouloudbejaoui20/best-models\"\nRESULTS_JSON = f\"{SAVE_DIR}/stage1_cv_results.json\"\n\n# ─────────────────────────────────────────────\n# 3. VERIFICATION\n# ─────────────────────────────────────────────\nprint(f\"\\n── Verification ────────────────────────\")\nassert os.path.exists(PREV_CKPT), f\"PREV_CKPT non trouve : {PREV_CKPT}\"\nprint(f\"  OK Stage 0 checkpoint trouve (QWK=0.7746)\")\n\nfor shard_id in range(SHARD_ID + 1):\n    path = SHARD_PATHS.get(shard_id, \"\")\n    if os.path.exists(path):\n        files = glob.glob(f\"{path}/*.png\")\n        if not files:\n            files = glob.glob(f\"{path}/**/*.png\", recursive=True)\n        print(f\"  OK Shard {shard_id}: {len(files)} patches -> {path}\")\n    else:\n        print(f\"  MANQUANT Shard {shard_id}: {path}\")\n\n# ─────────────────────────────────────────────\n# 4. CHARGEMENT PATCHES\n# ─────────────────────────────────────────────\nprint(f\"\\n── Patches cumulatifs (Shard 0 a {SHARD_ID}) ──\")\n\ndef load_patches(folder):\n    files = glob.glob(f\"{folder}/*.png\")\n    if not files:\n        files = glob.glob(f\"{folder}/**/*.png\", recursive=True)\n    records = []\n    for fp in files:\n        fname = os.path.basename(fp)\n        try:\n            parts  = fname.split(\"_slide\")\n            label  = int(parts[0].replace(\"grade\", \"\"))\n            wsi_id = parts[1].split(\"_\")[0]\n            records.append({\"path\": fp, \"wsi_id\": wsi_id, \"label\": label})\n        except:\n            continue\n    return records\n\nall_records = []\nfor shard_id in range(SHARD_ID + 1):\n    path = SHARD_PATHS.get(shard_id, \"\")\n    if path and os.path.exists(path):\n        r = load_patches(path)\n        all_records.extend(r)\n        print(f\"  Shard {shard_id}: {len(r)} patches\")\n    else:\n        print(f\"  Shard {shard_id}: NON TROUVE\")\n\n# Dedupliquer\nseen = set()\nunique_records = []\nfor r in all_records:\n    fname = os.path.basename(r[\"path\"])\n    if fname not in seen:\n        seen.add(fname)\n        unique_records.append(r)\n\nall_wsis = list(set(r[\"wsi_id\"] for r in unique_records))\nprint(f\"\\n  Total unique : {len(unique_records)} patches\")\nprint(f\"  Total WSIs  : {len(all_wsis)}\")\ncnt = Counter(r[\"label\"] for r in unique_records)\nfor c in range(6):\n    print(f\"    ISUP {c}: {cnt.get(c, 0)}\")\n\nassert len(all_wsis) > 500, f\"Trop peu de WSIs ({len(all_wsis)})\"\nprint(f\"\\n  OK {len(all_wsis)} WSIs uniques confirmes\")\n\n# ─────────────────────────────────────────────\n# 5. TRANSFORMS\n# ─────────────────────────────────────────────\nMEAN = [0.485, 0.456, 0.406]\nSTD  = [0.229, 0.224, 0.225]\n\ntrain_tf = T.Compose([\n    T.Resize((224, 224)),\n    T.RandomHorizontalFlip(),\n    T.RandomVerticalFlip(),\n    T.RandomRotation(90),\n    T.ColorJitter(brightness=0.2, contrast=0.2,\n                  saturation=0.1, hue=0.05),\n    T.RandomApply([T.GaussianBlur(3)], p=0.2),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n])\n\nval_tf = T.Compose([\n    T.Resize((224, 224)),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n])\n\nclass PatchDataset(Dataset):\n    def __init__(self, records, transform):\n        self.records   = records\n        self.transform = transform\n    def __len__(self):\n        return len(self.records)\n    def __getitem__(self, idx):\n        r = self.records[idx]\n        try:\n            img = Image.open(r[\"path\"]).convert(\"RGB\")\n        except:\n            img = Image.new(\"RGB\", (224, 224), (255, 255, 255))\n        return self.transform(img), int(r[\"label\"])\n\n# ─────────────────────────────────────────────\n# 6. MODELE\n# ─────────────────────────────────────────────\nclass HybridCNNTransformer(nn.Module):\n    def __init__(self, num_classes=6, d_model=512, nhead=8,\n                 num_layers=2, dim_feedforward=2048, dropout=0.1):\n        super().__init__()\n        backbone = models.resnet50(weights=models.ResNet50_Weights.DEFAULT)\n        self.cnn_backbone = nn.Sequential(*list(backbone.children())[:-2])\n        self.pool         = nn.AdaptiveAvgPool2d((1, 1))\n        self.proj         = nn.Sequential(\n            nn.Identity(), nn.Identity(),\n            nn.Linear(2048, d_model),\n            nn.LayerNorm(d_model)\n        )\n        self.cls_token = nn.Parameter(torch.zeros(1, 1, d_model))\n        self.pos_embed = nn.Parameter(torch.zeros(1, 2, d_model))\n        encoder_layer  = nn.TransformerEncoderLayer(\n            d_model=d_model, nhead=nhead,\n            dim_feedforward=dim_feedforward,\n            dropout=dropout, batch_first=True\n        )\n        self.transformer = nn.TransformerEncoder(\n            encoder_layer, num_layers=num_layers,\n            norm=nn.LayerNorm(d_model)\n        )\n        self.classifier = nn.Sequential(\n            nn.Linear(d_model, 256), nn.ReLU(),\n            nn.Dropout(dropout), nn.Linear(256, num_classes)\n        )\n        nn.init.trunc_normal_(self.cls_token, std=0.02)\n        nn.init.trunc_normal_(self.pos_embed, std=0.02)\n\n    def forward(self, x):\n        f   = self.pool(self.cnn_backbone(x)).flatten(1)\n        f   = self.proj(f).unsqueeze(1)\n        cls = self.cls_token.expand(x.size(0), -1, -1)\n        seq = torch.cat([cls, f], dim=1) + self.pos_embed\n        return self.classifier(self.transformer(seq)[:, 0])\n\ndef build_model(prev_ckpt=None):\n    model  = HybridCNNTransformer().to(DEVICE)\n    n_gpus = torch.cuda.device_count()\n    if n_gpus > 1:\n        model = nn.DataParallel(model)\n        print(f\"  OK DataParallel: {n_gpus} GPUs\")\n    n_param = sum(p.numel() for p in model.parameters())\n    assert n_param == 30_998_342, f\"Mismatch: {n_param}\"\n    print(f\"  OK Parameters: {n_param:,}\")\n    if prev_ckpt and os.path.exists(prev_ckpt):\n        ckpt  = torch.load(prev_ckpt, map_location=DEVICE, weights_only=False)\n        state = ckpt[\"model_state_dict\"]\n        if n_gpus > 1 and not list(state.keys())[0].startswith(\"module.\"):\n            state = {\"module.\" + k: v for k, v in state.items()}\n        elif n_gpus <= 1 and list(state.keys())[0].startswith(\"module.\"):\n            state = {k[7:]: v for k, v in state.items()}\n        model.load_state_dict(state)\n        print(f\"  OK Loaded Stage 0 (QWK={ckpt.get('kappa', '?')})\")\n    return model\n\n# ─────────────────────────────────────────────\n# 7. WSI-LEVEL QWK\n# ─────────────────────────────────────────────\ndef wsi_level_qwk(model, val_data, device):\n    model.eval()\n    wsi_preds  = {}\n    wsi_labels = {}\n    all_preds  = []\n    all_labels = []\n    loader = DataLoader(\n        PatchDataset(val_data, val_tf),\n        batch_size=BATCH_SIZE, shuffle=False,\n        num_workers=NUM_WORKERS\n    )\n    with torch.no_grad():\n        for imgs, labels in loader:\n            imgs  = imgs.to(device)\n            preds = model(imgs).argmax(1).cpu().numpy()\n            all_preds.extend(preds)\n            all_labels.extend(labels.numpy())\n    for i, r in enumerate(val_data):\n        wid = r[\"wsi_id\"]\n        if wid not in wsi_preds:\n            wsi_preds[wid]  = []\n            wsi_labels[wid] = r[\"label\"]\n        wsi_preds[wid].append(all_preds[i])\n    wf = []; wl = []\n    for wid in wsi_preds:\n        votes = wsi_preds[wid]\n        wf.append(max(set(votes), key=votes.count))\n        wl.append(wsi_labels[wid])\n    patch_qwk = cohen_kappa_score(all_labels, all_preds, weights=\"quadratic\")\n    wsi_qwk   = cohen_kappa_score(wl, wf, weights=\"quadratic\")\n    return patch_qwk, wsi_qwk, len(wf)\n\n# ─────────────────────────────────────────────\n# 8. PUSH KAGGLE\n# ─────────────────────────────────────────────\ndef push_to_kaggle(ckpt_path, message):\n    try:\n        tmp = f\"{SAVE_DIR}/_push_tmp\"\n        os.makedirs(tmp, exist_ok=True)\n        for c in glob.glob(f\"{SAVE_DIR}/stage{SHARD_ID}_fold*.pth\"):\n            shutil.copy(c, tmp)\n        shutil.copy(ckpt_path, tmp)\n        if os.path.exists(RESULTS_JSON):\n            shutil.copy(RESULTS_JSON, tmp)\n        json.dump({\n            \"title\": \"Memoire IHC Saved Models\",\n            \"id\": KAGGLE_DS,\n            \"licenses\": [{\"name\": \"CC0-1.0\"}]\n        }, open(f\"{tmp}/dataset-metadata.json\", \"w\"))\n        files_pushed = [os.path.basename(f)\n                        for f in glob.glob(f\"{tmp}/*.pth\")]\n        print(f\"   Pushing: {files_pushed}\")\n        r = subprocess.run(\n            [\"kaggle\", \"datasets\", \"version\",\n             \"-p\", tmp, \"-m\", message,\n             \"--dir-mode\", \"skip\"],\n            capture_output=True, text=True, timeout=300\n        )\n        print(f\"   {'OK Kaggle' if r.returncode==0 else 'ERREUR '+r.stderr[:100]}\")\n    except Exception as e:\n        print(f\"   ERREUR push: {e}\")\n\n# ─────────────────────────────────────────────\n# 9. TRAIN ONE FOLD\n# ─────────────────────────────────────────────\ndef train_one_fold(fold_id, train_data, val_data):\n    print(f\"\\n  Fold {fold_id+1}/{N_FOLDS} ──────────────────────────\")\n    print(f\"     Train: {len(train_data)} | Val: {len(val_data)}\")\n\n    train_loader = DataLoader(\n        PatchDataset(train_data, train_tf),\n        batch_size=BATCH_SIZE, shuffle=True,\n        num_workers=NUM_WORKERS, pin_memory=True\n    )\n\n    model = build_model(PREV_CKPT)\n\n    cnt_tr = Counter(r[\"label\"] for r in train_data)\n    w = torch.tensor(\n        [1.0 / max(cnt_tr.get(c, 1), 1) for c in range(6)],\n        dtype=torch.float\n    )\n    w = (w / w.sum() * 6).to(DEVICE)\n    criterion = nn.CrossEntropyLoss(weight=w)\n\n    base_model = model.module if hasattr(model, \"module\") else model\n    optimizer  = optim.AdamW([\n        {\"params\": base_model.cnn_backbone.parameters(), \"lr\": LR_CNN},\n        {\"params\": [p for n, p in base_model.named_parameters()\n                    if not n.startswith(\"cnn_backbone\")], \"lr\": LR_TR},\n    ], weight_decay=1e-4)\n\n    scheduler  = CosineAnnealingLR(optimizer, T_max=EPOCHS, eta_min=1e-7)\n\n    best_kappa = 0.0\n    no_improve = 0\n    best_ckpt  = f\"{SAVE_DIR}/stage{SHARD_ID}_fold{fold_id+1}_best.pth\"\n\n    for epoch in range(1, EPOCHS + 1):\n        model.train()\n        t_loss, t_correct, t_total = 0.0, 0, 0\n        for imgs, labels in train_loader:\n            imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n            optimizer.zero_grad()\n            out  = model(imgs)\n            loss = criterion(out, labels)\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n            optimizer.step()\n            t_loss    += loss.item() * imgs.size(0)\n            t_correct += (out.argmax(1) == labels).sum().item()\n            t_total   += imgs.size(0)\n        scheduler.step()\n\n        _, wsi_kappa, n_wsis_val = wsi_level_qwk(model, val_data, DEVICE)\n        is_best = wsi_kappa > best_kappa\n\n        print(f\"     Ep[{epoch:02d}/{EPOCHS}] \"\n              f\"loss={t_loss/t_total:.4f} acc={t_correct/t_total:.3f} | \"\n              f\"WSI_QWK={wsi_kappa:.4f} ({n_wsis_val} WSIs)\"\n              + (\" *\" if is_best else \"\"))\n\n        if is_best:\n            best_kappa = wsi_kappa\n            no_improve = 0\n            torch.save({\n                \"epoch\":            epoch,\n                \"model_state_dict\": model.state_dict() if not hasattr(model, \"module\")\n                                    else model.module.state_dict(),\n                \"kappa\":            wsi_kappa,\n                \"shard_id\":         SHARD_ID,\n                \"fold_id\":          fold_id,\n                \"n_wsis\":           len(all_wsis),\n                \"hyperparams\": {\n                    \"lr_cnn\":    LR_CNN,\n                    \"lr_tr\":     LR_TR,\n                    \"epochs\":    EPOCHS,\n                    \"patience\":  PATIENCE,\n                }\n            }, best_ckpt)\n        else:\n            no_improve += 1\n            if no_improve >= PATIENCE:\n                print(f\"     Early stopping epoch {epoch}\")\n                break\n\n    print(f\"  OK Fold {fold_id+1} best WSI_QWK: {best_kappa:.4f}\")\n    push_to_kaggle(best_ckpt,\n        f\"Stage{SHARD_ID} Fold{fold_id+1} WSI_QWK={best_kappa:.4f}\")\n    return best_kappa, best_ckpt\n\n# ─────────────────────────────────────────────\n# 10. 5-FOLD CV\n# ─────────────────────────────────────────────\nprint(f\"\\n{'='*60}\")\nprint(f\"Stage {SHARD_ID} -- {len(all_wsis)} WSIs -- 5-Fold CV\")\nprint(f\"Fine-tune depuis Stage 0 (QWK=0.7746)\")\nprint(f\"{'='*60}\")\n\nwsi_ids    = [r[\"wsi_id\"] for r in unique_records]\nlabels_arr = [r[\"label\"]  for r in unique_records]\nindices    = list(range(len(unique_records)))\n\nsgkf = StratifiedGroupKFold(n_splits=N_FOLDS, shuffle=True, random_state=SEED)\n\nfold_kappas       = DONE_KAPPAS.copy()\nfold_ckpts        = DONE_CKPTS.copy()\nbest_overall      = max(DONE_KAPPAS) if DONE_KAPPAS else -1.0\nbest_ckpt_overall = None\n\nfor fold_id, (train_idx, val_idx) in enumerate(\n        sgkf.split(indices, labels_arr, groups=wsi_ids)):\n\n    if fold_id < START_FOLD:\n        continue\n\n    train_data = [unique_records[i] for i in train_idx]\n    val_data   = [unique_records[i] for i in val_idx]\n\n    # Anti-leakage\n    tw = set(unique_records[i][\"wsi_id\"] for i in train_idx)\n    vw = set(unique_records[i][\"wsi_id\"] for i in val_idx)\n    assert len(tw & vw) == 0, f\"DATA LEAKAGE fold {fold_id+1}!\"\n\n    kappa, ckpt_path = train_one_fold(fold_id, train_data, val_data)\n    fold_kappas.append(kappa)\n    fold_ckpts.append(ckpt_path)\n\n    if kappa > best_overall:\n        best_overall      = kappa\n        best_ckpt_overall = ckpt_path\n\n# ─────────────────────────────────────────────\n# 11. RESUME FINAL\n# ─────────────────────────────────────────────\nmean_kappa = np.mean(fold_kappas)\nstd_kappa  = np.std(fold_kappas)\n\nprint(f\"\\n{'='*60}\")\nprint(f\"OK Stage {SHARD_ID} -- 5-Fold CV Results\")\nprint(f\"{'='*60}\")\nfor i, k in enumerate(fold_kappas):\n    print(f\"   Fold {i+1}: QWK = {k:.4f}\")\nprint(f\"{'─'*40}\")\nprint(f\"   Mean QWK : {mean_kappa:.4f} +/- {std_kappa:.4f}\")\nprint(f\"   Best QWK : {best_overall:.4f}\")\nprint(f\"   Stage 0  : 0.6514 +/- 0.1050\")\nprint(f\"   WSIs     : {len(all_wsis)}\")\nprint(f\"{'='*60}\")\n\nresults = {\n    \"stage\":       SHARD_ID,\n    \"n_wsis\":      len(all_wsis),\n    \"n_patches\":   len(unique_records),\n    \"fold_kappas\": fold_kappas,\n    \"mean_kappa\":  mean_kappa,\n    \"std_kappa\":   std_kappa,\n    \"best_kappa\":  best_overall,\n    \"timestamp\":   datetime.datetime.now().isoformat(),\n    \"prev_ckpt\":   PREV_CKPT,\n}\nwith open(RESULTS_JSON, \"w\") as f:\n    json.dump(results, f, indent=2)\n\nif best_ckpt_overall:\n    stage_best = f\"{SAVE_DIR}/stage{SHARD_ID}_best_qwk{best_overall:.4f}.pth\"\n    shutil.copy(best_ckpt_overall, stage_best)\n    print(f\"\\nOK Best checkpoint: {stage_best}\")\n    print(f\"\\n-> Pour Stage 2 :\")\n    print(f\"   SHARD_ID  = 2\")\n    print(f\"   PREV_CKPT = '{stage_best}'\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, glob, shutil, subprocess, json\n\n# Vérifier checkpoint\nckpts = glob.glob(\"/kaggle/working/stage1*.pth\")\nprint(\"Stage 1 checkpoints:\")\nfor c in ckpts:\n    mb = os.path.getsize(c)/(1024*1024)\n    print(f\"  {os.path.basename(c)} ({mb:.1f} MB)\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\n=============================================================\nStage 1 — Incremental 5-Fold CV\nFine-tune depuis Stage 0 (best QWK=0.7746)\nDonnées cumulatives :\n  Shard 0 : patches_400     (399 WSIs,  4,319 patches)\n  Shard 1 : panda-patches-run1 (1,000 WSIs, 14,842 patches)\nTotal     : ~1,399 WSIs, ~19,161 patches\n=============================================================\n\"\"\"\n\nimport os, glob, json, random, shutil, subprocess, datetime\nimport numpy as np\nfrom PIL import Image\nfrom collections import Counter\nfrom sklearn.model_selection import StratifiedGroupKFold\nfrom sklearn.metrics import cohen_kappa_score, accuracy_score\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as T\nimport torchvision.models as models\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\n\n# ─────────────────────────────────────────────\n# 1. CONFIG\n# ─────────────────────────────────────────────\nSHARD_ID    = 1\nSTART_FOLD  = 1    # Fold 2\nDONE_KAPPAS = [0.6659]\nDONE_CKPTS  = [\"/kaggle/working/stage1_fold1_best.pth\"]\n\nPREV_CKPT   = \"/kaggle/working/stage0_best_qwk0.7746.pth\"\n\n\n\n\n\n\n\n\nKAGGLE_DS   = \"khouloudbejaoui20/best-models\"\n\n\n\n\n\nN_FOLDS     = 5\nSEED        = 42\nBATCH_SIZE  = 32\nNUM_WORKERS = 2\nEPOCHS      = 20\nPATIENCE    = 5\nLR_CNN      = 1e-5\nLR_TR       = 2e-5\n\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\ntorch.backends.cudnn.deterministic = True\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Device    : {DEVICE}\")\nprint(f\"Stage     : {SHARD_ID} (fine-tune depuis Stage 0 QWK=0.7746)\")\nprint(f\"EPOCHS={EPOCHS} | PATIENCE={PATIENCE} | LR_CNN={LR_CNN} | LR_TR={LR_TR}\")\n\n# ─────────────────────────────────────────────\n# 2. CHEMINS PATCHES\n# ─────────────────────────────────────────────\nBASE = \"/kaggle/input/datasets/khouloudbejaoui20\"\n\nSHARD_PATHS = {\n    0: f\"{BASE}/dataset/patches_400-20260630T212405Z-3-001/patches_400\",\n    1: f\"{BASE}/panda-patches-run1/kaggle/working/patches_panda_new\",\n}\n\nSAVE_DIR     = \"/kaggle/working\"\nKAGGLE_DS    = \"khouloudbejaoui20/best-models\"\nRESULTS_JSON = f\"{SAVE_DIR}/stage1_cv_results.json\"\n\n# ─────────────────────────────────────────────\n# 3. VERIFICATION\n# ─────────────────────────────────────────────\nprint(f\"\\n── Verification ────────────────────────\")\nassert os.path.exists(PREV_CKPT), f\"PREV_CKPT non trouve : {PREV_CKPT}\"\nprint(f\"  OK Stage 0 checkpoint trouve (QWK=0.7746)\")\n\nfor shard_id in range(SHARD_ID + 1):\n    path = SHARD_PATHS.get(shard_id, \"\")\n    if os.path.exists(path):\n        files = glob.glob(f\"{path}/*.png\")\n        if not files:\n            files = glob.glob(f\"{path}/**/*.png\", recursive=True)\n        print(f\"  OK Shard {shard_id}: {len(files)} patches -> {path}\")\n    else:\n        print(f\"  MANQUANT Shard {shard_id}: {path}\")\n\n# ─────────────────────────────────────────────\n# 4. CHARGEMENT PATCHES\n# ─────────────────────────────────────────────\nprint(f\"\\n── Patches cumulatifs (Shard 0 a {SHARD_ID}) ──\")\n\ndef load_patches(folder):\n    files = glob.glob(f\"{folder}/*.png\")\n    if not files:\n        files = glob.glob(f\"{folder}/**/*.png\", recursive=True)\n    records = []\n    for fp in files:\n        fname = os.path.basename(fp)\n        try:\n            parts  = fname.split(\"_slide\")\n            label  = int(parts[0].replace(\"grade\", \"\"))\n            wsi_id = parts[1].split(\"_\")[0]\n            records.append({\"path\": fp, \"wsi_id\": wsi_id, \"label\": label})\n        except:\n            continue\n    return records\n\nall_records = []\nfor shard_id in range(SHARD_ID + 1):\n    path = SHARD_PATHS.get(shard_id, \"\")\n    if path and os.path.exists(path):\n        r = load_patches(path)\n        all_records.extend(r)\n        print(f\"  Shard {shard_id}: {len(r)} patches\")\n    else:\n        print(f\"  Shard {shard_id}: NON TROUVE\")\n\n# Dedupliquer\nseen = set()\nunique_records = []\nfor r in all_records:\n    fname = os.path.basename(r[\"path\"])\n    if fname not in seen:\n        seen.add(fname)\n        unique_records.append(r)\n\nall_wsis = list(set(r[\"wsi_id\"] for r in unique_records))\nprint(f\"\\n  Total unique : {len(unique_records)} patches\")\nprint(f\"  Total WSIs  : {len(all_wsis)}\")\ncnt = Counter(r[\"label\"] for r in unique_records)\nfor c in range(6):\n    print(f\"    ISUP {c}: {cnt.get(c, 0)}\")\n\nassert len(all_wsis) > 500, f\"Trop peu de WSIs ({len(all_wsis)})\"\nprint(f\"\\n  OK {len(all_wsis)} WSIs uniques confirmes\")\n\n# ─────────────────────────────────────────────\n# 5. TRANSFORMS\n# ─────────────────────────────────────────────\nMEAN = [0.485, 0.456, 0.406]\nSTD  = [0.229, 0.224, 0.225]\n\ntrain_tf = T.Compose([\n    T.Resize((224, 224)),\n    T.RandomHorizontalFlip(),\n    T.RandomVerticalFlip(),\n    T.RandomRotation(90),\n    T.ColorJitter(brightness=0.2, contrast=0.2,\n                  saturation=0.1, hue=0.05),\n    T.RandomApply([T.GaussianBlur(3)], p=0.2),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n])\n\nval_tf = T.Compose([\n    T.Resize((224, 224)),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n])\n\nclass PatchDataset(Dataset):\n    def __init__(self, records, transform):\n        self.records   = records\n        self.transform = transform\n    def __len__(self):\n        return len(self.records)\n    def __getitem__(self, idx):\n        r = self.records[idx]\n        try:\n            img = Image.open(r[\"path\"]).convert(\"RGB\")\n        except:\n            img = Image.new(\"RGB\", (224, 224), (255, 255, 255))\n        return self.transform(img), int(r[\"label\"])\n\n# ─────────────────────────────────────────────\n# 6. MODELE\n# ─────────────────────────────────────────────\nclass HybridCNNTransformer(nn.Module):\n    def __init__(self, num_classes=6, d_model=512, nhead=8,\n                 num_layers=2, dim_feedforward=2048, dropout=0.1):\n        super().__init__()\n        backbone = models.resnet50(weights=models.ResNet50_Weights.DEFAULT)\n        self.cnn_backbone = nn.Sequential(*list(backbone.children())[:-2])\n        self.pool         = nn.AdaptiveAvgPool2d((1, 1))\n        self.proj         = nn.Sequential(\n            nn.Identity(), nn.Identity(),\n            nn.Linear(2048, d_model),\n            nn.LayerNorm(d_model)\n        )\n        self.cls_token = nn.Parameter(torch.zeros(1, 1, d_model))\n        self.pos_embed = nn.Parameter(torch.zeros(1, 2, d_model))\n        encoder_layer  = nn.TransformerEncoderLayer(\n            d_model=d_model, nhead=nhead,\n            dim_feedforward=dim_feedforward,\n            dropout=dropout, batch_first=True\n        )\n        self.transformer = nn.TransformerEncoder(\n            encoder_layer, num_layers=num_layers,\n            norm=nn.LayerNorm(d_model)\n        )\n        self.classifier = nn.Sequential(\n            nn.Linear(d_model, 256), nn.ReLU(),\n            nn.Dropout(dropout), nn.Linear(256, num_classes)\n        )\n        nn.init.trunc_normal_(self.cls_token, std=0.02)\n        nn.init.trunc_normal_(self.pos_embed, std=0.02)\n\n    def forward(self, x):\n        f   = self.pool(self.cnn_backbone(x)).flatten(1)\n        f   = self.proj(f).unsqueeze(1)\n        cls = self.cls_token.expand(x.size(0), -1, -1)\n        seq = torch.cat([cls, f], dim=1) + self.pos_embed\n        return self.classifier(self.transformer(seq)[:, 0])\n\ndef build_model(prev_ckpt=None):\n    model  = HybridCNNTransformer().to(DEVICE)\n    n_gpus = torch.cuda.device_count()\n    if n_gpus > 1:\n        model = nn.DataParallel(model)\n        print(f\"  OK DataParallel: {n_gpus} GPUs\")\n    n_param = sum(p.numel() for p in model.parameters())\n    assert n_param == 30_998_342, f\"Mismatch: {n_param}\"\n    print(f\"  OK Parameters: {n_param:,}\")\n    if prev_ckpt and os.path.exists(prev_ckpt):\n        ckpt  = torch.load(prev_ckpt, map_location=DEVICE, weights_only=False)\n        state = ckpt[\"model_state_dict\"]\n        if n_gpus > 1 and not list(state.keys())[0].startswith(\"module.\"):\n            state = {\"module.\" + k: v for k, v in state.items()}\n        elif n_gpus <= 1 and list(state.keys())[0].startswith(\"module.\"):\n            state = {k[7:]: v for k, v in state.items()}\n        model.load_state_dict(state)\n        print(f\"  OK Loaded Stage 0 (QWK={ckpt.get('kappa', '?')})\")\n    return model\n\n# ─────────────────────────────────────────────\n# 7. WSI-LEVEL QWK\n# ─────────────────────────────────────────────\ndef wsi_level_qwk(model, val_data, device):\n    model.eval()\n    wsi_preds  = {}\n    wsi_labels = {}\n    all_preds  = []\n    all_labels = []\n    loader = DataLoader(\n        PatchDataset(val_data, val_tf),\n        batch_size=BATCH_SIZE, shuffle=False,\n        num_workers=NUM_WORKERS\n    )\n    with torch.no_grad():\n        for imgs, labels in loader:\n            imgs  = imgs.to(device)\n            preds = model(imgs).argmax(1).cpu().numpy()\n            all_preds.extend(preds)\n            all_labels.extend(labels.numpy())\n    for i, r in enumerate(val_data):\n        wid = r[\"wsi_id\"]\n        if wid not in wsi_preds:\n            wsi_preds[wid]  = []\n            wsi_labels[wid] = r[\"label\"]\n        wsi_preds[wid].append(all_preds[i])\n    wf = []; wl = []\n    for wid in wsi_preds:\n        votes = wsi_preds[wid]\n        wf.append(max(set(votes), key=votes.count))\n        wl.append(wsi_labels[wid])\n    patch_qwk = cohen_kappa_score(all_labels, all_preds, weights=\"quadratic\")\n    wsi_qwk   = cohen_kappa_score(wl, wf, weights=\"quadratic\")\n    return patch_qwk, wsi_qwk, len(wf)\n\n# ─────────────────────────────────────────────\n# 8. PUSH KAGGLE\n# ─────────────────────────────────────────────\ndef push_to_kaggle(ckpt_path, message):\n    try:\n        tmp = f\"{SAVE_DIR}/_push_tmp\"\n        os.makedirs(tmp, exist_ok=True)\n        for c in glob.glob(f\"{SAVE_DIR}/stage{SHARD_ID}_fold*.pth\"):\n            shutil.copy(c, tmp)\n        shutil.copy(ckpt_path, tmp)\n        if os.path.exists(RESULTS_JSON):\n            shutil.copy(RESULTS_JSON, tmp)\n        json.dump({\n            \"title\": \"Memoire IHC Saved Models\",\n            \"id\": KAGGLE_DS,\n            \"licenses\": [{\"name\": \"CC0-1.0\"}]\n        }, open(f\"{tmp}/dataset-metadata.json\", \"w\"))\n        files_pushed = [os.path.basename(f)\n                        for f in glob.glob(f\"{tmp}/*.pth\")]\n        print(f\"   Pushing: {files_pushed}\")\n        r = subprocess.run(\n            [\"kaggle\", \"datasets\", \"version\",\n             \"-p\", tmp, \"-m\", message,\n             \"--dir-mode\", \"skip\"],\n            capture_output=True, text=True, timeout=300\n        )\n        print(f\"   {'OK Kaggle' if r.returncode==0 else 'ERREUR '+r.stderr[:100]}\")\n    except Exception as e:\n        print(f\"   ERREUR push: {e}\")\n\n# ─────────────────────────────────────────────\n# 9. TRAIN ONE FOLD\n# ─────────────────────────────────────────────\ndef train_one_fold(fold_id, train_data, val_data):\n    print(f\"\\n  Fold {fold_id+1}/{N_FOLDS} ──────────────────────────\")\n    print(f\"     Train: {len(train_data)} | Val: {len(val_data)}\")\n\n    train_loader = DataLoader(\n        PatchDataset(train_data, train_tf),\n        batch_size=BATCH_SIZE, shuffle=True,\n        num_workers=NUM_WORKERS, pin_memory=True\n    )\n\n    model = build_model(PREV_CKPT)\n\n    cnt_tr = Counter(r[\"label\"] for r in train_data)\n    w = torch.tensor(\n        [1.0 / max(cnt_tr.get(c, 1), 1) for c in range(6)],\n        dtype=torch.float\n    )\n    w = (w / w.sum() * 6).to(DEVICE)\n    criterion = nn.CrossEntropyLoss(weight=w)\n\n    base_model = model.module if hasattr(model, \"module\") else model\n    optimizer  = optim.AdamW([\n        {\"params\": base_model.cnn_backbone.parameters(), \"lr\": LR_CNN},\n        {\"params\": [p for n, p in base_model.named_parameters()\n                    if not n.startswith(\"cnn_backbone\")], \"lr\": LR_TR},\n    ], weight_decay=1e-4)\n\n    scheduler  = CosineAnnealingLR(optimizer, T_max=EPOCHS, eta_min=1e-7)\n\n    best_kappa = 0.0\n    no_improve = 0\n    best_ckpt  = f\"{SAVE_DIR}/stage{SHARD_ID}_fold{fold_id+1}_best.pth\"\n\n    for epoch in range(1, EPOCHS + 1):\n        model.train()\n        t_loss, t_correct, t_total = 0.0, 0, 0\n        for imgs, labels in train_loader:\n            imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n            optimizer.zero_grad()\n            out  = model(imgs)\n            loss = criterion(out, labels)\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n            optimizer.step()\n            t_loss    += loss.item() * imgs.size(0)\n            t_correct += (out.argmax(1) == labels).sum().item()\n            t_total   += imgs.size(0)\n        scheduler.step()\n\n        _, wsi_kappa, n_wsis_val = wsi_level_qwk(model, val_data, DEVICE)\n        is_best = wsi_kappa > best_kappa\n\n        print(f\"     Ep[{epoch:02d}/{EPOCHS}] \"\n              f\"loss={t_loss/t_total:.4f} acc={t_correct/t_total:.3f} | \"\n              f\"WSI_QWK={wsi_kappa:.4f} ({n_wsis_val} WSIs)\"\n              + (\" *\" if is_best else \"\"))\n\n        if is_best:\n            best_kappa = wsi_kappa\n            no_improve = 0\n            torch.save({\n                \"epoch\":            epoch,\n                \"model_state_dict\": model.state_dict() if not hasattr(model, \"module\")\n                                    else model.module.state_dict(),\n                \"kappa\":            wsi_kappa,\n                \"shard_id\":         SHARD_ID,\n                \"fold_id\":          fold_id,\n                \"n_wsis\":           len(all_wsis),\n                \"hyperparams\": {\n                    \"lr_cnn\":    LR_CNN,\n                    \"lr_tr\":     LR_TR,\n                    \"epochs\":    EPOCHS,\n                    \"patience\":  PATIENCE,\n                }\n            }, best_ckpt)\n        else:\n            no_improve += 1\n            if no_improve >= PATIENCE:\n                print(f\"     Early stopping epoch {epoch}\")\n                break\n\n    print(f\"  OK Fold {fold_id+1} best WSI_QWK: {best_kappa:.4f}\")\n    push_to_kaggle(best_ckpt,\n        f\"Stage{SHARD_ID} Fold{fold_id+1} WSI_QWK={best_kappa:.4f}\")\n    return best_kappa, best_ckpt\n\n# ─────────────────────────────────────────────\n# 10. 5-FOLD CV\n# ─────────────────────────────────────────────\nprint(f\"\\n{'='*60}\")\nprint(f\"Stage {SHARD_ID} -- {len(all_wsis)} WSIs -- 5-Fold CV\")\nprint(f\"Fine-tune depuis Stage 0 (QWK=0.7746)\")\nprint(f\"{'='*60}\")\n\nwsi_ids    = [r[\"wsi_id\"] for r in unique_records]\nlabels_arr = [r[\"label\"]  for r in unique_records]\nindices    = list(range(len(unique_records)))\n\nsgkf = StratifiedGroupKFold(n_splits=N_FOLDS, shuffle=True, random_state=SEED)\n\nfold_kappas       = DONE_KAPPAS.copy()\nfold_ckpts        = DONE_CKPTS.copy()\nbest_overall      = max(DONE_KAPPAS) if DONE_KAPPAS else -1.0\nbest_ckpt_overall = None\n\nfor fold_id, (train_idx, val_idx) in enumerate(\n        sgkf.split(indices, labels_arr, groups=wsi_ids)):\n\n    if fold_id < START_FOLD:\n        continue\n\n    train_data = [unique_records[i] for i in train_idx]\n    val_data   = [unique_records[i] for i in val_idx]\n\n    # Anti-leakage\n    tw = set(unique_records[i][\"wsi_id\"] for i in train_idx)\n    vw = set(unique_records[i][\"wsi_id\"] for i in val_idx)\n    assert len(tw & vw) == 0, f\"DATA LEAKAGE fold {fold_id+1}!\"\n\n    kappa, ckpt_path = train_one_fold(fold_id, train_data, val_data)\n    fold_kappas.append(kappa)\n    fold_ckpts.append(ckpt_path)\n\n    if kappa > best_overall:\n        best_overall      = kappa\n        best_ckpt_overall = ckpt_path\n\n# ─────────────────────────────────────────────\n# 11. RESUME FINAL\n# ─────────────────────────────────────────────\nmean_kappa = np.mean(fold_kappas)\nstd_kappa  = np.std(fold_kappas)\n\nprint(f\"\\n{'='*60}\")\nprint(f\"OK Stage {SHARD_ID} -- 5-Fold CV Results\")\nprint(f\"{'='*60}\")\nfor i, k in enumerate(fold_kappas):\n    print(f\"   Fold {i+1}: QWK = {k:.4f}\")\nprint(f\"{'─'*40}\")\nprint(f\"   Mean QWK : {mean_kappa:.4f} +/- {std_kappa:.4f}\")\nprint(f\"   Best QWK : {best_overall:.4f}\")\nprint(f\"   Stage 0  : 0.6514 +/- 0.1050\")\nprint(f\"   WSIs     : {len(all_wsis)}\")\nprint(f\"{'='*60}\")\n\nresults = {\n    \"stage\":       SHARD_ID,\n    \"n_wsis\":      len(all_wsis),\n    \"n_patches\":   len(unique_records),\n    \"fold_kappas\": fold_kappas,\n    \"mean_kappa\":  mean_kappa,\n    \"std_kappa\":   std_kappa,\n    \"best_kappa\":  best_overall,\n    \"timestamp\":   datetime.datetime.now().isoformat(),\n    \"prev_ckpt\":   PREV_CKPT,\n}\nwith open(RESULTS_JSON, \"w\") as f:\n    json.dump(results, f, indent=2)\n\nif best_ckpt_overall:\n    stage_best = f\"{SAVE_DIR}/stage{SHARD_ID}_best_qwk{best_overall:.4f}.pth\"\n    shutil.copy(best_ckpt_overall, stage_best)\n    print(f\"\\nOK Best checkpoint: {stage_best}\")\n    print(f\"\\n-> Pour Stage 2 :\")\n    print(f\"   SHARD_ID  = 2\")\n    print(f\"   PREV_CKPT = '{stage_best}'\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nprint(os.path.exists(\"/kaggle/working/stage0_best_qwk0.7746.pth\"))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import subprocess\nr = subprocess.run([\n    \"kaggle\", \"datasets\", \"files\",\n    \"khouloudbejaoui20/best-models\"\n], capture_output=True, text=True)\nprint(r.stdout)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import subprocess\nr = subprocess.run([\n    \"kaggle\", \"datasets\", \"files\",\n    \"khouloudbejaoui20/panda-patches-run1\"\n], capture_output=True, text=True)\nprint(r.stdout)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import subprocess, os\n\nsubprocess.run([\n    \"kaggle\", \"datasets\", \"download\",\n    \"khouloudbejaoui20/best-models\",\n    \"--file\", \"stage3_fold5_best.pth\",\n    \"-p\", \"/kaggle/working\", \"--unzip\"\n], capture_output=True)\n\nprint(os.path.exists(\"/kaggle/working/stage3_fold5_best.pth\"))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"SHARD_ID    = 4\nSTART_FOLD  = 4    # Fold 5\n\nDONE_KAPPAS = [0.7220, 0.7151, 0.7030, 0.6904]\nDONE_CKPTS  = [\n    \"/kaggle/working/stage4_fold1_best.pth\",\n    \"/kaggle/working/stage4_fold2_best.pth\",\n    \"/kaggle/working/stage3_fold5_best.pth\",\n    \"/kaggle/working/stage4_fold4_best.pth\",\n]\n\nPREV_CKPT = \"/kaggle/working/stage3_fold5_best.pth\"\nKAGGLE_DS = \"khouloudbejaoui20/best-models\"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\n=============================================================\nStage 4 — Incremental 5-Fold CV — SCRIPT FINAL\nFine-tune depuis Stage 3 (best QWK=0.7845)\nDonnées cumulatives :\n  Shard 0 : patches_400              (399 WSIs)\n  Shard 1 : panda-patches-run1       (1,000 WSIs)\n  Shard 2 : panda-patches-run2       (999 WSIs)\n  Shard 3 : panda-patches-run3-correct (999 WSIs)\n  Shard 4 : panda-patches-run4       (1,000 WSIs)\nTotal     : ~4,397 WSIs, ~63,671 patches\nAméliorations :\n  - Label Smoothing = 0.1\n  - ColorJitter fort (0.3, 0.3, 0.2, 0.1)\n  - GaussianBlur kernel 5x5 (p=0.3)\n  - RandomErasing (p=0.2, scale 0.02-0.20)\n  - CosineAnnealingLR (OneCycleLR INTERDIT)\n  - Push TOUS les folds ensemble apres chaque fold\n  - Glob corrige: single-level first, recursive fallback\n  - Assert anti-leakage a chaque fold\nTarget    : WSI_QWK > 0.7845\n=============================================================\n\"\"\"\n\nimport os, glob, json, random, shutil, subprocess, datetime\nimport numpy as np\nfrom PIL import Image\nfrom collections import Counter\nfrom sklearn.model_selection import StratifiedGroupKFold\nfrom sklearn.metrics import cohen_kappa_score, accuracy_score\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as T\nimport torchvision.models as models\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\n\n# ─────────────────────────────────────────────\n# 1. CONFIG\n# ─────────────────────────────────────────────\nSHARD_ID    = 4\nSTART_FOLD  = 4    # Fold 5\nDONE_KAPPAS = [0.7220, 0.7151, 0.7030, 0.6904]\nDONE_CKPTS  = [\n    \"/kaggle/working/stage4_fold1_best.pth\",\n    \"/kaggle/working/stage4_fold2_best.pth\",\n    \"/kaggle/working/stage3_fold5_best.pth\",\n    \"/kaggle/working/stage4_fold4_best.pth\",\n]\n#PREV_CKPT = \"/kaggle/working/stage3_best_qwk0.7845.pth\"\n\n\n\n\n\n\n\n\n\nPREV_CKPT = \"/kaggle/working/stage3_fold5_best.pth\"\nKAGGLE_DS = \"khouloudbejaoui20/best-models\"\n\n\n\nN_FOLDS     = 5\nSEED        = 42\nBATCH_SIZE  = 32\nNUM_WORKERS = 2\nEPOCHS      = 30\nPATIENCE    = 7\nLR_CNN      = 2e-5\nLR_TR       = 4e-5\n\nrandom.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)\ntorch.backends.cudnn.deterministic = True\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Device    : {DEVICE}\")\nprint(f\"Stage     : {SHARD_ID} (fine-tune depuis Stage 3 QWK=0.7845)\")\nprint(f\"EPOCHS={EPOCHS} | PATIENCE={PATIENCE} | LR_CNN={LR_CNN} | LR_TR={LR_TR}\")\nprint(f\"Ameliorations : LabelSmoothing=0.1 | RandomErasing | ColorJitter+ | CosineAnnealingLR\")\n\n# ─────────────────────────────────────────────\n# 2. CHEMINS PATCHES\n# ─────────────────────────────────────────────\nBASE = \"/kaggle/input/datasets/khouloudbejaoui20\"\n\nSHARD_PATHS = {\n    0: f\"{BASE}/dataset/patches_400-20260630T212405Z-3-001/patches_400\",\n    1: f\"{BASE}/panda-patches-run1/kaggle/working/patches_panda_new\",\n    2: f\"{BASE}/panda-patches-run2/kaggle/working/patches_panda_run2\",\n    3: f\"{BASE}/panda-patches-run3/kaggle/working/patches_panda_run3_correct\",\n    4: f\"{BASE}/panda-patches-run4\",\n}\n\nSAVE_DIR     = \"/kaggle/working\"\nKAGGLE_DS    = \"khouloudbejaoui20/best_models\"\nRESULTS_JSON = f\"{SAVE_DIR}/incremental_cv_results_stage{SHARD_ID}.json\"\n\n# ─────────────────────────────────────────────\n# 3. VERIFICATION\n# ─────────────────────────────────────────────\nprint(f\"\\n── Verification ────────────────────────\")\nassert os.path.exists(PREV_CKPT), f\"PREV_CKPT non trouve : {PREV_CKPT}\"\nprint(f\"  OK Stage 3 checkpoint trouve (QWK=0.7845)\")\n\nfor shard_id in range(SHARD_ID + 1):\n    path = SHARD_PATHS.get(shard_id, \"\")\n    exists = os.path.exists(path)\n    if exists:\n        files = glob.glob(f\"{path}/*.png\")\n        if not files:\n            files = glob.glob(f\"{path}/**/*.png\", recursive=True)\n        n = len(files)\n    else:\n        n = 0\n    status = \"OK\" if (exists and n > 0) else \"MANQUANT\"\n    print(f\"  {status} Shard {shard_id}: {n} patches -> {path}\")\n\n# ─────────────────────────────────────────────\n# 4. CHARGEMENT PATCHES CORRIGES\n# ─────────────────────────────────────────────\nprint(f\"\\n── Patches cumulatifs (Shard 0 a {SHARD_ID}) ──\")\n\ndef load_patches(folder):\n    \"\"\"Charge patches sans duplication (single-level first, recursive fallback).\"\"\"\n    files = glob.glob(f\"{folder}/*.png\")\n    if not files:\n        files = glob.glob(f\"{folder}/**/*.png\", recursive=True)\n    records = []\n    for fp in files:\n        fname = os.path.basename(fp)\n        try:\n            parts  = fname.split(\"_slide\")\n            label  = int(parts[0].replace(\"grade\", \"\"))\n            wsi_id = parts[1].split(\"_\")[0]\n            records.append({\"path\": fp, \"wsi_id\": wsi_id, \"label\": label})\n        except:\n            continue\n    return records\n\nall_records = []\nfor shard_id in range(SHARD_ID + 1):\n    path = SHARD_PATHS.get(shard_id, \"\")\n    if path and os.path.exists(path):\n        r = load_patches(path)\n        all_records.extend(r)\n        print(f\"  Shard {shard_id}: {len(r)} patches\")\n    else:\n        print(f\"  Shard {shard_id}: NON TROUVE\")\n\n# Dedupliquer par filename\nseen = set()\nunique_records = []\nfor r in all_records:\n    fname = os.path.basename(r[\"path\"])\n    if fname not in seen:\n        seen.add(fname)\n        unique_records.append(r)\n\nall_wsis = list(set(r[\"wsi_id\"] for r in unique_records))\nprint(f\"\\n  Total unique : {len(unique_records)} patches\")\nprint(f\"  Total WSIs  : {len(all_wsis)}\")\ncnt = Counter(r[\"label\"] for r in unique_records)\nfor c in range(6):\n    print(f\"    ISUP {c}: {cnt.get(c, 0)}\")\n\nassert len(all_wsis) > 3500, \\\n    f\"Trop peu de WSIs ({len(all_wsis)}) — verifier SHARD_PATHS[4]\"\nprint(f\"\\n  OK {len(all_wsis)} WSIs uniques confirmes\")\n\n# ─────────────────────────────────────────────\n# 5. TRANSFORMS AMELIORES\n# ─────────────────────────────────────────────\nMEAN = [0.485, 0.456, 0.406]\nSTD  = [0.229, 0.224, 0.225]\n\ntrain_tf = T.Compose([\n    T.Resize((224, 224)),\n    T.RandomHorizontalFlip(),\n    T.RandomVerticalFlip(),\n    T.RandomRotation(90),\n    T.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.2, hue=0.1),\n    T.RandomApply([T.GaussianBlur(5)], p=0.3),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n    T.RandomApply([T.RandomErasing(p=0.5, scale=(0.02, 0.2))], p=0.2),\n])\n\nval_tf = T.Compose([\n    T.Resize((224, 224)),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n])\n\nclass PatchDataset(Dataset):\n    def __init__(self, records, transform):\n        self.records   = records\n        self.transform = transform\n    def __len__(self):\n        return len(self.records)\n    def __getitem__(self, idx):\n        r = self.records[idx]\n        try:\n            img = Image.open(r[\"path\"]).convert(\"RGB\")\n        except:\n            img = Image.new(\"RGB\", (224, 224), (255, 255, 255))\n        return self.transform(img), int(r[\"label\"])\n\n# ─────────────────────────────────────────────\n# 6. MODELE\n# ─────────────────────────────────────────────\nclass HybridCNNTransformer(nn.Module):\n    def __init__(self, num_classes=6, d_model=512, nhead=8,\n                 num_layers=2, dim_feedforward=2048, dropout=0.1):\n        super().__init__()\n        backbone = models.resnet50(weights=models.ResNet50_Weights.DEFAULT)\n        self.cnn_backbone = nn.Sequential(*list(backbone.children())[:-2])\n        self.pool         = nn.AdaptiveAvgPool2d((1, 1))\n        self.proj         = nn.Sequential(\n            nn.Identity(), nn.Identity(),\n            nn.Linear(2048, d_model),\n            nn.LayerNorm(d_model)\n        )\n        self.cls_token = nn.Parameter(torch.zeros(1, 1, d_model))\n        self.pos_embed = nn.Parameter(torch.zeros(1, 2, d_model))\n        encoder_layer  = nn.TransformerEncoderLayer(\n            d_model=d_model, nhead=nhead,\n            dim_feedforward=dim_feedforward,\n            dropout=dropout, batch_first=True\n        )\n        self.transformer = nn.TransformerEncoder(\n            encoder_layer, num_layers=num_layers,\n            norm=nn.LayerNorm(d_model)\n        )\n        self.classifier = nn.Sequential(\n            nn.Linear(d_model, 256), nn.ReLU(),\n            nn.Dropout(dropout), nn.Linear(256, num_classes)\n        )\n        nn.init.trunc_normal_(self.cls_token, std=0.02)\n        nn.init.trunc_normal_(self.pos_embed, std=0.02)\n\n    def forward(self, x):\n        f   = self.pool(self.cnn_backbone(x)).flatten(1)\n        f   = self.proj(f).unsqueeze(1)\n        cls = self.cls_token.expand(x.size(0), -1, -1)\n        seq = torch.cat([cls, f], dim=1) + self.pos_embed\n        return self.classifier(self.transformer(seq)[:, 0])\n\ndef build_model(prev_ckpt=None):\n    model  = HybridCNNTransformer().to(DEVICE)\n    n_gpus = torch.cuda.device_count()\n    if n_gpus > 1:\n        model = nn.DataParallel(model)\n        print(f\"  OK DataParallel: {n_gpus} GPUs\")\n    n_param = sum(p.numel() for p in model.parameters())\n    assert n_param == 30_998_342, f\"Mismatch: {n_param}\"\n    print(f\"  OK Parameters: {n_param:,}\")\n    if prev_ckpt and os.path.exists(prev_ckpt):\n        ckpt  = torch.load(prev_ckpt, map_location=DEVICE, weights_only=False)\n        state = ckpt[\"model_state_dict\"]\n        if n_gpus > 1 and not list(state.keys())[0].startswith(\"module.\"):\n            state = {\"module.\" + k: v for k, v in state.items()}\n        elif n_gpus <= 1 and list(state.keys())[0].startswith(\"module.\"):\n            state = {k[7:]: v for k, v in state.items()}\n        model.load_state_dict(state)\n        print(f\"  OK Loaded Stage 3 (QWK={ckpt.get('kappa', '?')})\")\n    return model\n\n# ─────────────────────────────────────────────\n# 7. WSI-LEVEL QWK\n# ─────────────────────────────────────────────\ndef wsi_level_qwk(model, val_data, device):\n    model.eval()\n    wsi_preds  = {}\n    wsi_labels = {}\n    all_preds  = []\n    all_labels = []\n    loader = DataLoader(\n        PatchDataset(val_data, val_tf),\n        batch_size=BATCH_SIZE, shuffle=False,\n        num_workers=NUM_WORKERS\n    )\n    with torch.no_grad():\n        for imgs, labels in loader:\n            imgs  = imgs.to(device)\n            preds = model(imgs).argmax(1).cpu().numpy()\n            all_preds.extend(preds)\n            all_labels.extend(labels.numpy())\n    for i, r in enumerate(val_data):\n        wid = r[\"wsi_id\"]\n        if wid not in wsi_preds:\n            wsi_preds[wid]  = []\n            wsi_labels[wid] = r[\"label\"]\n        wsi_preds[wid].append(all_preds[i])\n    wf = []; wl = []\n    for wid in wsi_preds:\n        votes = wsi_preds[wid]\n        wf.append(max(set(votes), key=votes.count))\n        wl.append(wsi_labels[wid])\n    patch_qwk = cohen_kappa_score(all_labels, all_preds, weights=\"quadratic\")\n    wsi_qwk   = cohen_kappa_score(wl, wf, weights=\"quadratic\")\n    return patch_qwk, wsi_qwk, len(wf)\n\n# ─────────────────────────────────────────────\n# 8. PUSH KAGGLE — TOUS LES FOLDS ENSEMBLE\n# ─────────────────────────────────────────────\ndef push_to_kaggle(ckpt_path, message):\n    try:\n        tmp = f\"{SAVE_DIR}/_push_tmp\"\n        os.makedirs(tmp, exist_ok=True)\n\n        # Copier TOUS les checkpoints stage4 existants\n        for c in glob.glob(f\"{SAVE_DIR}/stage{SHARD_ID}_fold*.pth\"):\n            shutil.copy(c, tmp)\n\n        # Copier le checkpoint actuel\n        shutil.copy(ckpt_path, tmp)\n\n        # Copier JSON resultats si existe\n        if os.path.exists(RESULTS_JSON):\n            shutil.copy(RESULTS_JSON, tmp)\n\n        json.dump({\n            \"title\": \"Memoire IHC Saved Models\",\n            \"id\": KAGGLE_DS,\n            \"licenses\": [{\"name\": \"CC0-1.0\"}]\n        }, open(f\"{tmp}/dataset-metadata.json\", \"w\"))\n\n        files_pushed = [os.path.basename(f) for f in glob.glob(f\"{tmp}/*.pth\")]\n        print(f\"   Pushing ALL : {files_pushed}\")\n\n        r = subprocess.run(\n            [\"kaggle\", \"datasets\", \"version\",\n             \"-p\", tmp, \"-m\", message, \"--dir-mode\", \"skip\"],\n            capture_output=True, text=True, timeout=300\n        )\n        print(f\"   {'OK Kaggle' if r.returncode == 0 else 'ERREUR ' + r.stderr[:100]}\")\n    except Exception as e:\n        print(f\"   ERREUR push: {e}\")\n\n# ─────────────────────────────────────────────\n# 9. TRAIN ONE FOLD\n# ─────────────────────────────────────────────\ndef train_one_fold(fold_id, train_data, val_data):\n    print(f\"\\n  Fold {fold_id+1}/{N_FOLDS} ──────────────────────────\")\n    print(f\"     Train: {len(train_data)} | Val: {len(val_data)}\")\n\n    train_loader = DataLoader(\n        PatchDataset(train_data, train_tf),\n        batch_size=BATCH_SIZE, shuffle=True,\n        num_workers=NUM_WORKERS, pin_memory=True\n    )\n    val_loader = DataLoader(\n        PatchDataset(val_data, val_tf),\n        batch_size=BATCH_SIZE, shuffle=False,\n        num_workers=NUM_WORKERS, pin_memory=True\n    )\n\n    model = build_model(PREV_CKPT)\n\n    # Class weights inversement proportionnels a la frequence\n    cnt_tr = Counter(r[\"label\"] for r in train_data)\n    w = torch.tensor(\n        [1.0 / max(cnt_tr.get(c, 1), 1) for c in range(6)],\n        dtype=torch.float\n    )\n    w = (w / w.sum() * 6).to(DEVICE)\n\n    # Label Smoothing = 0.1\n    criterion = nn.CrossEntropyLoss(weight=w, label_smoothing=0.1)\n\n    # Differential learning rates\n    base_model = model.module if hasattr(model, \"module\") else model\n    optimizer  = optim.AdamW([\n        {\"params\": base_model.cnn_backbone.parameters(), \"lr\": LR_CNN},\n        {\"params\": [p for n, p in base_model.named_parameters()\n                    if not n.startswith(\"cnn_backbone\")], \"lr\": LR_TR},\n    ], weight_decay=1e-4)\n\n    # CosineAnnealingLR (OneCycleLR cause divergence dans ce modele)\n    scheduler = CosineAnnealingLR(optimizer, T_max=EPOCHS, eta_min=1e-7)\n\n    best_kappa = 0.0\n    no_improve = 0\n    best_ckpt  = f\"{SAVE_DIR}/stage{SHARD_ID}_fold{fold_id+1}_best.pth\"\n\n    for epoch in range(1, EPOCHS + 1):\n        # Training\n        model.train()\n        t_loss, t_correct, t_total = 0.0, 0, 0\n        for imgs, labels in train_loader:\n            imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n            optimizer.zero_grad()\n            out  = model(imgs)\n            loss = criterion(out, labels)\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n            optimizer.step()\n            t_loss    += loss.item() * imgs.size(0)\n            t_correct += (out.argmax(1) == labels).sum().item()\n            t_total   += imgs.size(0)\n        # scheduler.step() APRES epoch (CosineAnnealingLR)\n        scheduler.step()\n\n        # Validation patch-level\n        model.eval()\n        preds_all, labels_all = [], []\n        with torch.no_grad():\n            for imgs, labels in val_loader:\n                imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n                out = model(imgs)\n                preds_all.extend(out.argmax(1).cpu().numpy())\n                labels_all.extend(labels.cpu().numpy())\n\n        val_acc     = accuracy_score(labels_all, preds_all)\n        patch_kappa = cohen_kappa_score(labels_all, preds_all, weights=\"quadratic\")\n\n        # Validation WSI-level (metrique principale)\n        _, wsi_kappa, n_wsis_val = wsi_level_qwk(model, val_data, DEVICE)\n        is_best = wsi_kappa > best_kappa\n\n        print(f\"     Ep[{epoch:02d}/{EPOCHS}] \"\n              f\"loss={t_loss/t_total:.4f} acc={t_correct/t_total:.3f} | \"\n              f\"patch_QWK={patch_kappa:.4f} \"\n              f\"WSI_QWK={wsi_kappa:.4f} ({n_wsis_val} WSIs)\"\n              + (\" *\" if is_best else \"\"))\n\n        if is_best:\n            best_kappa = wsi_kappa\n            no_improve = 0\n            torch.save({\n                \"epoch\":            epoch,\n                \"model_state_dict\": model.state_dict() if not hasattr(model, \"module\")\n                                    else model.module.state_dict(),\n                \"kappa\":            wsi_kappa,\n                \"patch_kappa\":      patch_kappa,\n                \"val_acc\":          val_acc,\n                \"shard_id\":         SHARD_ID,\n                \"fold_id\":          fold_id,\n                \"n_wsis\":           len(all_wsis),\n                \"hyperparams\": {\n                    \"lr_cnn\": LR_CNN, \"lr_tr\": LR_TR,\n                    \"epochs\": EPOCHS, \"patience\": PATIENCE,\n                    \"label_smoothing\": 0.1,\n                    \"scheduler\": \"CosineAnnealingLR\"\n                }\n            }, best_ckpt)\n        else:\n            no_improve += 1\n            if no_improve >= PATIENCE:\n                print(f\"     Early stopping epoch {epoch}\")\n                break\n\n    print(f\"  OK Fold {fold_id+1} best WSI_QWK : {best_kappa:.4f}\")\n    push_to_kaggle(best_ckpt,\n        f\"Stage{SHARD_ID} Fold{fold_id+1} WSI_QWK={best_kappa:.4f} [ALL folds]\")\n    return best_kappa, best_ckpt\n\n# ─────────────────────────────────────────────\n# 10. 5-FOLD CROSS-VALIDATION\n# ─────────────────────────────────────────────\nprint(f\"\\n{'='*60}\")\nprint(f\"Stage {SHARD_ID} -- {len(all_wsis)} WSIs -- 5-Fold CV\")\nprint(f\"Fine-tune depuis Stage 3 (QWK=0.7845)\")\nprint(f\"{'='*60}\")\n\nif DONE_KAPPAS:\n    print(f\"Folds deja completes : {DONE_KAPPAS}\")\n\nwsi_ids    = [r[\"wsi_id\"] for r in unique_records]\nlabels_arr = [r[\"label\"]  for r in unique_records]\nindices    = list(range(len(unique_records)))\n\nsgkf = StratifiedGroupKFold(n_splits=N_FOLDS, shuffle=True, random_state=SEED)\n\nfold_kappas       = DONE_KAPPAS.copy()\nfold_ckpts        = DONE_CKPTS.copy()\nbest_overall      = max(DONE_KAPPAS) if DONE_KAPPAS else -1.0\nbest_ckpt_overall = DONE_CKPTS[DONE_KAPPAS.index(max(DONE_KAPPAS))] \\\n                    if DONE_KAPPAS else None\n\nfor fold_id, (train_idx, val_idx) in enumerate(\n        sgkf.split(indices, labels_arr, groups=wsi_ids)):\n\n    if fold_id < START_FOLD:\n        continue\n\n    train_data = [unique_records[i] for i in train_idx]\n    val_data   = [unique_records[i] for i in val_idx]\n\n    # Anti-leakage assertion\n    tw = set(unique_records[i][\"wsi_id\"] for i in train_idx)\n    vw = set(unique_records[i][\"wsi_id\"] for i in val_idx)\n    assert len(tw & vw) == 0, f\"DATA LEAKAGE fold {fold_id+1} !\"\n\n    kappa, ckpt_path = train_one_fold(fold_id, train_data, val_data)\n    fold_kappas.append(kappa)\n    fold_ckpts.append(ckpt_path)\n\n    if kappa > best_overall:\n        best_overall      = kappa\n        best_ckpt_overall = ckpt_path\n\n# ─────────────────────────────────────────────\n# 11. RESUME FINAL\n# ─────────────────────────────────────────────\nmean_kappa = np.mean(fold_kappas)\nstd_kappa  = np.std(fold_kappas)\n\nprint(f\"\\n{'='*60}\")\nprint(f\"OK Stage {SHARD_ID} -- 5-Fold CV Results (WSI-level QWK)\")\nprint(f\"{'='*60}\")\nfor i, k in enumerate(fold_kappas):\n    print(f\"   Fold {i+1}: QWK = {k:.4f}\")\nprint(f\"{'─'*40}\")\nprint(f\"   Mean QWK : {mean_kappa:.4f} +/- {std_kappa:.4f}\")\nprint(f\"   Best QWK : {best_overall:.4f}\")\nprint(f\"   Stage 3  : 0.7280 +/- 0.0629 | Gain: {mean_kappa-0.7280:+.4f}\")\nprint(f\"   Stage 2  : 0.6832 +/- 0.0340\")\nprint(f\"   Stage 1  : 0.7054 +/- 0.0373\")\nprint(f\"   Stage 0  : 0.5764 +/- 0.0224\")\nprint(f\"   WSIs     : {len(all_wsis)}\")\nprint(f\"   SOTA [1] : 0.934 (Hao et al. 2025)\")\nprint(f\"{'='*60}\")\n\n# Sauvegarder JSON\nresults = {\n    \"stage\":        SHARD_ID,\n    \"n_wsis\":       len(all_wsis),\n    \"n_patches\":    len(unique_records),\n    \"fold_kappas\":  fold_kappas,\n    \"mean_kappa\":   mean_kappa,\n    \"std_kappa\":    std_kappa,\n    \"best_kappa\":   best_overall,\n    \"timestamp\":    datetime.datetime.now().isoformat(),\n    \"metric\":       \"WSI-level QWK (majority vote)\",\n    \"prev_ckpt\":    PREV_CKPT,\n    \"hyperparams\": {\n        \"lr_cnn\":          LR_CNN,\n        \"lr_tr\":           LR_TR,\n        \"epochs\":          EPOCHS,\n        \"patience\":        PATIENCE,\n        \"label_smoothing\": 0.1,\n        \"scheduler\":       \"CosineAnnealingLR\",\n        \"random_erasing\":  True,\n        \"color_jitter\":    \"0.3,0.3,0.2,0.1\",\n        \"gaussian_blur\":   \"kernel=5,p=0.3\",\n    }\n}\nwith open(RESULTS_JSON, \"w\") as f:\n    json.dump(results, f, indent=2)\n\n# Checkpoint final stage\nstage_best = f\"{SAVE_DIR}/stage{SHARD_ID}_best_qwk{best_overall:.4f}.pth\"\nassert best_ckpt_overall is not None, \"Aucun checkpoint trouve\"\nassert os.path.exists(best_ckpt_overall), f\"Fichier manquant: {best_ckpt_overall}\"\nshutil.copy(best_ckpt_overall, stage_best)\n\n# Push final avec TOUS les fichiers\ndef push_final():\n    tmp = f\"{SAVE_DIR}/_push_final\"\n    os.makedirs(tmp, exist_ok=True)\n    pushed = []\n    for c in glob.glob(f\"{SAVE_DIR}/stage{SHARD_ID}_fold*.pth\"):\n        shutil.copy(c, tmp); pushed.append(os.path.basename(c))\n    shutil.copy(stage_best, tmp); pushed.append(os.path.basename(stage_best))\n    if os.path.exists(RESULTS_JSON):\n        shutil.copy(RESULTS_JSON, tmp)\n    json.dump({\n        \"title\": \"Memoire IHC Saved Models\",\n        \"id\": KAGGLE_DS,\n        \"licenses\": [{\"name\": \"CC0-1.0\"}]\n    }, open(f\"{tmp}/dataset-metadata.json\", \"w\"))\n    print(f\"\\n   Pushing FINAL ALL : {pushed}\")\n    r = subprocess.run(\n        [\"kaggle\", \"datasets\", \"version\", \"-p\", tmp,\n         \"-m\", f\"Stage{SHARD_ID} FINAL WSI_QWK={mean_kappa:.4f}+/-{std_kappa:.4f} ({len(all_wsis)} WSIs)\",\n         \"--dir-mode\", \"skip\"],\n        capture_output=True, text=True, timeout=300\n    )\n    print(f\"   {'OK Kaggle' if r.returncode == 0 else 'ERREUR ' + r.stderr[:100]}\")\n\npush_final()\n\nprint(f\"\\nOK Checkpoint permanent : {stage_best}\")\nprint(f\"\\n-> Pour Stage 5 :\")\nprint(f\"   SHARD_ID  = 5\")\nprint(f\"   PREV_CKPT = '/kaggle/working/stage{SHARD_ID}_best_qwk{best_overall:.4f}.pth'\")\nprint(f\"\\nNe ferme pas avant fin du push Kaggle !\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import subprocess, os, shutil, json\n\n# Vérifier fichiers\nfor f in [\"stage4_fold5_best.pth\", \"stage4_best_qwk0.7257.pth\"]:\n    print(f\"{f}: {'OK' if os.path.exists(f'/kaggle/working/{f}') else 'MANQUANT'}\")\n\n# Push avec timeout 900s\ntmp = \"/kaggle/working/_push_emergency\"\nos.makedirs(tmp, exist_ok=True)\n\nfor f in [\"stage4_fold5_best.pth\", \"stage4_best_qwk0.7257.pth\",\n          \"stage3_fold5_best.pth\", \"stage4_fold4_best.pth\"]:\n    path = f\"/kaggle/working/{f}\"\n    if os.path.exists(path):\n        shutil.copy(path, tmp)\n        print(f\"Copied: {f}\")\n\njson.dump({\n    \"title\": \"Memoire IHC Saved Models\",\n    \"id\": \"khouloudbejaoui20/best-models\",\n    \"licenses\": [{\"name\": \"CC0-1.0\"}]\n}, open(f\"{tmp}/dataset-metadata.json\", \"w\"))\n\nr = subprocess.run([\n    \"kaggle\", \"datasets\", \"version\",\n    \"-p\", tmp,\n    \"-m\", \"Stage4 ALL folds final\",\n    \"--dir-mode\", \"skip\"\n], capture_output=True, text=True, timeout=900)\n\nprint(\"OK!\" if r.returncode == 0 else f\"ERREUR: {r.stderr[:300]}\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import subprocess, os, shutil, json\n\nfiles = [\n    \"stage4_fold5_best.pth\",\n    \"stage4_best_qwk0.7257.pth\",\n]\n\nfor fname in files:\n    tmp = f\"/kaggle/working/_push_{fname.replace('.pth','')}\"\n    os.makedirs(tmp, exist_ok=True)\n    shutil.copy(f\"/kaggle/working/{fname}\", tmp)\n    \n    json.dump({\n        \"title\": \"Memoire IHC Saved Models\",\n        \"id\": \"khouloudbejaoui20/best-models\",\n        \"licenses\": [{\"name\": \"CC0-1.0\"}]\n    }, open(f\"{tmp}/dataset-metadata.json\", \"w\"))\n    \n    r = subprocess.run([\n        \"kaggle\", \"datasets\", \"version\",\n        \"-p\", tmp,\n        \"-m\", f\"Stage4 {fname}\",\n        \"--dir-mode\", \"skip\"\n    ], capture_output=True, text=True, timeout=900)\n    \n    print(f\"{fname}: {'OK!' if r.returncode==0 else 'ERREUR'}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import json, os\n\nkaggle_creds = {\n    \"username\": \"khouloudbejaoui20\",\n    \"key\": \"af385e4bd0b6181b85b83134f6f1c405\"  # ← remplace ici\n}\n\nos.makedirs(\"/root/.config/kaggle\", exist_ok=True)\nwith open(\"/root/.config/kaggle/kaggle.json\", \"w\") as f:\n    json.dump(kaggle_creds, f)\nos.chmod(\"/root/.config/kaggle/kaggle.json\", 0o600)\nprint(\"OK!\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import subprocess, shutil, json, os\n\ntmp = \"/kaggle/working/_push_final\"\nos.makedirs(tmp, exist_ok=True)\n\nfor f in [\"stage4_fold5_best.pth\", \"stage4_best_qwk0.7257.pth\"]:\n    if os.path.exists(f\"/kaggle/working/{f}\"):\n        shutil.copy(f\"/kaggle/working/{f}\", tmp)\n        print(f\"Copied: {f}\")\n\njson.dump({\n    \"title\": \"Memoire IHC Saved Models\",\n    \"id\": \"khouloudbejaoui20/best-models\",\n    \"licenses\": [{\"name\": \"CC0-1.0\"}]\n}, open(f\"{tmp}/dataset-metadata.json\", \"w\"))\n\nr = subprocess.run([\n    \"kaggle\", \"datasets\", \"version\",\n    \"-p\", tmp,\n    \"-m\", \"Stage4 FINAL fold5=0.7257\",\n    \"--dir-mode\", \"skip\"\n], capture_output=True, text=True, timeout=900)\n\nprint(\"OK!\" if r.returncode == 0 else f\"ERREUR: {r.stderr[:300]}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import subprocess, shutil, json, os\n\n# Push seulement fold5\ntmp = \"/kaggle/working/_push_fold5_only\"\nos.makedirs(tmp, exist_ok=True)\n\nshutil.copy(\"/kaggle/working/stage4_fold5_best.pth\", tmp)\n\njson.dump({\n    \"title\": \"Memoire IHC Saved Models\",\n    \"id\": \"khouloudbejaoui20/best-models\",\n    \"licenses\": [{\"name\": \"CC0-1.0\"}]\n}, open(f\"{tmp}/dataset-metadata.json\", \"w\"))\n\nr = subprocess.run([\n    \"kaggle\", \"datasets\", \"version\",\n    \"-p\", tmp,\n    \"-m\", \"Stage4 Fold5 QWK=0.7257\",\n    \"--dir-mode\", \"skip\"\n], capture_output=True, text=True, timeout=900)\n\nprint(\"OK!\" if r.returncode == 0 else f\"ERREUR: {r.stderr[:300]}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import subprocess\n\n# Rendre le notebook public via API\nr = subprocess.run([\n    \"kaggle\", \"kernels\", \"push\",\n    \"-p\", \"/kaggle/working\"\n], capture_output=True, text=True)\nprint(r.stdout)\nprint(r.stderr[:200])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import subprocess, shutil, json, os, glob\n\n# Créer nouveau dataset avec compte bejaouikhouloud\ntmp = \"/kaggle/working/_new_ds\"\nos.makedirs(tmp, exist_ok=True)\n\n# Copier fichiers stage4\nfor f in glob.glob(\"/kaggle/working/stage4*.pth\"):\n    shutil.copy(f, tmp)\n    print(f\"Copied: {os.path.basename(f)}\")\n\n# Metadata pour NOUVEAU dataset\njson.dump({\n    \"title\": \"Stage4 Results Bejaoui 2026\",\n    \"id\": \"bejaouikhouloud/stage4-results-bejaoui-2026\",\n    \"licenses\": [{\"name\": \"CC0-1.0\"}]\n}, open(f\"{tmp}/dataset-metadata.json\", \"w\"))\n\n# Créer nouveau dataset\nr = subprocess.run([\n    \"kaggle\", \"datasets\", \"create\",\n    \"-p\", tmp,\n    \"--dir-mode\", \"skip\"\n], capture_output=True, text=True, timeout=900)\n\nprint(\"OK!\" if r.returncode == 0 else f\"ERREUR: {r.stderr[:300]}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import subprocess, os\n\n# Option 1 — depuis best-models\nsubprocess.run([\n    \"kaggle\", \"datasets\", \"download\",\n    \"khouloudbejaoui20/best-models\",\n    \"--file\", \"stage0_fold3_best.pth\",\n    \"-p\", \"/kaggle/working\", \"--unzip\"\n], capture_output=True)\n\nprint(os.path.exists(\"/kaggle/working/stage0_fold3_best.pth\"))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import subprocess\nr = subprocess.run([\n    \"kaggle\", \"datasets\", \"files\",\n    \"khouloudbejaoui20/best-models\"\n], capture_output=True, text=True)\nprint(r.stdout)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import subprocess\nr = subprocess.run([\n    \"kaggle\", \"datasets\", \"files\",\n    \"khouloudbejaoui20/best-models\"\n], capture_output=True, text=True)\nprint(r.stdout)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\n=============================================================\nPatch Extraction — Run 5 (WSI 4000 → 5000)\nWITH AUTO-SAVE every 50 WSIs using VERSION (dataset exists!)\nDataset: prostate-cancer-grade-assessment\nOutput: /kaggle/working/patches_panda_run5/\n=============================================================\nAttache dans Kaggle :\n  prostate-cancer-grade-assessment (competition)\n  khouloudbejaoui20/panda-patches-run5 (existing dataset)\n\"\"\"\n\nimport os, glob, random, numpy as np, pandas as pd\nfrom PIL import Image\nfrom concurrent.futures import ThreadPoolExecutor, as_completed\nimport subprocess, json\n\n# ─── Config ───────────────────────────────────────────────\nWSI_START   = 4000\nWSI_END     = 5000\nPATCH_SIZE  = 512\nTISSUE_THR  = 0.65\nMAX_PATCHES = 15\nN_WORKERS   = 2\nSEED        = 42\nSAVE_EVERY  = 50   # Push toutes les 50 WSIs\n\nrandom.seed(SEED)\nnp.random.seed(SEED)\n\nOUT_DIR        = \"/kaggle/working/patches_panda_run5\"\nPANDA_IMG_DIR  = \"/kaggle/input/competitions/prostate-cancer-grade-assessment/train_images\"\nPANDA_CSV_PATH = \"/kaggle/input/competitions/prostate-cancer-grade-assessment/train.csv\"\nKAGGLE_DS      = \"khouloudbejaoui20/panda-patches-run5\"\n\nos.makedirs(OUT_DIR, exist_ok=True)\n\n# ─── Verify ───────────────────────────────────────────────\nprint(\"Checking PANDA dataset...\")\nassert os.path.exists(PANDA_CSV_PATH), \"CSV not found!\"\nassert os.path.exists(PANDA_IMG_DIR),  \"Images not found!\"\nprint(f\"  OK CSV    : {PANDA_CSV_PATH}\")\nprint(f\"  OK Images : {PANDA_IMG_DIR}\")\n\n# ─── Check existing on Kaggle dataset ─────────────────────\nprint(\"\\nChecking existing patches on Kaggle dataset...\")\nr = subprocess.run([\n    \"kaggle\", \"datasets\", \"files\", KAGGLE_DS\n], capture_output=True, text=True)\n\nkaggle_files = set()\nfor line in r.stdout.split(\"\\n\"):\n    if \".png\" in line:\n        fname = line.strip().split()[0].split(\"/\")[-1]\n        kaggle_files.add(fname)\nprint(f\"  Patches already on Kaggle: {len(kaggle_files)}\")\n\n# ─── Load metadata ────────────────────────────────────────\ndf      = pd.read_csv(PANDA_CSV_PATH)\ndf_run5 = df.iloc[WSI_START:WSI_END].reset_index(drop=True)\nprint(f\"\\nRun 5: WSI {WSI_START} → {WSI_END} ({len(df_run5)} WSIs)\")\n\n# ─── Check already done locally ───────────────────────────\nexisting = glob.glob(f\"{OUT_DIR}/*.png\")\ndone_wsis = set()\nfor f in existing:\n    try:\n        wsi_id = os.path.basename(f).split(\"_slide\")[1].split(\"_\")[0]\n        done_wsis.add(wsi_id)\n    except:\n        pass\n\n# Also mark WSIs already on Kaggle as done\nfor fname in kaggle_files:\n    try:\n        wsi_id = fname.split(\"_slide\")[1].split(\"_\")[0]\n        done_wsis.add(wsi_id)\n    except:\n        pass\n\nprint(f\"Already done WSIs: {len(done_wsis)}\")\nprint(f\"Local patches: {len(existing)}\")\n\n# ─── Push function (VERSION — dataset already exists) ─────\ndef push_to_kaggle(message):\n    try:\n        json.dump({\n            \"title\": \"PANDA Patches Run5\",\n            \"id\": KAGGLE_DS,\n            \"licenses\": [{\"name\": \"CC0-1.0\"}]\n        }, open(f\"{OUT_DIR}/dataset-metadata.json\", \"w\"))\n\n        files_count = len(glob.glob(f\"{OUT_DIR}/*.png\"))\n\n        r = subprocess.run([\n            \"kaggle\", \"datasets\", \"version\",\n            \"-p\", OUT_DIR,\n            \"-m\", message,\n            \"--dir-mode\", \"skip\"\n        ], capture_output=True, text=True, timeout=300)\n\n        status = \"OK\" if r.returncode == 0 else f\"ERREUR: {r.stderr[:100]}\"\n        print(f\"   → Push {status} ({files_count} patches)\")\n    except Exception as e:\n        print(f\"   → Push ERREUR: {e}\")\n\n# ─── Tissue detection ─────────────────────────────────────\ndef has_tissue(patch_arr):\n    gray = np.mean(patch_arr, axis=2)\n    return (gray < 220).sum() / gray.size >= TISSUE_THR\n\n# ─── Extract one WSI ──────────────────────────────────────\ndef extract_wsi(row):\n    import openslide\n\n    wsi_id = str(row['image_id'])\n    label  = int(row['isup_grade'])\n\n    if wsi_id in done_wsis:\n        return 0, label, wsi_id, \"skipped\"\n\n    wsi_path = os.path.join(PANDA_IMG_DIR, f\"{wsi_id}.tiff\")\n    if not os.path.exists(wsi_path):\n        return 0, label, wsi_id, \"not_found\"\n\n    try:\n        slide  = openslide.OpenSlide(wsi_path)\n        w, h   = slide.dimensions\n        saved  = 0\n\n        positions = [\n            (x, y)\n            for y in range(0, h - PATCH_SIZE, PATCH_SIZE)\n            for x in range(0, w - PATCH_SIZE, PATCH_SIZE)\n        ]\n        random.shuffle(positions)\n\n        for x, y in positions:\n            if saved >= MAX_PATCHES:\n                break\n            try:\n                patch     = slide.read_region((x, y), 0, (PATCH_SIZE, PATCH_SIZE))\n                patch_rgb = np.array(patch.convert(\"RGB\"))\n                if has_tissue(patch_rgb):\n                    fname = f\"grade{label}_slide{wsi_id}_{saved:03d}.png\"\n                    Image.fromarray(patch_rgb).save(\n                        os.path.join(OUT_DIR, fname), \"PNG\"\n                    )\n                    saved += 1\n            except:\n                continue\n        slide.close()\n        return saved, label, wsi_id, \"ok\"\n    except Exception as e:\n        return 0, label, wsi_id, \"error\"\n\n# ─── Main loop ────────────────────────────────────────────\nprint(f\"\\nStarting extraction — AUTO-SAVE every {SAVE_EVERY} WSIs...\")\nprint(f\"Config: {PATCH_SIZE}px | tissue>={TISSUE_THR*100:.0f}% | max {MAX_PATCHES}/WSI\")\nprint(\"-\" * 60)\n\ntotal_patches = len(existing)\ntotal_ok      = len(done_wsis)\ntotal_failed  = 0\n\nrows = [row for _, row in df_run5.iterrows()\n        if str(row['image_id']) not in done_wsis]\nprint(f\"WSIs remaining: {len(rows)}\")\n\nwith ThreadPoolExecutor(max_workers=N_WORKERS) as executor:\n    futures = {executor.submit(extract_wsi, row): row for row in rows}\n\n    for i, future in enumerate(as_completed(futures)):\n        try:\n            n, grade, wsi_id, status = future.result()\n            if status in [\"ok\", \"skipped\"]:\n                total_patches += n\n                total_ok      += 1\n                done_wsis.add(wsi_id)\n            else:\n                total_failed += 1\n\n            processed = i + 1\n\n            if processed % 25 == 0 or processed == len(rows):\n                print(f\"  [{processed:4d}/{len(rows)}] \"\n                      f\"OK={total_ok} | Failed={total_failed} | \"\n                      f\"Patches={total_patches}\")\n\n            # AUTO-SAVE every SAVE_EVERY WSIs\n            if processed % SAVE_EVERY == 0:\n                print(f\"\\n  ── AUTO-SAVE at {processed} WSIs ──\")\n                push_to_kaggle(f\"Run5 {processed}/{len(rows)} WSIs done\")\n                print()\n\n        except Exception as e:\n            total_failed += 1\n\n# ─── Final push ───────────────────────────────────────────\nfiles = glob.glob(f\"{OUT_DIR}/*.png\")\nprint(f\"\\n{'='*60}\")\nprint(f\"Run 5 Complete!\")\nprint(f\"WSIs OK      : {total_ok}\")\nprint(f\"WSIs failed  : {total_failed}\")\nprint(f\"Local patches: {len(files)}\")\n\nprint(f\"\\nFinal push...\")\npush_to_kaggle(f\"Run5 FINAL {len(files)} patches\")\n\nprint(f\"\\nNext: SHARD_ID=5, PREV_CKPT=stage4_best_qwk0.7257.pth\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import subprocess, json, glob\n\njson.dump({\n    \"title\": \"PANDA Patches Run5\",\n    \"id\": \"khouloudbejaoui/panda-patches-run5\",\n    \"licenses\": [{\"name\": \"CC0-1.0\"}]\n}, open(\"/kaggle/working/patches_panda_run5/dataset-metadata.json\", \"w\"))\n\nr = subprocess.run([\n    \"kaggle\", \"datasets\", \"version\",\n    \"-p\", \"/kaggle/working/patches_panda_run5\",\n    \"-m\", \"Run5 FINAL 14833 patches 1000 WSIs\",\n    \"--dir-mode\", \"skip\"\n], capture_output=True, text=True, timeout=7200)\n\nprint(\"OK!\" if r.returncode == 0 else f\"ERREUR: {r.stderr[:300]}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import subprocess, json, os, shutil, glob\n\n# Copier 1 seul fichier\ntmp = \"/kaggle/working/run5_one\"\nos.makedirs(tmp, exist_ok=True)\n\nfiles = glob.glob(\"/kaggle/working/patches_panda_run5/*.png\")\nshutil.copy(files[0], tmp)\nprint(f\"Copied: {os.path.basename(files[0])}\")\n\njson.dump({\n    \"title\": \"PANDA Patches Run5\",\n    \"id\": \"khouloudbejaoui/panda-patches-run5\",\n    \"licenses\": [{\"name\": \"CC0-1.0\"}]\n}, open(f\"{tmp}/dataset-metadata.json\", \"w\"))\n\nr = subprocess.run([\n    \"kaggle\", \"datasets\", \"create\",\n    \"-p\", tmp, \"--dir-mode\", \"skip\"\n], capture_output=True, text=True, timeout=120)\n\nprint(\"OK!\" if r.returncode == 0 else f\"ERREUR: {r.stderr[:200]}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import subprocess, json, os, shutil, glob\n\n# 1 seul fichier\ntmp = \"/kaggle/working/run5_one\"\nfiles = glob.glob(\"/kaggle/working/patches_panda_run5/*.png\")\nprint(f\"Fichier: {os.path.basename(files[0])}\")\n\njson.dump({\n    \"title\": \"PANDA Patches Run5\",\n    \"id\": \"khouloudbejaoui/panda-patches-run5\",\n    \"licenses\": [{\"name\": \"CC0-1.0\"}]\n}, open(f\"{tmp}/dataset-metadata.json\", \"w\"))\n\nr = subprocess.run([\n    \"kaggle\", \"datasets\", \"create\",\n    \"-p\", tmp, \"--dir-mode\", \"skip\"\n], capture_output=True, text=True, timeout=120)\n\nprint(\"OK!\" if r.returncode == 0 else f\"ERREUR: {r.stderr[:200]}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import subprocess, json, os, shutil, glob\n\nall_files = sorted(glob.glob(\"/kaggle/working/patches_panda_run5/grade*.png\"))\nprint(f\"Total: {len(all_files)} patches\")\n\nchunk_size = 1000\n\nfor i in range(0, len(all_files), chunk_size):\n    chunk = all_files[i:i+chunk_size]\n    tmp = f\"/kaggle/working/chunk_{i}\"\n    os.makedirs(tmp, exist_ok=True)\n    \n    for f in chunk:\n        shutil.copy(f, tmp)\n    \n    json.dump({\n        \"title\": \"PANDA Patches Run5\",\n        \"id\": \"khouloudbejaoui/panda-patches-run5\",\n        \"licenses\": [{\"name\": \"CC0-1.0\"}]\n    }, open(f\"{tmp}/dataset-metadata.json\", \"w\"))\n    \n    r = subprocess.run([\n        \"kaggle\", \"datasets\", \"version\",\n        \"-p\", tmp,\n        \"-m\", f\"chunk {i}-{i+chunk_size}\",\n        \"--dir-mode\", \"skip\"\n    ], capture_output=True, text=True, timeout=600)\n    \n    status = \"OK!\" if r.returncode == 0 else f\"ERREUR: {r.stderr[:100]}\"\n    print(f\"Chunk {i}-{i+chunk_size}: {status}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import subprocess, json, glob, os, shutil\n\nall_files = sorted(glob.glob(\"/kaggle/working/patches_panda_run5/grade*.png\"))\nprint(f\"Total patches: {len(all_files)}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import subprocess, json, os, shutil, glob\n\nall_files = sorted(glob.glob(\"/kaggle/working/patches_panda_run5/grade*.png\"))\nprint(f\"Total: {len(all_files)} patches\")\n\nchunk_size = 500  # petit chunk = plus rapide !\n\nfor i in range(0, len(all_files), chunk_size):\n    chunk = all_files[i:i+chunk_size]\n    tmp = f\"/kaggle/working/chunk_{i}\"\n    os.makedirs(tmp, exist_ok=True)\n    \n    for f in chunk:\n        shutil.copy(f, tmp)\n    \n    json.dump({\n        \"title\": \"PANDA Patches Run5\",\n        \"id\": \"khouloudbejaoui/panda-patches-run5\",\n        \"licenses\": [{\"name\": \"CC0-1.0\"}]\n    }, open(f\"{tmp}/dataset-metadata.json\", \"w\"))\n    \n    r = subprocess.run([\n        \"kaggle\", \"datasets\", \"version\",\n        \"-p\", tmp,\n        \"-m\", f\"chunk {i}-{i+chunk_size}\",\n        \"--dir-mode\", \"skip\"\n    ], capture_output=True, text=True, timeout=600)\n    \n    status = \"OK!\" if r.returncode == 0 else f\"ERREUR: {r.stderr[:100]}\"\n    print(f\"Chunk {i}-{i+chunk_size}: {status}\")\n    \n    # Nettoyer pour libérer espace\n    shutil.rmtree(tmp)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import subprocess\n# Zipper tous les patches\nsubprocess.run([\n    \"zip\", \"-r\", \n    \"/kaggle/working/patches_run5.zip\",\n    \"/kaggle/working/patches_panda_run5/\"\n], capture_output=True)\nprint(\"Done!\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nsize = os.path.getsize(\"/kaggle/working/patches_run5.zip\") / (1024*1024)\nprint(f\"ZIP size: {size:.1f} MB\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import subprocess, glob, os\n\nfiles = sorted(glob.glob(\"/kaggle/working/patches_panda_run5/grade*.png\"))\nchunk_size = 1000\n\nfor i in range(0, len(files), chunk_size):\n    chunk = files[i:i+chunk_size]\n    zip_path = f\"/kaggle/working/run5_part{i//chunk_size+1}.zip\"\n    \n    r = subprocess.run(\n        [\"zip\", zip_path] + chunk,\n        capture_output=True\n    )\n    size = os.path.getsize(zip_path)/(1024*1024)\n    print(f\"Part {i//chunk_size+1}: {size:.0f} MB — {'OK' if r.returncode==0 else 'ERREUR'}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import subprocess, json, os, shutil\n\ntmp = \"/kaggle/working/zip_push\"\nos.makedirs(tmp, exist_ok=True)\nshutil.copy(\"/kaggle/working/patches_run5.zip\", tmp)\n\njson.dump({\n    \"title\": \"PANDA Patches Run5\",\n    \"id\": \"khouloudbejaoui/panda-patches-run5\",\n    \"licenses\": [{\"name\": \"CC0-1.0\"}]\n}, open(f\"{tmp}/dataset-metadata.json\", \"w\"))\n\nr = subprocess.run([\n    \"kaggle\", \"datasets\", \"version\",\n    \"-p\", tmp,\n    \"-m\", \"Run5 ZIP 14833 patches\",\n    \"--dir-mode\", \"skip\"\n], capture_output=True, text=True, timeout=7200)\n\nprint(\"OK!\" if r.returncode == 0 else f\"ERREUR: {r.stderr[:300]}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import subprocess, json, os\n\n# Push le dossier ZIP directement\njson.dump({\n    \"title\": \"PANDA Patches Run5\",\n    \"id\": \"khouloudbejaoui/panda-patches-run5\",\n    \"licenses\": [{\"name\": \"CC0-1.0\"}]\n}, open(\"/kaggle/working/dataset-metadata.json\", \"w\"))\n\nr = subprocess.run([\n    \"kaggle\", \"datasets\", \"version\",\n    \"-p\", \"/kaggle/working\",\n    \"-m\", \"Run5 ZIP 14833 patches\",\n    \"--dir-mode\", \"skip\"\n], capture_output=True, text=True, timeout=7200)\n\nprint(\"OK!\" if r.returncode == 0 else f\"ERREUR: {r.stderr[:300]}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import shutil, os, glob\n\n# Supprimer tout sauf les patches\nfor d in glob.glob(\"/kaggle/working/chunk_*\"):\n    shutil.rmtree(d)\nfor d in glob.glob(\"/kaggle/working/run5_*\"):\n    shutil.rmtree(d)\nfor f in glob.glob(\"/kaggle/working/*.zip\"):\n    os.remove(f)\n    print(f\"Deleted: {f}\")\n\n# Espace libre\nimport subprocess\nr = subprocess.run([\"df\", \"-h\", \"/kaggle/working\"], \n                   capture_output=True, text=True)\nprint(r.stdout)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import subprocess, json, os, shutil, glob\n\nall_files = sorted(glob.glob(\"/kaggle/working/patches_panda_run5/grade*.png\"))\nprint(f\"Total: {len(all_files)} patches\")\n\n# Push par chunks de 500\nchunk_size = 500\n\nfor i in range(0, len(all_files), chunk_size):\n    chunk = all_files[i:i+chunk_size]\n    tmp = f\"/kaggle/working/chunk_{i}\"\n    os.makedirs(tmp, exist_ok=True)\n    \n    for f in chunk:\n        shutil.copy(f, tmp)\n    \n    json.dump({\n        \"title\": \"PANDA Patches Run5\",\n        \"id\": \"khouloudbejaoui/panda-patches-run5\",\n        \"licenses\": [{\"name\": \"CC0-1.0\"}]\n    }, open(f\"{tmp}/dataset-metadata.json\", \"w\"))\n    \n    r = subprocess.run([\n        \"kaggle\", \"datasets\", \"version\",\n        \"-p\", tmp,\n        \"-m\", f\"chunk {i}-{i+chunk_size}\",\n        \"--dir-mode\", \"skip\"\n    ], capture_output=True, text=True, timeout=600)\n    \n    status = \"OK!\" if r.returncode == 0 else f\"ERREUR: {r.stderr[:100]}\"\n    print(f\"Chunk {i}-{i+chunk_size}: {status}\")\n    \n    # Nettoyer pour libérer espace\n    shutil.rmtree(tmp)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import subprocess, glob, os\n\nfiles = sorted(glob.glob(\"/kaggle/working/patches_panda_run5/grade*.png\"))\nprint(f\"Total: {len(files)} patches\")\n\nchunk_size = 1000\n\nfor i in range(0, len(files), chunk_size):\n    chunk = files[i:i+chunk_size]\n    part_num = i//chunk_size + 1\n    zip_path = f\"/kaggle/working/run5_part{part_num}.zip\"\n    \n    r = subprocess.run(\n        [\"zip\", zip_path] + chunk,\n        capture_output=True\n    )\n    size = os.path.getsize(zip_path)/(1024*1024)\n    print(f\"Part {part_num}: {size:.0f} MB {'OK' if r.returncode==0 else 'ERREUR'}\")\n\nprint(\"\\nTous les ZIPs créés!\")\nprint(\"Va dans Output tab → télécharge chaque run5_partX.zip\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import subprocess, os\n\nsubprocess.run([\n    \"kaggle\", \"datasets\", \"download\",\n    \"khouloudbejaoui20/best-models\",\n    \"--file\", \"stage4_best_qwk0.7257.pth\",\n    \"-p\", \"/kaggle/working\", \"--unzip\"\n], capture_output=True)\n\nprint(os.path.exists(\"/kaggle/working/stage4_best_qwk0.7257.pth\"))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import subprocess, os\n\nsubprocess.run([\n    \"kaggle\", \"datasets\", \"download\",\n    \"khouloudbejaoui20/best-models\",\n    \"--file\", \"stage4_best_qwk0.7257.pth\",\n    \"-p\", \"/kaggle/working\", \"--unzip\"\n], capture_output=True)\n\nprint(os.path.exists(\"/kaggle/working/stage4_best_qwk0.7257.pth\"))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import subprocess\nr = subprocess.run([\n    \"kaggle\", \"datasets\", \"files\",\n    \"khouloudbejaoui20/best-models\"\n], capture_output=True, text=True)\nprint(r.stdout)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import subprocess, os, json\n\n# Configurer avec bejaouikhouloud\nkaggle_creds = {\n    \"username\": \"bejaouikhouloud\",\n    \"key\": \"af385e4bd0b6181b85b83134f6f1c405\"\n}\nos.makedirs(\"/root/.config/kaggle\", exist_ok=True)\nwith open(\"/root/.config/kaggle/kaggle.json\", \"w\") as f:\n    json.dump(kaggle_creds, f)\nos.chmod(\"/root/.config/kaggle/kaggle.json\", 0o600)\n\n# Télécharger\nsubprocess.run([\n    \"kaggle\", \"datasets\", \"download\",\n    \"bejaouikhouloud/stage4-results-bejaoui-2026\",\n    \"--file\", \"stage4_fold5_best.pth\",\n    \"-p\", \"/kaggle/working\", \"--unzip\"\n], capture_output=True)\n\nprint(os.path.exists(\"/kaggle/working/stage4_fold5_best.pth\"))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import glob, os\n\n# Chercher patches run5\nfor root, dirs, files in os.walk(\"/kaggle/input\"):\n    for f in files:\n        if f.endswith(\".png\") and \"run5\" in root.lower():\n            print(os.path.join(root, f))\n            break\n    if any(f.endswith(\".png\") for f in files) and \"run5\" in root.lower():\n        print(f\"Run5 path: {root}\")\n        print(f\"Files: {len([f for f in files if f.endswith('.png')])}\")\n        break","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\n=============================================================\nStage 5 — Incremental 5-Fold CV\nFine-tune depuis Stage 4 (best QWK=0.7257)\nDonnées cumulatives :\n  Shard 0 : patches_400              (399  WSIs,  4,319 patches)\n  Shard 1 : panda-patches-run1       (1000 WSIs, 14,842 patches)\n  Shard 2 : panda-patches-run2       (999  WSIs, 14,867 patches)\n  Shard 3 : panda-patches-run3       (999  WSIs, 14,835 patches)\n  Shard 4 : panda-patches-run4       (1000 WSIs, 14,808 patches)\n  Shard 5 : panda-patches-run5-parts (1000 WSIs, 14,833 patches)\nTotal     : ~5,397 WSIs, ~78,504 patches\nTarget    : WSI_QWK > 0.7257\n=============================================================\nAttache dans Kaggle :\n  khouloudbejaoui20/dataset\n  khouloudbejaoui20/panda-patches-run1\n  khouloudbejaoui20/panda-patches-run2\n  khouloudbejaoui20/panda-patches-run3\n  khouloudbejaoui20/panda-patches-run4\n  bejaouikhouloud/panda-patches-run5-parts\n  khouloudbejaoui20/best-models\n\"\"\"\n\nimport os, glob, json, random, shutil, subprocess, datetime\nimport numpy as np\nfrom PIL import Image\nfrom collections import Counter\nfrom sklearn.model_selection import StratifiedGroupKFold\nfrom sklearn.metrics import cohen_kappa_score, accuracy_score\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as T\nimport torchvision.models as models\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\n\n# ─────────────────────────────────────────────\n# 1. CONFIG\n# ─────────────────────────────────────────────\nSHARD_ID    = 5\nSTART_FOLD  = 0\nDONE_KAPPAS = []\nDONE_CKPTS  = []\n\nPREV_CKPT   = \"/kaggle/working/stage4_fold5_best.pth\"\n\nN_FOLDS     = 5\nSEED        = 42\nBATCH_SIZE  = 32\nNUM_WORKERS = 2\nEPOCHS      = 30\nPATIENCE    = 7\nLR_CNN      = 2e-5\nLR_TR       = 4e-5\n\nrandom.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)\ntorch.backends.cudnn.deterministic = True\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Device    : {DEVICE}\")\nprint(f\"Stage     : {SHARD_ID} (fine-tune depuis Stage 4 QWK=0.7257)\")\nprint(f\"EPOCHS={EPOCHS} | PATIENCE={PATIENCE} | LR_CNN={LR_CNN} | LR_TR={LR_TR}\")\n\n# ─────────────────────────────────────────────\n# 2. CHEMINS PATCHES\n# ─────────────────────────────────────────────\nBASE     = \"/kaggle/input/datasets/khouloudbejaoui20\"\nRUN5_BASE = \"/kaggle/input/datasets/bejaouikhouloud/panda-patches-run5-parts\"\n\nSHARD_PATHS = {\n    0: f\"{BASE}/dataset/patches_400-20260630T212405Z-3-001/patches_400\",\n    1: f\"{BASE}/panda-patches-run1/kaggle/working/patches_panda_new\",\n    2: f\"{BASE}/panda-patches-run2/kaggle/working/patches_panda_run2\",\n    3: f\"{BASE}/panda-patches-run3/kaggle/working/patches_panda_run3_correct\",\n    4: f\"{BASE}/panda-patches-run4\",\n}\n\nSAVE_DIR     = \"/kaggle/working\"\nKAGGLE_DS    = \"khouloudbejaoui20/best-models\"\nRESULTS_JSON = f\"{SAVE_DIR}/stage5_cv_results.json\"\n\n# ─────────────────────────────────────────────\n# 3. VERIFICATION\n# ─────────────────────────────────────────────\nprint(f\"\\n── Verification ────────────────────────\")\n\n# Download PREV_CKPT si nécessaire\nif not os.path.exists(PREV_CKPT):\n    print(\"  Downloading stage4_fold5_best.pth...\")\n    # Essayer depuis bejaouikhouloud\n    subprocess.run([\n        \"kaggle\", \"datasets\", \"download\",\n        \"bejaouikhouloud/stage4-results-bejaoui-2026\",\n        \"--file\", \"stage4_fold5_best.pth\",\n        \"-p\", SAVE_DIR, \"--unzip\"\n    ], capture_output=True)\n\nassert os.path.exists(PREV_CKPT), f\"PREV_CKPT non trouve: {PREV_CKPT}\"\nprint(f\"  OK Stage 4 checkpoint (QWK=0.7257)\")\n\n# Vérifier Shards 0-4\nfor shard_id, path in SHARD_PATHS.items():\n    if os.path.exists(path):\n        files = glob.glob(f\"{path}/*.png\")\n        if not files:\n            files = glob.glob(f\"{path}/**/*.png\", recursive=True)\n        print(f\"  OK Shard {shard_id}: {len(files)} patches\")\n    else:\n        print(f\"  MANQUANT Shard {shard_id}: {path}\")\n\n# Vérifier Run5\nrun5_count = 0\nfor part in range(1, 16):\n    path = f\"{RUN5_BASE}/run5_part{part}/kaggle/working/patches_panda_run5\"\n    if os.path.exists(path):\n        files = glob.glob(f\"{path}/*.png\")\n        run5_count += len(files)\nprint(f\"  OK Shard 5 (Run5): {run5_count} patches (15 parties)\")\n\n# ─────────────────────────────────────────────\n# 4. CHARGEMENT PATCHES\n# ─────────────────────────────────────────────\nprint(f\"\\n── Patches cumulatifs (Shard 0 à {SHARD_ID}) ──\")\n\ndef load_patches(folder):\n    files = glob.glob(f\"{folder}/*.png\")\n    if not files:\n        files = glob.glob(f\"{folder}/**/*.png\", recursive=True)\n    records = []\n    for fp in files:\n        fname = os.path.basename(fp)\n        try:\n            parts  = fname.split(\"_slide\")\n            label  = int(parts[0].replace(\"grade\", \"\"))\n            wsi_id = parts[1].split(\"_\")[0]\n            records.append({\"path\": fp, \"wsi_id\": wsi_id, \"label\": label})\n        except:\n            continue\n    return records\n\nall_records = []\n\n# Shards 0-4\nfor shard_id, path in SHARD_PATHS.items():\n    if path and os.path.exists(path):\n        r = load_patches(path)\n        all_records.extend(r)\n        print(f\"  Shard {shard_id}: {len(r)} patches\")\n    else:\n        print(f\"  Shard {shard_id}: NON TROUVE\")\n\n# Shard 5 — 15 parties\nrun5_total = 0\nfor part in range(1, 16):\n    path = f\"{RUN5_BASE}/run5_part{part}/kaggle/working/patches_panda_run5\"\n    if os.path.exists(path):\n        r = load_patches(path)\n        all_records.extend(r)\n        run5_total += len(r)\nprint(f\"  Shard 5 (15 parts): {run5_total} patches\")\n\n# Deduplicate\nseen = set()\nunique_records = []\nfor r in all_records:\n    fname = os.path.basename(r[\"path\"])\n    if fname not in seen:\n        seen.add(fname)\n        unique_records.append(r)\n\nall_wsis = list(set(r[\"wsi_id\"] for r in unique_records))\nprint(f\"\\n  Total unique : {len(unique_records)} patches\")\nprint(f\"  Total WSIs  : {len(all_wsis)}\")\ncnt = Counter(r[\"label\"] for r in unique_records)\nfor c in range(6):\n    print(f\"    ISUP {c}: {cnt.get(c, 0)}\")\n\nassert len(all_wsis) > 4500, f\"Trop peu de WSIs ({len(all_wsis)})\"\nprint(f\"\\n  OK {len(all_wsis)} WSIs uniques confirmes\")\n\n# ─────────────────────────────────────────────\n# 5. TRANSFORMS\n# ─────────────────────────────────────────────\nMEAN = [0.485, 0.456, 0.406]\nSTD  = [0.229, 0.224, 0.225]\n\ntrain_tf = T.Compose([\n    T.Resize((224, 224)),\n    T.RandomHorizontalFlip(),\n    T.RandomVerticalFlip(),\n    T.RandomRotation(90),\n    T.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.2, hue=0.1),\n    T.RandomApply([T.GaussianBlur(5)], p=0.3),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n    T.RandomApply([T.RandomErasing(p=0.5, scale=(0.02, 0.2))], p=0.2),\n])\n\nval_tf = T.Compose([\n    T.Resize((224, 224)),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n])\n\nclass PatchDataset(Dataset):\n    def __init__(self, records, transform):\n        self.records   = records\n        self.transform = transform\n    def __len__(self): return len(self.records)\n    def __getitem__(self, idx):\n        r = self.records[idx]\n        try:\n            img = Image.open(r[\"path\"]).convert(\"RGB\")\n        except:\n            img = Image.new(\"RGB\", (224, 224), (255, 255, 255))\n        return self.transform(img), int(r[\"label\"])\n\n# ─────────────────────────────────────────────\n# 6. MODELE\n# ─────────────────────────────────────────────\nclass HybridCNNTransformer(nn.Module):\n    def __init__(self, num_classes=6, d_model=512, nhead=8,\n                 num_layers=2, dim_feedforward=2048, dropout=0.1):\n        super().__init__()\n        backbone = models.resnet50(weights=models.ResNet50_Weights.DEFAULT)\n        self.cnn_backbone = nn.Sequential(*list(backbone.children())[:-2])\n        self.pool         = nn.AdaptiveAvgPool2d((1, 1))\n        self.proj         = nn.Sequential(\n            nn.Identity(), nn.Identity(),\n            nn.Linear(2048, d_model), nn.LayerNorm(d_model)\n        )\n        self.cls_token = nn.Parameter(torch.zeros(1, 1, d_model))\n        self.pos_embed = nn.Parameter(torch.zeros(1, 2, d_model))\n        encoder_layer  = nn.TransformerEncoderLayer(\n            d_model=d_model, nhead=nhead,\n            dim_feedforward=dim_feedforward,\n            dropout=dropout, batch_first=True\n        )\n        self.transformer = nn.TransformerEncoder(\n            encoder_layer, num_layers=num_layers,\n            norm=nn.LayerNorm(d_model)\n        )\n        self.classifier = nn.Sequential(\n            nn.Linear(d_model, 256), nn.ReLU(),\n            nn.Dropout(dropout), nn.Linear(256, num_classes)\n        )\n        nn.init.trunc_normal_(self.cls_token, std=0.02)\n        nn.init.trunc_normal_(self.pos_embed, std=0.02)\n\n    def forward(self, x):\n        f   = self.pool(self.cnn_backbone(x)).flatten(1)\n        f   = self.proj(f).unsqueeze(1)\n        cls = self.cls_token.expand(x.size(0), -1, -1)\n        seq = torch.cat([cls, f], dim=1) + self.pos_embed\n        return self.classifier(self.transformer(seq)[:, 0])\n\ndef build_model(prev_ckpt=None):\n    model  = HybridCNNTransformer().to(DEVICE)\n    n_gpus = torch.cuda.device_count()\n    if n_gpus > 1:\n        model = nn.DataParallel(model)\n        print(f\"  OK DataParallel: {n_gpus} GPUs\")\n    n_param = sum(p.numel() for p in model.parameters())\n    assert n_param == 30_998_342, f\"Mismatch: {n_param}\"\n    print(f\"  OK Parameters: {n_param:,}\")\n    if prev_ckpt and os.path.exists(prev_ckpt):\n        ckpt  = torch.load(prev_ckpt, map_location=DEVICE, weights_only=False)\n        state = ckpt[\"model_state_dict\"]\n        if n_gpus > 1 and not list(state.keys())[0].startswith(\"module.\"):\n            state = {\"module.\" + k: v for k, v in state.items()}\n        elif n_gpus <= 1 and list(state.keys())[0].startswith(\"module.\"):\n            state = {k[7:]: v for k, v in state.items()}\n        model.load_state_dict(state)\n        print(f\"  OK Loaded Stage 4 (QWK={ckpt.get('kappa', '?')})\")\n    return model\n\n# ─────────────────────────────────────────────\n# 7. WSI-LEVEL QWK\n# ─────────────────────────────────────────────\ndef wsi_level_qwk(model, val_data, device):\n    model.eval()\n    wsi_preds  = {}\n    wsi_labels = {}\n    all_preds  = []\n    all_labels = []\n    loader = DataLoader(\n        PatchDataset(val_data, val_tf),\n        batch_size=BATCH_SIZE, shuffle=False,\n        num_workers=NUM_WORKERS\n    )\n    with torch.no_grad():\n        for imgs, labels in loader:\n            imgs  = imgs.to(device)\n            preds = model(imgs).argmax(1).cpu().numpy()\n            all_preds.extend(preds)\n            all_labels.extend(labels.numpy())\n    for i, r in enumerate(val_data):\n        wid = r[\"wsi_id\"]\n        if wid not in wsi_preds:\n            wsi_preds[wid]  = []\n            wsi_labels[wid] = r[\"label\"]\n        wsi_preds[wid].append(all_preds[i])\n    wf = []; wl = []\n    for wid in wsi_preds:\n        votes = wsi_preds[wid]\n        wf.append(max(set(votes), key=votes.count))\n        wl.append(wsi_labels[wid])\n    patch_qwk = cohen_kappa_score(all_labels, all_preds, weights=\"quadratic\")\n    wsi_qwk   = cohen_kappa_score(wl, wf, weights=\"quadratic\")\n    return patch_qwk, wsi_qwk, len(wf)\n\n# ─────────────────────────────────────────────\n# 8. PUSH KAGGLE\n# ─────────────────────────────────────────────\ndef push_to_kaggle(ckpt_path, message):\n    try:\n        tmp = f\"{SAVE_DIR}/_push_tmp\"\n        os.makedirs(tmp, exist_ok=True)\n        for c in glob.glob(f\"{SAVE_DIR}/stage{SHARD_ID}_fold*.pth\"):\n            shutil.copy(c, tmp)\n        shutil.copy(ckpt_path, tmp)\n        if os.path.exists(RESULTS_JSON):\n            shutil.copy(RESULTS_JSON, tmp)\n        json.dump({\n            \"title\": \"Memoire IHC Saved Models\",\n            \"id\": KAGGLE_DS,\n            \"licenses\": [{\"name\": \"CC0-1.0\"}]\n        }, open(f\"{tmp}/dataset-metadata.json\", \"w\"))\n        files_pushed = [os.path.basename(f) for f in glob.glob(f\"{tmp}/*.pth\")]\n        print(f\"   Pushing: {files_pushed}\")\n        r = subprocess.run(\n            [\"kaggle\", \"datasets\", \"version\",\n             \"-p\", tmp, \"-m\", message, \"--dir-mode\", \"skip\"],\n            capture_output=True, text=True, timeout=300\n        )\n        print(f\"   {'OK Kaggle' if r.returncode==0 else 'ERREUR '+r.stderr[:100]}\")\n    except Exception as e:\n        print(f\"   ERREUR push: {e}\")\n\n# ─────────────────────────────────────────────\n# 9. TRAIN ONE FOLD\n# ─────────────────────────────────────────────\ndef train_one_fold(fold_id, train_data, val_data):\n    print(f\"\\n  Fold {fold_id+1}/{N_FOLDS} ──────────────────────────\")\n    print(f\"     Train: {len(train_data)} | Val: {len(val_data)}\")\n\n    train_loader = DataLoader(\n        PatchDataset(train_data, train_tf),\n        batch_size=BATCH_SIZE, shuffle=True,\n        num_workers=NUM_WORKERS, pin_memory=True\n    )\n\n    model = build_model(PREV_CKPT)\n\n    cnt_tr = Counter(r[\"label\"] for r in train_data)\n    w = torch.tensor(\n        [1.0 / max(cnt_tr.get(c, 1), 1) for c in range(6)],\n        dtype=torch.float\n    )\n    w = (w / w.sum() * 6).to(DEVICE)\n    criterion = nn.CrossEntropyLoss(weight=w, label_smoothing=0.1)\n\n    base_model = model.module if hasattr(model, \"module\") else model\n    optimizer  = optim.AdamW([\n        {\"params\": base_model.cnn_backbone.parameters(), \"lr\": LR_CNN},\n        {\"params\": [p for n, p in base_model.named_parameters()\n                    if not n.startswith(\"cnn_backbone\")], \"lr\": LR_TR},\n    ], weight_decay=1e-4)\n\n    scheduler = CosineAnnealingLR(optimizer, T_max=EPOCHS, eta_min=1e-7)\n\n    best_kappa = 0.0\n    no_improve = 0\n    best_ckpt  = f\"{SAVE_DIR}/stage{SHARD_ID}_fold{fold_id+1}_best.pth\"\n\n    for epoch in range(1, EPOCHS + 1):\n        model.train()\n        t_loss, t_correct, t_total = 0.0, 0, 0\n        for imgs, labels in train_loader:\n            imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n            optimizer.zero_grad()\n            out  = model(imgs)\n            loss = criterion(out, labels)\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n            optimizer.step()\n            t_loss    += loss.item() * imgs.size(0)\n            t_correct += (out.argmax(1) == labels).sum().item()\n            t_total   += imgs.size(0)\n        scheduler.step()\n\n        _, wsi_kappa, n_wsis_val = wsi_level_qwk(model, val_data, DEVICE)\n        is_best = wsi_kappa > best_kappa\n\n        print(f\"     Ep[{epoch:02d}/{EPOCHS}] \"\n              f\"loss={t_loss/t_total:.4f} acc={t_correct/t_total:.3f} | \"\n              f\"WSI_QWK={wsi_kappa:.4f} ({n_wsis_val} WSIs)\"\n              + (\" *\" if is_best else \"\"))\n\n        if is_best:\n            best_kappa = wsi_kappa\n            no_improve = 0\n            torch.save({\n                \"epoch\":            epoch,\n                \"model_state_dict\": model.state_dict() if not hasattr(model, \"module\")\n                                    else model.module.state_dict(),\n                \"kappa\":            wsi_kappa,\n                \"shard_id\":         SHARD_ID,\n                \"fold_id\":          fold_id,\n                \"n_wsis\":           len(all_wsis),\n                \"hyperparams\": {\n                    \"lr_cnn\": LR_CNN, \"lr_tr\": LR_TR,\n                    \"epochs\": EPOCHS, \"patience\": PATIENCE,\n                    \"label_smoothing\": 0.1,\n                }\n            }, best_ckpt)\n        else:\n            no_improve += 1\n            if no_improve >= PATIENCE:\n                print(f\"     Early stopping epoch {epoch}\")\n                break\n\n    print(f\"  OK Fold {fold_id+1} best WSI_QWK: {best_kappa:.4f}\")\n    push_to_kaggle(best_ckpt,\n        f\"Stage{SHARD_ID} Fold{fold_id+1} WSI_QWK={best_kappa:.4f}\")\n    return best_kappa, best_ckpt\n\n# ─────────────────────────────────────────────\n# 10. 5-FOLD CV\n# ─────────────────────────────────────────────\nprint(f\"\\n{'='*60}\")\nprint(f\"Stage {SHARD_ID} -- {len(all_wsis)} WSIs -- 5-Fold CV\")\nprint(f\"Fine-tune depuis Stage 4 (QWK=0.7257)\")\nprint(f\"{'='*60}\")\nif DONE_KAPPAS:\n    print(f\"Folds deja completes : {DONE_KAPPAS}\")\n\nwsi_ids    = [r[\"wsi_id\"] for r in unique_records]\nlabels_arr = [r[\"label\"]  for r in unique_records]\nindices    = list(range(len(unique_records)))\n\nsgkf = StratifiedGroupKFold(n_splits=N_FOLDS, shuffle=True, random_state=SEED)\n\nfold_kappas       = DONE_KAPPAS.copy()\nfold_ckpts        = DONE_CKPTS.copy()\nbest_overall      = max(DONE_KAPPAS) if DONE_KAPPAS else -1.0\nbest_ckpt_overall = None\n\nfor fold_id, (train_idx, val_idx) in enumerate(\n        sgkf.split(indices, labels_arr, groups=wsi_ids)):\n\n    if fold_id < START_FOLD:\n        continue\n\n    train_data = [unique_records[i] for i in train_idx]\n    val_data   = [unique_records[i] for i in val_idx]\n\n    tw = set(unique_records[i][\"wsi_id\"] for i in train_idx)\n    vw = set(unique_records[i][\"wsi_id\"] for i in val_idx)\n    assert len(tw & vw) == 0, f\"DATA LEAKAGE fold {fold_id+1}!\"\n\n    kappa, ckpt_path = train_one_fold(fold_id, train_data, val_data)\n    fold_kappas.append(kappa)\n    fold_ckpts.append(ckpt_path)\n\n    if kappa > best_overall:\n        best_overall      = kappa\n        best_ckpt_overall = ckpt_path\n\n# ─────────────────────────────────────────────\n# 11. RESUME FINAL\n# ─────────────────────────────────────────────\nmean_kappa = np.mean(fold_kappas)\nstd_kappa  = np.std(fold_kappas)\n\nprint(f\"\\n{'='*60}\")\nprint(f\"OK Stage {SHARD_ID} -- 5-Fold CV Results\")\nprint(f\"{'='*60}\")\nfor i, k in enumerate(fold_kappas):\n    print(f\"   Fold {i+1}: QWK = {k:.4f}\")\nprint(f\"{'─'*40}\")\nprint(f\"   Mean QWK : {mean_kappa:.4f} +/- {std_kappa:.4f}\")\nprint(f\"   Best QWK : {best_overall:.4f}\")\nprint(f\"   Stage 4  : 0.7112 +/- 0.0130\")\nprint(f\"   Stage 3  : 0.7280 +/- 0.0629\")\nprint(f\"   WSIs     : {len(all_wsis)}\")\nprint(f\"   SOTA [1] : 0.934\")\nprint(f\"{'='*60}\")\n\nresults = {\n    \"stage\":       SHARD_ID,\n    \"n_wsis\":      len(all_wsis),\n    \"n_patches\":   len(unique_records),\n    \"fold_kappas\": fold_kappas,\n    \"mean_kappa\":  mean_kappa,\n    \"std_kappa\":   std_kappa,\n    \"best_kappa\":  best_overall,\n    \"timestamp\":   datetime.datetime.now().isoformat(),\n    \"prev_ckpt\":   PREV_CKPT,\n}\nwith open(RESULTS_JSON, \"w\") as f:\n    json.dump(results, f, indent=2)\n\nif best_ckpt_overall:\n    stage_best = f\"{SAVE_DIR}/stage{SHARD_ID}_best_qwk{best_overall:.4f}.pth\"\n    shutil.copy(best_ckpt_overall, stage_best)\n    print(f\"\\nOK Best checkpoint: {stage_best}\")\n    print(f\"\\n-> Pour Stage 6 :\")\n    print(f\"   SHARD_ID  = 6\")\n    print(f\"   PREV_CKPT = '{stage_best}'\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, glob\n\nckpts = glob.glob(\"/kaggle/working/stage5*.pth\")\nprint(\"Checkpoints:\")\nfor c in ckpts:\n    mb = os.path.getsize(c)/(1024*1024)\n    print(f\"  {os.path.basename(c)} ({mb:.1f} MB)\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n\"\"\"\n=============================================================\nStage 5 — Incremental 5-Fold CV    fold2\nFine-tune depuis Stage 4 (best QWK=0.7257)\nDonnées cumulatives :\n  Shard 0 : patches_400              (399  WSIs,  4,319 patches)\n  Shard 1 : panda-patches-run1       (1000 WSIs, 14,842 patches)\n  Shard 2 : panda-patches-run2       (999  WSIs, 14,867 patches)\n  Shard 3 : panda-patches-run3       (999  WSIs, 14,835 patches)\n  Shard 4 : panda-patches-run4       (1000 WSIs, 14,808 patches)\n  Shard 5 : panda-patches-run5-parts (1000 WSIs, 14,833 patches)\nTotal     : ~5,397 WSIs, ~78,504 patches\nTarget    : WSI_QWK > 0.7257\n=============================================================\nAttache dans Kaggle :\n  khouloudbejaoui20/dataset\n  khouloudbejaoui20/panda-patches-run1\n  khouloudbejaoui20/panda-patches-run2\n  khouloudbejaoui20/panda-patches-run3\n  khouloudbejaoui20/panda-patches-run4\n  bejaouikhouloud/panda-patches-run5-parts\n  khouloudbejaoui20/best-models\n\"\"\"\n\nimport os, glob, json, random, shutil, subprocess, datetime\nimport numpy as np\nfrom PIL import Image\nfrom collections import Counter\nfrom sklearn.model_selection import StratifiedGroupKFold\nfrom sklearn.metrics import cohen_kappa_score, accuracy_score\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as T\nimport torchvision.models as models\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\n\n# ─────────────────────────────────────────────\n# 1. CONFIG\n# ─────────────────────────────────────────────\nSHARD_ID    = 5\nSTART_FOLD  = 1   # Fold 2 ← change 0 → 1\nDONE_KAPPAS = [0.7415]  # Fold 1 résultat\nDONE_CKPTS  = [\"/kaggle/working/stage5_fold1_best.pth\"]  # perdu mais noté\n\n\n\n\n\nPREV_CKPT = \"/kaggle/working/stage4_fold5_best.pth\"\n\n\n\nN_FOLDS     = 5\nSEED        = 42\nBATCH_SIZE  = 32\nNUM_WORKERS = 2\nEPOCHS      = 30\nPATIENCE    = 7\nLR_CNN      = 2e-5\nLR_TR       = 4e-5\n\nrandom.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)\ntorch.backends.cudnn.deterministic = True\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Device    : {DEVICE}\")\nprint(f\"Stage     : {SHARD_ID} (fine-tune depuis Stage 4 QWK=0.7257)\")\nprint(f\"EPOCHS={EPOCHS} | PATIENCE={PATIENCE} | LR_CNN={LR_CNN} | LR_TR={LR_TR}\")\n\n# ─────────────────────────────────────────────\n# 2. CHEMINS PATCHES\n# ─────────────────────────────────────────────\nBASE     = \"/kaggle/input/datasets/khouloudbejaoui20\"\nRUN5_BASE = \"/kaggle/input/datasets/bejaouikhouloud/panda-patches-run5-parts\"\n\nSHARD_PATHS = {\n    0: f\"{BASE}/dataset/patches_400-20260630T212405Z-3-001/patches_400\",\n    1: f\"{BASE}/panda-patches-run1/kaggle/working/patches_panda_new\",\n    2: f\"{BASE}/panda-patches-run2/kaggle/working/patches_panda_run2\",\n    3: f\"{BASE}/panda-patches-run3/kaggle/working/patches_panda_run3_correct\",\n    4: f\"{BASE}/panda-patches-run4\",\n}\n\nSAVE_DIR     = \"/kaggle/working\"\nKAGGLE_DS    = \"khouloudbejaoui20/best-models\"\nRESULTS_JSON = f\"{SAVE_DIR}/stage5_cv_results.json\"\n\n# ─────────────────────────────────────────────\n# 3. VERIFICATION\n# ─────────────────────────────────────────────\nprint(f\"\\n── Verification ────────────────────────\")\n\n# Download PREV_CKPT si nécessaire\nif not os.path.exists(PREV_CKPT):\n    print(\"  Downloading stage4_fold5_best.pth...\")\n    # Essayer depuis bejaouikhouloud\n    subprocess.run([\n        \"kaggle\", \"datasets\", \"download\",\n        \"bejaouikhouloud/stage4-results-bejaoui-2026\",\n        \"--file\", \"stage4_fold5_best.pth\",\n        \"-p\", SAVE_DIR, \"--unzip\"\n    ], capture_output=True)\n\nassert os.path.exists(PREV_CKPT), f\"PREV_CKPT non trouve: {PREV_CKPT}\"\nprint(f\"  OK Stage 4 checkpoint (QWK=0.7257)\")\n\n# Vérifier Shards 0-4\nfor shard_id, path in SHARD_PATHS.items():\n    if os.path.exists(path):\n        files = glob.glob(f\"{path}/*.png\")\n        if not files:\n            files = glob.glob(f\"{path}/**/*.png\", recursive=True)\n        print(f\"  OK Shard {shard_id}: {len(files)} patches\")\n    else:\n        print(f\"  MANQUANT Shard {shard_id}: {path}\")\n\n# Vérifier Run5\nrun5_count = 0\nfor part in range(1, 16):\n    path = f\"{RUN5_BASE}/run5_part{part}/kaggle/working/patches_panda_run5\"\n    if os.path.exists(path):\n        files = glob.glob(f\"{path}/*.png\")\n        run5_count += len(files)\nprint(f\"  OK Shard 5 (Run5): {run5_count} patches (15 parties)\")\n\n# ─────────────────────────────────────────────\n# 4. CHARGEMENT PATCHES\n# ─────────────────────────────────────────────\nprint(f\"\\n── Patches cumulatifs (Shard 0 à {SHARD_ID}) ──\")\n\ndef load_patches(folder):\n    files = glob.glob(f\"{folder}/*.png\")\n    if not files:\n        files = glob.glob(f\"{folder}/**/*.png\", recursive=True)\n    records = []\n    for fp in files:\n        fname = os.path.basename(fp)\n        try:\n            parts  = fname.split(\"_slide\")\n            label  = int(parts[0].replace(\"grade\", \"\"))\n            wsi_id = parts[1].split(\"_\")[0]\n            records.append({\"path\": fp, \"wsi_id\": wsi_id, \"label\": label})\n        except:\n            continue\n    return records\n\nall_records = []\n\n# Shards 0-4\nfor shard_id, path in SHARD_PATHS.items():\n    if path and os.path.exists(path):\n        r = load_patches(path)\n        all_records.extend(r)\n        print(f\"  Shard {shard_id}: {len(r)} patches\")\n    else:\n        print(f\"  Shard {shard_id}: NON TROUVE\")\n\n# Shard 5 — 15 parties\nrun5_total = 0\nfor part in range(1, 16):\n    path = f\"{RUN5_BASE}/run5_part{part}/kaggle/working/patches_panda_run5\"\n    if os.path.exists(path):\n        r = load_patches(path)\n        all_records.extend(r)\n        run5_total += len(r)\nprint(f\"  Shard 5 (15 parts): {run5_total} patches\")\n\n# Deduplicate\nseen = set()\nunique_records = []\nfor r in all_records:\n    fname = os.path.basename(r[\"path\"])\n    if fname not in seen:\n        seen.add(fname)\n        unique_records.append(r)\n\nall_wsis = list(set(r[\"wsi_id\"] for r in unique_records))\nprint(f\"\\n  Total unique : {len(unique_records)} patches\")\nprint(f\"  Total WSIs  : {len(all_wsis)}\")\ncnt = Counter(r[\"label\"] for r in unique_records)\nfor c in range(6):\n    print(f\"    ISUP {c}: {cnt.get(c, 0)}\")\n\nassert len(all_wsis) > 4500, f\"Trop peu de WSIs ({len(all_wsis)})\"\nprint(f\"\\n  OK {len(all_wsis)} WSIs uniques confirmes\")\n\n# ─────────────────────────────────────────────\n# 5. TRANSFORMS\n# ─────────────────────────────────────────────\nMEAN = [0.485, 0.456, 0.406]\nSTD  = [0.229, 0.224, 0.225]\n\ntrain_tf = T.Compose([\n    T.Resize((224, 224)),\n    T.RandomHorizontalFlip(),\n    T.RandomVerticalFlip(),\n    T.RandomRotation(90),\n    T.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.2, hue=0.1),\n    T.RandomApply([T.GaussianBlur(5)], p=0.3),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n    T.RandomApply([T.RandomErasing(p=0.5, scale=(0.02, 0.2))], p=0.2),\n])\n\nval_tf = T.Compose([\n    T.Resize((224, 224)),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n])\n\nclass PatchDataset(Dataset):\n    def __init__(self, records, transform):\n        self.records   = records\n        self.transform = transform\n    def __len__(self): return len(self.records)\n    def __getitem__(self, idx):\n        r = self.records[idx]\n        try:\n            img = Image.open(r[\"path\"]).convert(\"RGB\")\n        except:\n            img = Image.new(\"RGB\", (224, 224), (255, 255, 255))\n        return self.transform(img), int(r[\"label\"])\n\n# ─────────────────────────────────────────────\n# 6. MODELE\n# ─────────────────────────────────────────────\nclass HybridCNNTransformer(nn.Module):\n    def __init__(self, num_classes=6, d_model=512, nhead=8,\n                 num_layers=2, dim_feedforward=2048, dropout=0.1):\n        super().__init__()\n        backbone = models.resnet50(weights=models.ResNet50_Weights.DEFAULT)\n        self.cnn_backbone = nn.Sequential(*list(backbone.children())[:-2])\n        self.pool         = nn.AdaptiveAvgPool2d((1, 1))\n        self.proj         = nn.Sequential(\n            nn.Identity(), nn.Identity(),\n            nn.Linear(2048, d_model), nn.LayerNorm(d_model)\n        )\n        self.cls_token = nn.Parameter(torch.zeros(1, 1, d_model))\n        self.pos_embed = nn.Parameter(torch.zeros(1, 2, d_model))\n        encoder_layer  = nn.TransformerEncoderLayer(\n            d_model=d_model, nhead=nhead,\n            dim_feedforward=dim_feedforward,\n            dropout=dropout, batch_first=True\n        )\n        self.transformer = nn.TransformerEncoder(\n            encoder_layer, num_layers=num_layers,\n            norm=nn.LayerNorm(d_model)\n        )\n        self.classifier = nn.Sequential(\n            nn.Linear(d_model, 256), nn.ReLU(),\n            nn.Dropout(dropout), nn.Linear(256, num_classes)\n        )\n        nn.init.trunc_normal_(self.cls_token, std=0.02)\n        nn.init.trunc_normal_(self.pos_embed, std=0.02)\n\n    def forward(self, x):\n        f   = self.pool(self.cnn_backbone(x)).flatten(1)\n        f   = self.proj(f).unsqueeze(1)\n        cls = self.cls_token.expand(x.size(0), -1, -1)\n        seq = torch.cat([cls, f], dim=1) + self.pos_embed\n        return self.classifier(self.transformer(seq)[:, 0])\n\ndef build_model(prev_ckpt=None):\n    model  = HybridCNNTransformer().to(DEVICE)\n    n_gpus = torch.cuda.device_count()\n    if n_gpus > 1:\n        model = nn.DataParallel(model)\n        print(f\"  OK DataParallel: {n_gpus} GPUs\")\n    n_param = sum(p.numel() for p in model.parameters())\n    assert n_param == 30_998_342, f\"Mismatch: {n_param}\"\n    print(f\"  OK Parameters: {n_param:,}\")\n    if prev_ckpt and os.path.exists(prev_ckpt):\n        ckpt  = torch.load(prev_ckpt, map_location=DEVICE, weights_only=False)\n        state = ckpt[\"model_state_dict\"]\n        if n_gpus > 1 and not list(state.keys())[0].startswith(\"module.\"):\n            state = {\"module.\" + k: v for k, v in state.items()}\n        elif n_gpus <= 1 and list(state.keys())[0].startswith(\"module.\"):\n            state = {k[7:]: v for k, v in state.items()}\n        model.load_state_dict(state)\n        print(f\"  OK Loaded Stage 4 (QWK={ckpt.get('kappa', '?')})\")\n    return model\n\n# ─────────────────────────────────────────────\n# 7. WSI-LEVEL QWK\n# ─────────────────────────────────────────────\ndef wsi_level_qwk(model, val_data, device):\n    model.eval()\n    wsi_preds  = {}\n    wsi_labels = {}\n    all_preds  = []\n    all_labels = []\n    loader = DataLoader(\n        PatchDataset(val_data, val_tf),\n        batch_size=BATCH_SIZE, shuffle=False,\n        num_workers=NUM_WORKERS\n    )\n    with torch.no_grad():\n        for imgs, labels in loader:\n            imgs  = imgs.to(device)\n            preds = model(imgs).argmax(1).cpu().numpy()\n            all_preds.extend(preds)\n            all_labels.extend(labels.numpy())\n    for i, r in enumerate(val_data):\n        wid = r[\"wsi_id\"]\n        if wid not in wsi_preds:\n            wsi_preds[wid]  = []\n            wsi_labels[wid] = r[\"label\"]\n        wsi_preds[wid].append(all_preds[i])\n    wf = []; wl = []\n    for wid in wsi_preds:\n        votes = wsi_preds[wid]\n        wf.append(max(set(votes), key=votes.count))\n        wl.append(wsi_labels[wid])\n    patch_qwk = cohen_kappa_score(all_labels, all_preds, weights=\"quadratic\")\n    wsi_qwk   = cohen_kappa_score(wl, wf, weights=\"quadratic\")\n    return patch_qwk, wsi_qwk, len(wf)\n\n# ─────────────────────────────────────────────\n# 8. PUSH KAGGLE\n# ─────────────────────────────────────────────\ndef push_to_kaggle(ckpt_path, message):\n    try:\n        tmp = f\"{SAVE_DIR}/_push_tmp\"\n        os.makedirs(tmp, exist_ok=True)\n        for c in glob.glob(f\"{SAVE_DIR}/stage{SHARD_ID}_fold*.pth\"):\n            shutil.copy(c, tmp)\n        shutil.copy(ckpt_path, tmp)\n        if os.path.exists(RESULTS_JSON):\n            shutil.copy(RESULTS_JSON, tmp)\n        json.dump({\n            \"title\": \"Memoire IHC Saved Models\",\n            \"id\": KAGGLE_DS,\n            \"licenses\": [{\"name\": \"CC0-1.0\"}]\n        }, open(f\"{tmp}/dataset-metadata.json\", \"w\"))\n        files_pushed = [os.path.basename(f) for f in glob.glob(f\"{tmp}/*.pth\")]\n        print(f\"   Pushing: {files_pushed}\")\n        r = subprocess.run(\n            [\"kaggle\", \"datasets\", \"version\",\n             \"-p\", tmp, \"-m\", message, \"--dir-mode\", \"skip\"],\n            capture_output=True, text=True, timeout=300\n        )\n        print(f\"   {'OK Kaggle' if r.returncode==0 else 'ERREUR '+r.stderr[:100]}\")\n    except Exception as e:\n        print(f\"   ERREUR push: {e}\")\n\n# ─────────────────────────────────────────────\n# 9. TRAIN ONE FOLD\n# ─────────────────────────────────────────────\ndef train_one_fold(fold_id, train_data, val_data):\n    print(f\"\\n  Fold {fold_id+1}/{N_FOLDS} ──────────────────────────\")\n    print(f\"     Train: {len(train_data)} | Val: {len(val_data)}\")\n\n    train_loader = DataLoader(\n        PatchDataset(train_data, train_tf),\n        batch_size=BATCH_SIZE, shuffle=True,\n        num_workers=NUM_WORKERS, pin_memory=True\n    )\n\n    model = build_model(PREV_CKPT)\n\n    cnt_tr = Counter(r[\"label\"] for r in train_data)\n    w = torch.tensor(\n        [1.0 / max(cnt_tr.get(c, 1), 1) for c in range(6)],\n        dtype=torch.float\n    )\n    w = (w / w.sum() * 6).to(DEVICE)\n    criterion = nn.CrossEntropyLoss(weight=w, label_smoothing=0.1)\n\n    base_model = model.module if hasattr(model, \"module\") else model\n    optimizer  = optim.AdamW([\n        {\"params\": base_model.cnn_backbone.parameters(), \"lr\": LR_CNN},\n        {\"params\": [p for n, p in base_model.named_parameters()\n                    if not n.startswith(\"cnn_backbone\")], \"lr\": LR_TR},\n    ], weight_decay=1e-4)\n\n    scheduler = CosineAnnealingLR(optimizer, T_max=EPOCHS, eta_min=1e-7)\n\n    best_kappa = 0.0\n    no_improve = 0\n    best_ckpt  = f\"{SAVE_DIR}/stage{SHARD_ID}_fold{fold_id+1}_best.pth\"\n\n    for epoch in range(1, EPOCHS + 1):\n        model.train()\n        t_loss, t_correct, t_total = 0.0, 0, 0\n        for imgs, labels in train_loader:\n            imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n            optimizer.zero_grad()\n            out  = model(imgs)\n            loss = criterion(out, labels)\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n            optimizer.step()\n            t_loss    += loss.item() * imgs.size(0)\n            t_correct += (out.argmax(1) == labels).sum().item()\n            t_total   += imgs.size(0)\n        scheduler.step()\n\n        _, wsi_kappa, n_wsis_val = wsi_level_qwk(model, val_data, DEVICE)\n        is_best = wsi_kappa > best_kappa\n\n        print(f\"     Ep[{epoch:02d}/{EPOCHS}] \"\n              f\"loss={t_loss/t_total:.4f} acc={t_correct/t_total:.3f} | \"\n              f\"WSI_QWK={wsi_kappa:.4f} ({n_wsis_val} WSIs)\"\n              + (\" *\" if is_best else \"\"))\n\n        if is_best:\n            best_kappa = wsi_kappa\n            no_improve = 0\n            torch.save({\n                \"epoch\":            epoch,\n                \"model_state_dict\": model.state_dict() if not hasattr(model, \"module\")\n                                    else model.module.state_dict(),\n                \"kappa\":            wsi_kappa,\n                \"shard_id\":         SHARD_ID,\n                \"fold_id\":          fold_id,\n                \"n_wsis\":           len(all_wsis),\n                \"hyperparams\": {\n                    \"lr_cnn\": LR_CNN, \"lr_tr\": LR_TR,\n                    \"epochs\": EPOCHS, \"patience\": PATIENCE,\n                    \"label_smoothing\": 0.1,\n                }\n            }, best_ckpt)\n        else:\n            no_improve += 1\n            if no_improve >= PATIENCE:\n                print(f\"     Early stopping epoch {epoch}\")\n                break\n\n    print(f\"  OK Fold {fold_id+1} best WSI_QWK: {best_kappa:.4f}\")\n    push_to_kaggle(best_ckpt,\n        f\"Stage{SHARD_ID} Fold{fold_id+1} WSI_QWK={best_kappa:.4f}\")\n    return best_kappa, best_ckpt\n\n# ─────────────────────────────────────────────\n# 10. 5-FOLD CV\n# ─────────────────────────────────────────────\nprint(f\"\\n{'='*60}\")\nprint(f\"Stage {SHARD_ID} -- {len(all_wsis)} WSIs -- 5-Fold CV\")\nprint(f\"Fine-tune depuis Stage 4 (QWK=0.7257)\")\nprint(f\"{'='*60}\")\nif DONE_KAPPAS:\n    print(f\"Folds deja completes : {DONE_KAPPAS}\")\n\nwsi_ids    = [r[\"wsi_id\"] for r in unique_records]\nlabels_arr = [r[\"label\"]  for r in unique_records]\nindices    = list(range(len(unique_records)))\n\nsgkf = StratifiedGroupKFold(n_splits=N_FOLDS, shuffle=True, random_state=SEED)\n\nfold_kappas       = DONE_KAPPAS.copy()\nfold_ckpts        = DONE_CKPTS.copy()\nbest_overall      = max(DONE_KAPPAS) if DONE_KAPPAS else -1.0\nbest_ckpt_overall = None\n\nfor fold_id, (train_idx, val_idx) in enumerate(\n        sgkf.split(indices, labels_arr, groups=wsi_ids)):\n\n    if fold_id < START_FOLD:\n        continue\n\n    train_data = [unique_records[i] for i in train_idx]\n    val_data   = [unique_records[i] for i in val_idx]\n\n    tw = set(unique_records[i][\"wsi_id\"] for i in train_idx)\n    vw = set(unique_records[i][\"wsi_id\"] for i in val_idx)\n    assert len(tw & vw) == 0, f\"DATA LEAKAGE fold {fold_id+1}!\"\n\n    kappa, ckpt_path = train_one_fold(fold_id, train_data, val_data)\n    fold_kappas.append(kappa)\n    fold_ckpts.append(ckpt_path)\n\n    if kappa > best_overall:\n        best_overall      = kappa\n        best_ckpt_overall = ckpt_path\n\n# ─────────────────────────────────────────────\n# 11. RESUME FINAL\n# ─────────────────────────────────────────────\nmean_kappa = np.mean(fold_kappas)\nstd_kappa  = np.std(fold_kappas)\n\nprint(f\"\\n{'='*60}\")\nprint(f\"OK Stage {SHARD_ID} -- 5-Fold CV Results\")\nprint(f\"{'='*60}\")\nfor i, k in enumerate(fold_kappas):\n    print(f\"   Fold {i+1}: QWK = {k:.4f}\")\nprint(f\"{'─'*40}\")\nprint(f\"   Mean QWK : {mean_kappa:.4f} +/- {std_kappa:.4f}\")\nprint(f\"   Best QWK : {best_overall:.4f}\")\nprint(f\"   Stage 4  : 0.7112 +/- 0.0130\")\nprint(f\"   Stage 3  : 0.7280 +/- 0.0629\")\nprint(f\"   WSIs     : {len(all_wsis)}\")\nprint(f\"   SOTA [1] : 0.934\")\nprint(f\"{'='*60}\")\n\nresults = {\n    \"stage\":       SHARD_ID,\n    \"n_wsis\":      len(all_wsis),\n    \"n_patches\":   len(unique_records),\n    \"fold_kappas\": fold_kappas,\n    \"mean_kappa\":  mean_kappa,\n    \"std_kappa\":   std_kappa,\n    \"best_kappa\":  best_overall,\n    \"timestamp\":   datetime.datetime.now().isoformat(),\n    \"prev_ckpt\":   PREV_CKPT,\n}\nwith open(RESULTS_JSON, \"w\") as f:\n    json.dump(results, f, indent=2)\n\nif best_ckpt_overall:\n    stage_best = f\"{SAVE_DIR}/stage{SHARD_ID}_best_qwk{best_overall:.4f}.pth\"\n    shutil.copy(best_ckpt_overall, stage_best)\n    print(f\"\\nOK Best checkpoint: {stage_best}\")\n    print(f\"\\n-> Pour Stage 6 :\")\n    print(f\"   SHARD_ID  = 6\")\n    print(f\"   PREV_CKPT = '{stage_best}'\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, glob, shutil, subprocess, json\n\nckpt = \"/kaggle/working/stage5_fold2_best.pth\"\nprint(f\"Fold2: {'OK' if os.path.exists(ckpt) else 'MANQUANT'}\")\n\nif os.path.exists(ckpt):\n    tmp = \"/kaggle/working/_push_s5f2\"\n    os.makedirs(tmp, exist_ok=True)\n    \n    # Copier fold1 et fold2\n    for f in glob.glob(\"/kaggle/working/stage5_fold*.pth\"):\n        shutil.copy(f, tmp)\n        print(f\"  Copied: {os.path.basename(f)}\")\n    \n    json.dump({\n        \"title\": \"Memoire IHC Saved Models\",\n        \"id\": \"khouloudbejaoui20/best-models\",\n        \"licenses\": [{\"name\": \"CC0-1.0\"}]\n    }, open(f\"{tmp}/dataset-metadata.json\", \"w\"))\n    \n    r = subprocess.run([\n        \"kaggle\", \"datasets\", \"version\",\n        \"-p\", tmp,\n        \"-m\", \"Stage5 Fold2 done\",\n        \"--dir-mode\", \"skip\"\n    ], capture_output=True, text=True, timeout=300)\n    \n    print(\"OK!\" if r.returncode == 0 else f\"ERREUR: {r.stderr[:200]}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, glob, shutil, subprocess, json\n\ntmp = \"/kaggle/working/_push_stage5\"\nos.makedirs(tmp, exist_ok=True)\n\n# Copier folds\nfor f in glob.glob(\"/kaggle/working/stage5_fold*.pth\"):\n    shutil.copy(f, tmp)\n    print(f\"  Copied: {os.path.basename(f)}\")\n\njson.dump({\n    \"title\": \"Stage5 Results Bejaoui 2026\",\n    \"id\": \"khouloudbejaoui20/stage5-results-bejaoui-2026\",\n    \"licenses\": [{\"name\": \"CC0-1.0\"}]\n}, open(f\"{tmp}/dataset-metadata.json\", \"w\"))\n\nr = subprocess.run([\n    \"kaggle\", \"datasets\", \"create\",\n    \"-p\", tmp,\n    \"--dir-mode\", \"skip\"\n], capture_output=True, text=True, timeout=600)\n\nprint(\"OK!\" if r.returncode == 0 else f\"ERREUR: {r.stderr[:200]}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import subprocess, json\n\njson.dump({\n    \"title\": \"Stage5 Results Bejaoui 2026\",\n    \"id\": \"khouloudbejaoui20/stage5-results-bejaoui-2026\",\n    \"licenses\": [{\"name\": \"CC0-1.0\"}]\n}, open(\"/kaggle/working/_push_s5f3/dataset-metadata.json\", \"w\"))\n\nr = subprocess.run([\n    \"kaggle\", \"datasets\", \"version\",\n    \"-p\", \"/kaggle/working/_push_s5f3\",\n    \"-m\", \"Stage5 Fold1+2+3\",\n    \"--dir-mode\", \"skip\"\n], capture_output=True, text=True, timeout=600)\n\nprint(\"OK!\" if r.returncode == 0 else f\"ERREUR: {r.stderr[:200]}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import glob, os\n\nfor f in glob.glob(\"/kaggle/working/stage5*.pth\"):\n    mb = os.path.getsize(f)/(1024*1024)\n    print(f\"{os.path.basename(f)} ({mb:.1f} MB)\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import subprocess\nr = subprocess.run([\n    \"kaggle\", \"datasets\", \"files\",\n    \"khouloudbejaoui20/stage5-results-bejaoui-2026\"\n], capture_output=True, text=True)\nprint(r.stdout)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import json, os\n\nkaggle_creds = {\n    \"username\": \"khouloudbejaoui20\",\n    \"key\": \"COLLE_TA_KEY_ICI\"\n}\nos.makedirs(\"/root/.config/kaggle\", exist_ok=True)\nwith open(\"/root/.config/kaggle/kaggle.json\", \"w\") as f:\n    json.dump(kaggle_creds, f)\nos.chmod(\"/root/.config/kaggle/kaggle.json\", 0o600)\n\nimport subprocess\nr = subprocess.run([\n    \"kaggle\", \"datasets\", \"files\",\n    \"khouloudbejaoui20/stage5-results-bejaoui-2026\"\n], capture_output=True, text=True)\nprint(r.stdout)\nprint(r.stderr[:200])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\n=============================================================\nStage 5 — Fold 4+5 ONLY\nFine-tune depuis Stage 4 (best QWK=0.7257)\nAMELIORATION : Cost-Sensitive Ordinal Loss\n→ Pénalise plus les erreurs entre grades éloignés\n→ Réduit la confusion entre grades adjacents\n=============================================================\nAttache dans Kaggle :\n  khouloudbejaoui20/dataset\n  khouloudbejaoui20/panda-patches-run1\n  khouloudbejaoui20/panda-patches-run2\n  khouloudbejaoui20/panda-patches-run3\n  khouloudbejaoui20/panda-patches-run4\n  bejaouikhouloud/panda-patches-run5-parts\n  bejaouikhouloud/stage4-results-bejaoui-2026\n\"\"\"\n\nimport os, glob, json, random, shutil, subprocess, datetime\nimport numpy as np\nfrom PIL import Image\nfrom collections import Counter\nfrom sklearn.model_selection import StratifiedGroupKFold\nfrom sklearn.metrics import cohen_kappa_score\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as T\nimport torchvision.models as models\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\n\n# ─────────────────────────────────────────────\n# 1. CONFIG\n# ─────────────────────────────────────────────\nSHARD_ID    = 5\nSTART_FOLD  = 3    # Fold 4 !\n\nDONE_KAPPAS = [0.7415, 0.6992, 0.7288]\nDONE_CKPTS  = []\n\nPREV_CKPT   = \"/kaggle/working/stage4_fold5_best.pth\"\n\nN_FOLDS     = 5\nSEED        = 42\nBATCH_SIZE  = 32\nNUM_WORKERS = 2\nEPOCHS      = 30\nPATIENCE    = 7\nLR_CNN      = 2e-5\nLR_TR       = 4e-5\n\nrandom.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)\ntorch.backends.cudnn.deterministic = True\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Device    : {DEVICE}\")\nprint(f\"Stage     : {SHARD_ID} Folds 4+5 (Cost-Sensitive Ordinal Loss)\")\nprint(f\"EPOCHS={EPOCHS} | PATIENCE={PATIENCE} | LR_CNN={LR_CNN} | LR_TR={LR_TR}\")\n\n# ─────────────────────────────────────────────\n# 2. COST-SENSITIVE ORDINAL LOSS\n# ─────────────────────────────────────────────\n\nclass OrdinalCostSensitiveLoss(nn.Module):\n    def __init__(self, num_classes=6, class_weights=None,\n                 label_smoothing=0.1, ordinal_weight=0.5):\n        super().__init__()\n        self.num_classes     = num_classes\n        self.label_smoothing = label_smoothing\n        self.ordinal_weight  = ordinal_weight\n\n        cost = torch.zeros(num_classes, num_classes)\n        for i in range(num_classes):\n            for j in range(num_classes):\n                cost[i][j] = abs(i - j)\n        self.register_buffer('cost_matrix', cost)\n\n        self.ce = nn.CrossEntropyLoss(\n            weight=class_weights,\n            label_smoothing=label_smoothing\n        )\n\n    def forward(self, logits, targets):\n        ce_loss = self.ce(logits, targets)\n        probs = F.softmax(logits, dim=1)\n        # Fix: cost_matrix sur même device que targets\n        cost = self.cost_matrix.to(targets.device)\n        target_costs = cost[targets]\n        ordinal_loss = (probs * target_costs).sum(dim=1).mean()\n        return ((1 - self.ordinal_weight) * ce_loss +\n                 self.ordinal_weight * ordinal_loss)\n      \n\n# ─────────────────────────────────────────────\n# 3. DOWNLOAD PREV_CKPT\n# ─────────────────────────────────────────────\nif not os.path.exists(PREV_CKPT):\n    print(\"Downloading stage4_fold5_best.pth...\")\n    # Config bejaouikhouloud\n    kaggle_creds = {\n        \"username\": \"bejaouikhouloud\",\n        \"key\": \"af385e4bd0b6181b85b83134f6f1c405\"\n    }\n    os.makedirs(\"/root/.config/kaggle\", exist_ok=True)\n    with open(\"/root/.config/kaggle/kaggle.json\", \"w\") as f:\n        json.dump(kaggle_creds, f)\n    os.chmod(\"/root/.config/kaggle/kaggle.json\", 0o600)\n\n    subprocess.run([\n        \"kaggle\", \"datasets\", \"download\",\n        \"bejaouikhouloud/stage4-results-bejaoui-2026\",\n        \"--file\", \"stage4_fold5_best.pth\",\n        \"-p\", \"/kaggle/working\", \"--unzip\"\n    ], capture_output=True)\n\nassert os.path.exists(PREV_CKPT), f\"PREV_CKPT non trouve: {PREV_CKPT}\"\nprint(f\"OK Stage 4 checkpoint (QWK=0.7257)\")\n\n# ─────────────────────────────────────────────\n# 4. CHEMINS PATCHES\n# ─────────────────────────────────────────────\nBASE      = \"/kaggle/input/datasets/khouloudbejaoui20\"\nRUN5_BASE = \"/kaggle/input/datasets/bejaouikhouloud/panda-patches-run5-parts\"\n\nSHARD_PATHS = {\n    0: f\"{BASE}/dataset/patches_400-20260630T212405Z-3-001/patches_400\",\n    1: f\"{BASE}/panda-patches-run1/kaggle/working/patches_panda_new\",\n    2: f\"{BASE}/panda-patches-run2/kaggle/working/patches_panda_run2\",\n    3: f\"{BASE}/panda-patches-run3/kaggle/working/patches_panda_run3_correct\",\n    4: f\"{BASE}/panda-patches-run4\",\n}\n\nSAVE_DIR     = \"/kaggle/working\"\nKAGGLE_DS    = \"khouloudbejaoui20/stage5-results-bejaoui-2026\"\nRESULTS_JSON = f\"{SAVE_DIR}/stage5_cv_results.json\"\n\n# ─────────────────────────────────────────────\n# 5. LOAD PATCHES\n# ─────────────────────────────────────────────\nprint(f\"\\n── Loading patches ──────────────────────\")\n\ndef load_patches(folder):\n    files = glob.glob(f\"{folder}/*.png\")\n    if not files:\n        files = glob.glob(f\"{folder}/**/*.png\", recursive=True)\n    records = []\n    for fp in files:\n        fname = os.path.basename(fp)\n        try:\n            parts  = fname.split(\"_slide\")\n            label  = int(parts[0].replace(\"grade\", \"\"))\n            wsi_id = parts[1].split(\"_\")[0]\n            records.append({\"path\": fp, \"wsi_id\": wsi_id, \"label\": label})\n        except:\n            continue\n    return records\n\nall_records = []\nfor shard_id, path in SHARD_PATHS.items():\n    if path and os.path.exists(path):\n        r = load_patches(path)\n        all_records.extend(r)\n        print(f\"  Shard {shard_id}: {len(r)} patches\")\n    else:\n        print(f\"  Shard {shard_id}: NON TROUVE\")\n\n# Shard 5 — 15 parties\nrun5_total = 0\nfor part in range(1, 16):\n    path = f\"{RUN5_BASE}/run5_part{part}/kaggle/working/patches_panda_run5\"\n    if os.path.exists(path):\n        r = load_patches(path)\n        all_records.extend(r)\n        run5_total += len(r)\nprint(f\"  Shard 5 (15 parts): {run5_total} patches\")\n\n# Deduplicate\nseen = set()\nunique_records = []\nfor r in all_records:\n    fname = os.path.basename(r[\"path\"])\n    if fname not in seen:\n        seen.add(fname)\n        unique_records.append(r)\n\nall_wsis = list(set(r[\"wsi_id\"] for r in unique_records))\nprint(f\"\\n  Total unique : {len(unique_records)} patches\")\nprint(f\"  Total WSIs  : {len(all_wsis)}\")\ncnt = Counter(r[\"label\"] for r in unique_records)\nfor c in range(6):\n    print(f\"    ISUP {c}: {cnt.get(c, 0)}\")\n\nassert len(all_wsis) > 4500, f\"Trop peu de WSIs ({len(all_wsis)})\"\nprint(f\"\\n  OK {len(all_wsis)} WSIs confirmes\")\n\n# ─────────────────────────────────────────────\n# 6. TRANSFORMS\n# ─────────────────────────────────────────────\nMEAN = [0.485, 0.456, 0.406]\nSTD  = [0.229, 0.224, 0.225]\n\ntrain_tf = T.Compose([\n    T.Resize((224, 224)),\n    T.RandomHorizontalFlip(),\n    T.RandomVerticalFlip(),\n    T.RandomRotation(90),\n    T.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.2, hue=0.1),\n    T.RandomApply([T.GaussianBlur(5)], p=0.3),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n    T.RandomApply([T.RandomErasing(p=0.5, scale=(0.02, 0.2))], p=0.2),\n])\n\nval_tf = T.Compose([\n    T.Resize((224, 224)),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n])\n\nclass PatchDataset(Dataset):\n    def __init__(self, records, transform):\n        self.records   = records\n        self.transform = transform\n    def __len__(self): return len(self.records)\n    def __getitem__(self, idx):\n        r = self.records[idx]\n        try:\n            img = Image.open(r[\"path\"]).convert(\"RGB\")\n        except:\n            img = Image.new(\"RGB\", (224, 224), (255, 255, 255))\n        return self.transform(img), int(r[\"label\"])\n\n# ─────────────────────────────────────────────\n# 7. MODEL\n# ─────────────────────────────────────────────\nclass HybridCNNTransformer(nn.Module):\n    def __init__(self, num_classes=6, d_model=512, nhead=8,\n                 num_layers=2, dim_feedforward=2048, dropout=0.1):\n        super().__init__()\n        backbone = models.resnet50(weights=models.ResNet50_Weights.DEFAULT)\n        self.cnn_backbone = nn.Sequential(*list(backbone.children())[:-2])\n        self.pool         = nn.AdaptiveAvgPool2d((1, 1))\n        self.proj         = nn.Sequential(\n            nn.Identity(), nn.Identity(),\n            nn.Linear(2048, d_model), nn.LayerNorm(d_model)\n        )\n        self.cls_token = nn.Parameter(torch.zeros(1, 1, d_model))\n        self.pos_embed = nn.Parameter(torch.zeros(1, 2, d_model))\n        encoder_layer  = nn.TransformerEncoderLayer(\n            d_model=d_model, nhead=nhead,\n            dim_feedforward=dim_feedforward,\n            dropout=dropout, batch_first=True\n        )\n        self.transformer = nn.TransformerEncoder(\n            encoder_layer, num_layers=num_layers,\n            norm=nn.LayerNorm(d_model)\n        )\n        self.classifier = nn.Sequential(\n            nn.Linear(d_model, 256), nn.ReLU(),\n            nn.Dropout(dropout), nn.Linear(256, num_classes)\n        )\n        nn.init.trunc_normal_(self.cls_token, std=0.02)\n        nn.init.trunc_normal_(self.pos_embed, std=0.02)\n\n    def forward(self, x):\n        f   = self.pool(self.cnn_backbone(x)).flatten(1)\n        f   = self.proj(f).unsqueeze(1)\n        cls = self.cls_token.expand(x.size(0), -1, -1)\n        seq = torch.cat([cls, f], dim=1) + self.pos_embed\n        return self.classifier(self.transformer(seq)[:, 0])\n\ndef build_model():\n    model  = HybridCNNTransformer().to(DEVICE)\n    n_gpus = torch.cuda.device_count()\n    if n_gpus > 1:\n        model = nn.DataParallel(model)\n        print(f\"  OK DataParallel: {n_gpus} GPUs\")\n    n_param = sum(p.numel() for p in model.parameters())\n    assert n_param == 30_998_342, f\"Mismatch: {n_param}\"\n    print(f\"  OK Parameters: {n_param:,}\")\n    ckpt  = torch.load(PREV_CKPT, map_location=DEVICE, weights_only=False)\n    state = ckpt[\"model_state_dict\"]\n    if n_gpus > 1 and not list(state.keys())[0].startswith(\"module.\"):\n        state = {\"module.\" + k: v for k, v in state.items()}\n    elif n_gpus <= 1 and list(state.keys())[0].startswith(\"module.\"):\n        state = {k[7:]: v for k, v in state.items()}\n    model.load_state_dict(state)\n    print(f\"  OK Loaded Stage 4 (QWK={ckpt.get('kappa', '?')})\")\n    return model\n\n# ─────────────────────────────────────────────\n# 8. WSI-LEVEL QWK\n# ─────────────────────────────────────────────\ndef wsi_level_qwk(model, val_data, device):\n    model.eval()\n    wsi_preds  = {}\n    wsi_labels = {}\n    all_preds  = []\n    all_labels = []\n    loader = DataLoader(\n        PatchDataset(val_data, val_tf),\n        batch_size=BATCH_SIZE, shuffle=False,\n        num_workers=NUM_WORKERS\n    )\n    with torch.no_grad():\n        for imgs, labels in loader:\n            imgs  = imgs.to(device)\n            preds = model(imgs).argmax(1).cpu().numpy()\n            all_preds.extend(preds)\n            all_labels.extend(labels.numpy())\n    for i, r in enumerate(val_data):\n        wid = r[\"wsi_id\"]\n        if wid not in wsi_preds:\n            wsi_preds[wid]  = []\n            wsi_labels[wid] = r[\"label\"]\n        wsi_preds[wid].append(all_preds[i])\n    wf = []; wl = []\n    for wid in wsi_preds:\n        votes = wsi_preds[wid]\n        wf.append(max(set(votes), key=votes.count))\n        wl.append(wsi_labels[wid])\n    wsi_qwk = cohen_kappa_score(wl, wf, weights=\"quadratic\")\n    return wsi_qwk, len(wf)\n\n# ─────────────────────────────────────────────\n# 9. PUSH KAGGLE\n# ─────────────────────────────────────────────\ndef push_to_kaggle(ckpt_path, message):\n    try:\n        tmp = f\"{SAVE_DIR}/_push_tmp\"\n        os.makedirs(tmp, exist_ok=True)\n        for c in glob.glob(f\"{SAVE_DIR}/stage{SHARD_ID}_fold*.pth\"):\n            shutil.copy(c, tmp)\n        shutil.copy(ckpt_path, tmp)\n        json.dump({\n            \"title\": \"Stage5 Results Bejaoui 2026\",\n            \"id\": KAGGLE_DS,\n            \"licenses\": [{\"name\": \"CC0-1.0\"}]\n        }, open(f\"{tmp}/dataset-metadata.json\", \"w\"))\n        files_pushed = [os.path.basename(f) for f in glob.glob(f\"{tmp}/*.pth\")]\n        print(f\"   Pushing: {files_pushed}\")\n        r = subprocess.run(\n            [\"kaggle\", \"datasets\", \"version\",\n             \"-p\", tmp, \"-m\", message, \"--dir-mode\", \"skip\"],\n            capture_output=True, text=True, timeout=300\n        )\n        print(f\"   {'OK!' if r.returncode==0 else 'ERREUR '+r.stderr[:100]}\")\n    except Exception as e:\n        print(f\"   ERREUR: {e}\")\n\n# ─────────────────────────────────────────────\n# 10. TRAIN ONE FOLD\n# ─────────────────────────────────────────────\ndef train_one_fold(fold_id, train_data, val_data):\n    print(f\"\\n  Fold {fold_id+1}/{N_FOLDS} ──────────────────────────\")\n    print(f\"     Train: {len(train_data)} | Val: {len(val_data)}\")\n    print(f\"     Loss: OrdinalCostSensitive (ordinal_weight=0.5)\")\n\n    train_loader = DataLoader(\n        PatchDataset(train_data, train_tf),\n        batch_size=BATCH_SIZE, shuffle=True,\n        num_workers=NUM_WORKERS, pin_memory=True\n    )\n\n    model = build_model()\n\n    # Class weights\n    cnt_tr = Counter(r[\"label\"] for r in train_data)\n    w = torch.tensor(\n        [1.0 / max(cnt_tr.get(c, 1), 1) for c in range(6)],\n        dtype=torch.float\n    )\n    w = (w / w.sum() * 6).to(DEVICE)\n\n    # Ordinal Cost-Sensitive Loss\n    criterion = OrdinalCostSensitiveLoss(\n        num_classes=6,\n        class_weights=w,\n        label_smoothing=0.1,\n        ordinal_weight=0.5\n    )\n\n    base_model = model.module if hasattr(model, \"module\") else model\n    optimizer  = optim.AdamW([\n        {\"params\": base_model.cnn_backbone.parameters(), \"lr\": LR_CNN},\n        {\"params\": [p for n, p in base_model.named_parameters()\n                    if not n.startswith(\"cnn_backbone\")], \"lr\": LR_TR},\n    ], weight_decay=1e-4)\n\n    scheduler = CosineAnnealingLR(optimizer, T_max=EPOCHS, eta_min=1e-7)\n\n    best_kappa = 0.0\n    no_improve = 0\n    best_ckpt  = f\"{SAVE_DIR}/stage{SHARD_ID}_fold{fold_id+1}_best.pth\"\n\n    for epoch in range(1, EPOCHS + 1):\n        model.train()\n        t_loss, t_correct, t_total = 0.0, 0, 0\n        for imgs, labels in train_loader:\n            imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n            optimizer.zero_grad()\n            out  = model(imgs)\n            loss = criterion(out, labels)\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n            optimizer.step()\n            t_loss    += loss.item() * imgs.size(0)\n            t_correct += (out.argmax(1) == labels).sum().item()\n            t_total   += imgs.size(0)\n        scheduler.step()\n\n        wsi_kappa, n_wsis_val = wsi_level_qwk(model, val_data, DEVICE)\n        is_best = wsi_kappa > best_kappa\n\n        print(f\"     Ep[{epoch:02d}/{EPOCHS}] \"\n              f\"loss={t_loss/t_total:.4f} acc={t_correct/t_total:.3f} | \"\n              f\"WSI_QWK={wsi_kappa:.4f} ({n_wsis_val} WSIs)\"\n              + (\" *\" if is_best else \"\"))\n\n        if is_best:\n            best_kappa = wsi_kappa\n            no_improve = 0\n            torch.save({\n                \"epoch\":            epoch,\n                \"model_state_dict\": model.state_dict() if not hasattr(model, \"module\")\n                                    else model.module.state_dict(),\n                \"kappa\":            wsi_kappa,\n                \"shard_id\":         SHARD_ID,\n                \"fold_id\":          fold_id,\n                \"n_wsis\":           len(all_wsis),\n                \"loss_type\":        \"OrdinalCostSensitive\",\n            }, best_ckpt)\n        else:\n            no_improve += 1\n            if no_improve >= PATIENCE:\n                print(f\"     Early stopping epoch {epoch}\")\n                break\n\n    print(f\"  OK Fold {fold_id+1} best WSI_QWK: {best_kappa:.4f}\")\n    push_to_kaggle(best_ckpt,\n        f\"Stage{SHARD_ID} Fold{fold_id+1} OrdinalLoss QWK={best_kappa:.4f}\")\n    return best_kappa, best_ckpt\n\n# ─────────────────────────────────────────────\n# 11. 5-FOLD CV\n# ─────────────────────────────────────────────\nprint(f\"\\n{'='*60}\")\nprint(f\"Stage {SHARD_ID} Folds 4+5 — {len(all_wsis)} WSIs\")\nprint(f\"Loss: OrdinalCostSensitive (ordinal_weight=0.5)\")\nprint(f\"Folds deja completes : {DONE_KAPPAS}\")\nprint(f\"{'='*60}\")\n\nwsi_ids    = [r[\"wsi_id\"] for r in unique_records]\nlabels_arr = [r[\"label\"]  for r in unique_records]\nindices    = list(range(len(unique_records)))\n\nsgkf = StratifiedGroupKFold(n_splits=N_FOLDS, shuffle=True, random_state=SEED)\n\nfold_kappas       = DONE_KAPPAS.copy()\nfold_ckpts        = DONE_CKPTS.copy()\nbest_overall      = max(DONE_KAPPAS) if DONE_KAPPAS else -1.0\nbest_ckpt_overall = None\n\nfor fold_id, (train_idx, val_idx) in enumerate(\n        sgkf.split(indices, labels_arr, groups=wsi_ids)):\n\n    if fold_id < START_FOLD:\n        continue\n\n    train_data = [unique_records[i] for i in train_idx]\n    val_data   = [unique_records[i] for i in val_idx]\n\n    tw = set(unique_records[i][\"wsi_id\"] for i in train_idx)\n    vw = set(unique_records[i][\"wsi_id\"] for i in val_idx)\n    assert len(tw & vw) == 0, f\"DATA LEAKAGE fold {fold_id+1}!\"\n\n    kappa, ckpt_path = train_one_fold(fold_id, train_data, val_data)\n    fold_kappas.append(kappa)\n    fold_ckpts.append(ckpt_path)\n\n    if kappa > best_overall:\n        best_overall      = kappa\n        best_ckpt_overall = ckpt_path\n\n# ─────────────────────────────────────────────\n# 12. RESUME FINAL\n# ─────────────────────────────────────────────\nmean_kappa = np.mean(fold_kappas)\nstd_kappa  = np.std(fold_kappas)\n\nprint(f\"\\n{'='*60}\")\nprint(f\"OK Stage {SHARD_ID} — 5-Fold CV Results\")\nprint(f\"{'='*60}\")\nfor i, k in enumerate(fold_kappas):\n    print(f\"   Fold {i+1}: QWK = {k:.4f}\")\nprint(f\"{'─'*40}\")\nprint(f\"   Mean QWK : {mean_kappa:.4f} +/- {std_kappa:.4f}\")\nprint(f\"   Best QWK : {best_overall:.4f}\")\nprint(f\"   Stage 4  : 0.7112 +/- 0.0130\")\nprint(f\"   Stage 3  : 0.7280 +/- 0.0629\")\nprint(f\"   SOTA     : 0.934\")\nprint(f\"{'='*60}\")\n\nif best_ckpt_overall:\n    stage_best = f\"{SAVE_DIR}/stage{SHARD_ID}_best_qwk{best_overall:.4f}.pth\"\n    shutil.copy(best_ckpt_overall, stage_best)\n    print(f\"\\nOK Best checkpoint: {stage_best}\")\n    print(f\"\\n→ Pour Stage 6 :\")\n    print(f\"   SHARD_ID  = 6\")\n    print(f\"   PREV_CKPT = '{stage_best}'\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\n=============================================================\nStage 5 — Fold 4+5 ONLY\nFine-tune depuis Stage 4 (best QWK=0.7257)\nAMELIORATION : Cost-Sensitive Ordinal Loss\n→ Pénalise plus les erreurs entre grades éloignés\n→ Réduit la confusion entre grades adjacents\n=============================================================\nAttache dans Kaggle :\n  khouloudbejaoui20/dataset\n  khouloudbejaoui20/panda-patches-run1\n  khouloudbejaoui20/panda-patches-run2\n  khouloudbejaoui20/panda-patches-run3\n  khouloudbejaoui20/panda-patches-run4\n  bejaouikhouloud/panda-patches-run5-parts\n  bejaouikhouloud/stage4-results-bejaoui-2026\n\"\"\"\n\nimport os, glob, json, random, shutil, subprocess, datetime\nimport numpy as np\nfrom PIL import Image\nfrom collections import Counter\nfrom sklearn.model_selection import StratifiedGroupKFold\nfrom sklearn.metrics import cohen_kappa_score\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as T\nimport torchvision.models as models\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\n\n# ─────────────────────────────────────────────\n# 1. CONFIG\n# ─────────────────────────────────────────────\nSHARD_ID    = 5\nSTART_FOLD  = 4    # Fold 5 !\n\nDONE_KAPPAS = [0.7415, 0.6992, 0.7288, 0.7054]\nDONE_CKPTS  = []\n\nPREV_CKPT   = \"/kaggle/working/stage4_fold5_best.pth\"\n\nN_FOLDS     = 5\nSEED        = 42\nBATCH_SIZE  = 32\nNUM_WORKERS = 2\nEPOCHS      = 30\nPATIENCE    = 7\nLR_CNN      = 2e-5\nLR_TR       = 4e-5\n\nrandom.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)\ntorch.backends.cudnn.deterministic = True\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Device    : {DEVICE}\")\nprint(f\"Stage     : {SHARD_ID} Folds 4+5 (Cost-Sensitive Ordinal Loss)\")\nprint(f\"EPOCHS={EPOCHS} | PATIENCE={PATIENCE} | LR_CNN={LR_CNN} | LR_TR={LR_TR}\")\n\n# ─────────────────────────────────────────────\n# 2. COST-SENSITIVE ORDINAL LOSS\n# ─────────────────────────────────────────────\n\nclass OrdinalCostSensitiveLoss(nn.Module):\n    def __init__(self, num_classes=6, class_weights=None,\n                 label_smoothing=0.1, ordinal_weight=0.5):\n        super().__init__()\n        self.num_classes     = num_classes\n        self.label_smoothing = label_smoothing\n        self.ordinal_weight  = ordinal_weight\n\n        cost = torch.zeros(num_classes, num_classes)\n        for i in range(num_classes):\n            for j in range(num_classes):\n                cost[i][j] = abs(i - j)\n        self.register_buffer('cost_matrix', cost)\n\n        self.ce = nn.CrossEntropyLoss(\n            weight=class_weights,\n            label_smoothing=label_smoothing\n        )\n\n    def forward(self, logits, targets):\n        ce_loss = self.ce(logits, targets)\n        probs = F.softmax(logits, dim=1)\n        # Fix: cost_matrix sur même device que targets\n        cost = self.cost_matrix.to(targets.device)\n        target_costs = cost[targets]\n        ordinal_loss = (probs * target_costs).sum(dim=1).mean()\n        return ((1 - self.ordinal_weight) * ce_loss +\n                 self.ordinal_weight * ordinal_loss)\n      \n\n# ─────────────────────────────────────────────\n# 3. DOWNLOAD PREV_CKPT\n# ─────────────────────────────────────────────\nif not os.path.exists(PREV_CKPT):\n    print(\"Downloading stage4_fold5_best.pth...\")\n    # Config bejaouikhouloud\n    kaggle_creds = {\n        \"username\": \"bejaouikhouloud\",\n        \"key\": \"af385e4bd0b6181b85b83134f6f1c405\"\n    }\n    os.makedirs(\"/root/.config/kaggle\", exist_ok=True)\n    with open(\"/root/.config/kaggle/kaggle.json\", \"w\") as f:\n        json.dump(kaggle_creds, f)\n    os.chmod(\"/root/.config/kaggle/kaggle.json\", 0o600)\n\n    subprocess.run([\n        \"kaggle\", \"datasets\", \"download\",\n        \"bejaouikhouloud/stage4-results-bejaoui-2026\",\n        \"--file\", \"stage4_fold5_best.pth\",\n        \"-p\", \"/kaggle/working\", \"--unzip\"\n    ], capture_output=True)\n\nassert os.path.exists(PREV_CKPT), f\"PREV_CKPT non trouve: {PREV_CKPT}\"\nprint(f\"OK Stage 4 checkpoint (QWK=0.7257)\")\n\n# ─────────────────────────────────────────────\n# 4. CHEMINS PATCHES\n# ─────────────────────────────────────────────\nBASE      = \"/kaggle/input/datasets/khouloudbejaoui20\"\nRUN5_BASE = \"/kaggle/input/datasets/bejaouikhouloud/panda-patches-run5-parts\"\n\nSHARD_PATHS = {\n    0: f\"{BASE}/dataset/patches_400-20260630T212405Z-3-001/patches_400\",\n    1: f\"{BASE}/panda-patches-run1/kaggle/working/patches_panda_new\",\n    2: f\"{BASE}/panda-patches-run2/kaggle/working/patches_panda_run2\",\n    3: f\"{BASE}/panda-patches-run3/kaggle/working/patches_panda_run3_correct\",\n    4: f\"{BASE}/panda-patches-run4\",\n}\n\nSAVE_DIR     = \"/kaggle/working\"\nKAGGLE_DS    = \"khouloudbejaoui20/stage5-results-bejaoui-2026\"\nRESULTS_JSON = f\"{SAVE_DIR}/stage5_cv_results.json\"\n\n# ─────────────────────────────────────────────\n# 5. LOAD PATCHES\n# ─────────────────────────────────────────────\nprint(f\"\\n── Loading patches ──────────────────────\")\n\ndef load_patches(folder):\n    files = glob.glob(f\"{folder}/*.png\")\n    if not files:\n        files = glob.glob(f\"{folder}/**/*.png\", recursive=True)\n    records = []\n    for fp in files:\n        fname = os.path.basename(fp)\n        try:\n            parts  = fname.split(\"_slide\")\n            label  = int(parts[0].replace(\"grade\", \"\"))\n            wsi_id = parts[1].split(\"_\")[0]\n            records.append({\"path\": fp, \"wsi_id\": wsi_id, \"label\": label})\n        except:\n            continue\n    return records\n\nall_records = []\nfor shard_id, path in SHARD_PATHS.items():\n    if path and os.path.exists(path):\n        r = load_patches(path)\n        all_records.extend(r)\n        print(f\"  Shard {shard_id}: {len(r)} patches\")\n    else:\n        print(f\"  Shard {shard_id}: NON TROUVE\")\n\n# Shard 5 — 15 parties\nrun5_total = 0\nfor part in range(1, 16):\n    path = f\"{RUN5_BASE}/run5_part{part}/kaggle/working/patches_panda_run5\"\n    if os.path.exists(path):\n        r = load_patches(path)\n        all_records.extend(r)\n        run5_total += len(r)\nprint(f\"  Shard 5 (15 parts): {run5_total} patches\")\n\n# Deduplicate\nseen = set()\nunique_records = []\nfor r in all_records:\n    fname = os.path.basename(r[\"path\"])\n    if fname not in seen:\n        seen.add(fname)\n        unique_records.append(r)\n\nall_wsis = list(set(r[\"wsi_id\"] for r in unique_records))\nprint(f\"\\n  Total unique : {len(unique_records)} patches\")\nprint(f\"  Total WSIs  : {len(all_wsis)}\")\ncnt = Counter(r[\"label\"] for r in unique_records)\nfor c in range(6):\n    print(f\"    ISUP {c}: {cnt.get(c, 0)}\")\n\nassert len(all_wsis) > 4500, f\"Trop peu de WSIs ({len(all_wsis)})\"\nprint(f\"\\n  OK {len(all_wsis)} WSIs confirmes\")\n\n# ─────────────────────────────────────────────\n# 6. TRANSFORMS\n# ─────────────────────────────────────────────\nMEAN = [0.485, 0.456, 0.406]\nSTD  = [0.229, 0.224, 0.225]\n\ntrain_tf = T.Compose([\n    T.Resize((224, 224)),\n    T.RandomHorizontalFlip(),\n    T.RandomVerticalFlip(),\n    T.RandomRotation(90),\n    T.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.2, hue=0.1),\n    T.RandomApply([T.GaussianBlur(5)], p=0.3),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n    T.RandomApply([T.RandomErasing(p=0.5, scale=(0.02, 0.2))], p=0.2),\n])\n\nval_tf = T.Compose([\n    T.Resize((224, 224)),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n])\n\nclass PatchDataset(Dataset):\n    def __init__(self, records, transform):\n        self.records   = records\n        self.transform = transform\n    def __len__(self): return len(self.records)\n    def __getitem__(self, idx):\n        r = self.records[idx]\n        try:\n            img = Image.open(r[\"path\"]).convert(\"RGB\")\n        except:\n            img = Image.new(\"RGB\", (224, 224), (255, 255, 255))\n        return self.transform(img), int(r[\"label\"])\n\n# ─────────────────────────────────────────────\n# 7. MODEL\n# ─────────────────────────────────────────────\nclass HybridCNNTransformer(nn.Module):\n    def __init__(self, num_classes=6, d_model=512, nhead=8,\n                 num_layers=2, dim_feedforward=2048, dropout=0.1):\n        super().__init__()\n        backbone = models.resnet50(weights=models.ResNet50_Weights.DEFAULT)\n        self.cnn_backbone = nn.Sequential(*list(backbone.children())[:-2])\n        self.pool         = nn.AdaptiveAvgPool2d((1, 1))\n        self.proj         = nn.Sequential(\n            nn.Identity(), nn.Identity(),\n            nn.Linear(2048, d_model), nn.LayerNorm(d_model)\n        )\n        self.cls_token = nn.Parameter(torch.zeros(1, 1, d_model))\n        self.pos_embed = nn.Parameter(torch.zeros(1, 2, d_model))\n        encoder_layer  = nn.TransformerEncoderLayer(\n            d_model=d_model, nhead=nhead,\n            dim_feedforward=dim_feedforward,\n            dropout=dropout, batch_first=True\n        )\n        self.transformer = nn.TransformerEncoder(\n            encoder_layer, num_layers=num_layers,\n            norm=nn.LayerNorm(d_model)\n        )\n        self.classifier = nn.Sequential(\n            nn.Linear(d_model, 256), nn.ReLU(),\n            nn.Dropout(dropout), nn.Linear(256, num_classes)\n        )\n        nn.init.trunc_normal_(self.cls_token, std=0.02)\n        nn.init.trunc_normal_(self.pos_embed, std=0.02)\n\n    def forward(self, x):\n        f   = self.pool(self.cnn_backbone(x)).flatten(1)\n        f   = self.proj(f).unsqueeze(1)\n        cls = self.cls_token.expand(x.size(0), -1, -1)\n        seq = torch.cat([cls, f], dim=1) + self.pos_embed\n        return self.classifier(self.transformer(seq)[:, 0])\n\ndef build_model():\n    model  = HybridCNNTransformer().to(DEVICE)\n    n_gpus = torch.cuda.device_count()\n    if n_gpus > 1:\n        model = nn.DataParallel(model)\n        print(f\"  OK DataParallel: {n_gpus} GPUs\")\n    n_param = sum(p.numel() for p in model.parameters())\n    assert n_param == 30_998_342, f\"Mismatch: {n_param}\"\n    print(f\"  OK Parameters: {n_param:,}\")\n    ckpt  = torch.load(PREV_CKPT, map_location=DEVICE, weights_only=False)\n    state = ckpt[\"model_state_dict\"]\n    if n_gpus > 1 and not list(state.keys())[0].startswith(\"module.\"):\n        state = {\"module.\" + k: v for k, v in state.items()}\n    elif n_gpus <= 1 and list(state.keys())[0].startswith(\"module.\"):\n        state = {k[7:]: v for k, v in state.items()}\n    model.load_state_dict(state)\n    print(f\"  OK Loaded Stage 4 (QWK={ckpt.get('kappa', '?')})\")\n    return model\n\n# ─────────────────────────────────────────────\n# 8. WSI-LEVEL QWK\n# ─────────────────────────────────────────────\ndef wsi_level_qwk(model, val_data, device):\n    model.eval()\n    wsi_preds  = {}\n    wsi_labels = {}\n    all_preds  = []\n    all_labels = []\n    loader = DataLoader(\n        PatchDataset(val_data, val_tf),\n        batch_size=BATCH_SIZE, shuffle=False,\n        num_workers=NUM_WORKERS\n    )\n    with torch.no_grad():\n        for imgs, labels in loader:\n            imgs  = imgs.to(device)\n            preds = model(imgs).argmax(1).cpu().numpy()\n            all_preds.extend(preds)\n            all_labels.extend(labels.numpy())\n    for i, r in enumerate(val_data):\n        wid = r[\"wsi_id\"]\n        if wid not in wsi_preds:\n            wsi_preds[wid]  = []\n            wsi_labels[wid] = r[\"label\"]\n        wsi_preds[wid].append(all_preds[i])\n    wf = []; wl = []\n    for wid in wsi_preds:\n        votes = wsi_preds[wid]\n        wf.append(max(set(votes), key=votes.count))\n        wl.append(wsi_labels[wid])\n    wsi_qwk = cohen_kappa_score(wl, wf, weights=\"quadratic\")\n    return wsi_qwk, len(wf)\n\n# ─────────────────────────────────────────────\n# 9. PUSH KAGGLE\n# ─────────────────────────────────────────────\ndef push_to_kaggle(ckpt_path, message):\n    try:\n        tmp = f\"{SAVE_DIR}/_push_tmp\"\n        os.makedirs(tmp, exist_ok=True)\n        for c in glob.glob(f\"{SAVE_DIR}/stage{SHARD_ID}_fold*.pth\"):\n            shutil.copy(c, tmp)\n        shutil.copy(ckpt_path, tmp)\n        json.dump({\n            \"title\": \"Stage5 Results Bejaoui 2026\",\n            \"id\": KAGGLE_DS,\n            \"licenses\": [{\"name\": \"CC0-1.0\"}]\n        }, open(f\"{tmp}/dataset-metadata.json\", \"w\"))\n        files_pushed = [os.path.basename(f) for f in glob.glob(f\"{tmp}/*.pth\")]\n        print(f\"   Pushing: {files_pushed}\")\n        r = subprocess.run(\n            [\"kaggle\", \"datasets\", \"version\",\n             \"-p\", tmp, \"-m\", message, \"--dir-mode\", \"skip\"],\n            capture_output=True, text=True, timeout=300\n        )\n        print(f\"   {'OK!' if r.returncode==0 else 'ERREUR '+r.stderr[:100]}\")\n    except Exception as e:\n        print(f\"   ERREUR: {e}\")\n\n# ─────────────────────────────────────────────\n# 10. TRAIN ONE FOLD\n# ─────────────────────────────────────────────\ndef train_one_fold(fold_id, train_data, val_data):\n    print(f\"\\n  Fold {fold_id+1}/{N_FOLDS} ──────────────────────────\")\n    print(f\"     Train: {len(train_data)} | Val: {len(val_data)}\")\n    print(f\"     Loss: OrdinalCostSensitive (ordinal_weight=0.5)\")\n\n    train_loader = DataLoader(\n        PatchDataset(train_data, train_tf),\n        batch_size=BATCH_SIZE, shuffle=True,\n        num_workers=NUM_WORKERS, pin_memory=True\n    )\n\n    model = build_model()\n\n    # Class weights\n    cnt_tr = Counter(r[\"label\"] for r in train_data)\n    w = torch.tensor(\n        [1.0 / max(cnt_tr.get(c, 1), 1) for c in range(6)],\n        dtype=torch.float\n    )\n    w = (w / w.sum() * 6).to(DEVICE)\n\n    # Ordinal Cost-Sensitive Loss\n    criterion = OrdinalCostSensitiveLoss(\n        num_classes=6,\n        class_weights=w,\n        label_smoothing=0.1,\n        ordinal_weight=0.5\n    )\n\n    base_model = model.module if hasattr(model, \"module\") else model\n    optimizer  = optim.AdamW([\n        {\"params\": base_model.cnn_backbone.parameters(), \"lr\": LR_CNN},\n        {\"params\": [p for n, p in base_model.named_parameters()\n                    if not n.startswith(\"cnn_backbone\")], \"lr\": LR_TR},\n    ], weight_decay=1e-4)\n\n    scheduler = CosineAnnealingLR(optimizer, T_max=EPOCHS, eta_min=1e-7)\n\n    best_kappa = 0.0\n    no_improve = 0\n    best_ckpt  = f\"{SAVE_DIR}/stage{SHARD_ID}_fold{fold_id+1}_best.pth\"\n\n    for epoch in range(1, EPOCHS + 1):\n        model.train()\n        t_loss, t_correct, t_total = 0.0, 0, 0\n        for imgs, labels in train_loader:\n            imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n            optimizer.zero_grad()\n            out  = model(imgs)\n            loss = criterion(out, labels)\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n            optimizer.step()\n            t_loss    += loss.item() * imgs.size(0)\n            t_correct += (out.argmax(1) == labels).sum().item()\n            t_total   += imgs.size(0)\n        scheduler.step()\n\n        wsi_kappa, n_wsis_val = wsi_level_qwk(model, val_data, DEVICE)\n        is_best = wsi_kappa > best_kappa\n\n        print(f\"     Ep[{epoch:02d}/{EPOCHS}] \"\n              f\"loss={t_loss/t_total:.4f} acc={t_correct/t_total:.3f} | \"\n              f\"WSI_QWK={wsi_kappa:.4f} ({n_wsis_val} WSIs)\"\n              + (\" *\" if is_best else \"\"))\n\n        if is_best:\n            best_kappa = wsi_kappa\n            no_improve = 0\n            torch.save({\n                \"epoch\":            epoch,\n                \"model_state_dict\": model.state_dict() if not hasattr(model, \"module\")\n                                    else model.module.state_dict(),\n                \"kappa\":            wsi_kappa,\n                \"shard_id\":         SHARD_ID,\n                \"fold_id\":          fold_id,\n                \"n_wsis\":           len(all_wsis),\n                \"loss_type\":        \"OrdinalCostSensitive\",\n            }, best_ckpt)\n        else:\n            no_improve += 1\n            if no_improve >= PATIENCE:\n                print(f\"     Early stopping epoch {epoch}\")\n                break\n\n    print(f\"  OK Fold {fold_id+1} best WSI_QWK: {best_kappa:.4f}\")\n    push_to_kaggle(best_ckpt,\n        f\"Stage{SHARD_ID} Fold{fold_id+1} OrdinalLoss QWK={best_kappa:.4f}\")\n    return best_kappa, best_ckpt\n\n# ─────────────────────────────────────────────\n# 11. 5-FOLD CV\n# ─────────────────────────────────────────────\nprint(f\"\\n{'='*60}\")\nprint(f\"Stage {SHARD_ID} Folds 4+5 — {len(all_wsis)} WSIs\")\nprint(f\"Loss: OrdinalCostSensitive (ordinal_weight=0.5)\")\nprint(f\"Folds deja completes : {DONE_KAPPAS}\")\nprint(f\"{'='*60}\")\n\nwsi_ids    = [r[\"wsi_id\"] for r in unique_records]\nlabels_arr = [r[\"label\"]  for r in unique_records]\nindices    = list(range(len(unique_records)))\n\nsgkf = StratifiedGroupKFold(n_splits=N_FOLDS, shuffle=True, random_state=SEED)\n\nfold_kappas       = DONE_KAPPAS.copy()\nfold_ckpts        = DONE_CKPTS.copy()\nbest_overall      = max(DONE_KAPPAS) if DONE_KAPPAS else -1.0\nbest_ckpt_overall = None\n\nfor fold_id, (train_idx, val_idx) in enumerate(\n        sgkf.split(indices, labels_arr, groups=wsi_ids)):\n\n    if fold_id < START_FOLD:\n        continue\n\n    train_data = [unique_records[i] for i in train_idx]\n    val_data   = [unique_records[i] for i in val_idx]\n\n    tw = set(unique_records[i][\"wsi_id\"] for i in train_idx)\n    vw = set(unique_records[i][\"wsi_id\"] for i in val_idx)\n    assert len(tw & vw) == 0, f\"DATA LEAKAGE fold {fold_id+1}!\"\n\n    kappa, ckpt_path = train_one_fold(fold_id, train_data, val_data)\n    fold_kappas.append(kappa)\n    fold_ckpts.append(ckpt_path)\n\n    if kappa > best_overall:\n        best_overall      = kappa\n        best_ckpt_overall = ckpt_path\n\n# ─────────────────────────────────────────────\n# 12. RESUME FINAL\n# ─────────────────────────────────────────────\nmean_kappa = np.mean(fold_kappas)\nstd_kappa  = np.std(fold_kappas)\n\nprint(f\"\\n{'='*60}\")\nprint(f\"OK Stage {SHARD_ID} — 5-Fold CV Results\")\nprint(f\"{'='*60}\")\nfor i, k in enumerate(fold_kappas):\n    print(f\"   Fold {i+1}: QWK = {k:.4f}\")\nprint(f\"{'─'*40}\")\nprint(f\"   Mean QWK : {mean_kappa:.4f} +/- {std_kappa:.4f}\")\nprint(f\"   Best QWK : {best_overall:.4f}\")\nprint(f\"   Stage 4  : 0.7112 +/- 0.0130\")\nprint(f\"   Stage 3  : 0.7280 +/- 0.0629\")\nprint(f\"   SOTA     : 0.934\")\nprint(f\"{'='*60}\")\n\nif best_ckpt_overall:\n    stage_best = f\"{SAVE_DIR}/stage{SHARD_ID}_best_qwk{best_overall:.4f}.pth\"\n    shutil.copy(best_ckpt_overall, stage_best)\n    print(f\"\\nOK Best checkpoint: {stage_best}\")\n    print(f\"\\n→ Pour Stage 6 :\")\n    print(f\"   SHARD_ID  = 6\")\n    print(f\"   PREV_CKPT = '{stage_best}'\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, glob\n\nfor f in glob.glob(\"/kaggle/working/stage5*.pth\"):\n    mb = os.path.getsize(f)/(1024*1024)\n    print(f\"{os.path.basename(f)} ({mb:.1f} MB)\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, glob, shutil, subprocess, json\n\ntmp = \"/kaggle/working/_push_s5_final\"\nos.makedirs(tmp, exist_ok=True)\n\nfor f in glob.glob(\"/kaggle/working/stage5_fold*.pth\"):\n    shutil.copy(f, tmp)\n    print(f\"  Copied: {os.path.basename(f)}\")\n\njson.dump({\n    \"title\": \"Stage5 Results Bejaoui 2026\",\n    \"id\": \"khouloudbejaoui20/stage5-results-bejaoui-2026\",\n    \"licenses\": [{\"name\": \"CC0-1.0\"}]\n}, open(f\"{tmp}/dataset-metadata.json\", \"w\"))\n\nr = subprocess.run([\n    \"kaggle\", \"datasets\", \"version\",\n    \"-p\", tmp,\n    \"-m\", \"Stage5 Fold4+5 OrdinalLoss\",\n    \"--dir-mode\", \"skip\"\n], capture_output=True, text=True, timeout=300)\n\nprint(\"OK!\" if r.returncode == 0 else f\"ERREUR: {r.stderr[:200]}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import subprocess, json, os\n\njson.dump({\n    \"title\": \"Stage5 Results Bejaoui 2026\",\n    \"id\": \"bejaouikhouloud/stage5-results-bejaoui-2026\",\n    \"licenses\": [{\"name\": \"CC0-1.0\"}]\n}, open(\"/kaggle/working/_push_s5_final/dataset-metadata.json\", \"w\"))\n\nr = subprocess.run([\n    \"kaggle\", \"datasets\", \"create\",\n    \"-p\", \"/kaggle/working/_push_s5_final\",\n    \"--dir-mode\", \"skip\"\n], capture_output=True, text=True, timeout=600)\n\nprint(\"OK!\" if r.returncode == 0 else f\"ERREUR: {r.stderr[:200]}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import subprocess, json, os, glob, shutil\n\n# Dataset qui a marché : bejaouikhouloud/stage5-results-bejaoui-2026\njson.dump({\n    \"title\": \"Stage5 Results Bejaoui 2026\",\n    \"id\": \"bejaouikhouloud/stage5-results-bejaoui-2026\",\n    \"licenses\": [{\"name\": \"CC0-1.0\"}]\n}, open(\"/kaggle/working/_push_s5_final/dataset-metadata.json\", \"w\"))\n\nr = subprocess.run([\n    \"kaggle\", \"datasets\", \"version\",\n    \"-p\", \"/kaggle/working/_push_s5_final\",\n    \"-m\", \"Stage5 Fold4+5 final\",\n    \"--dir-mode\", \"skip\"\n], capture_output=True, text=True, timeout=600)\n\nprint(\"OK!\" if r.returncode == 0 else f\"ERREUR: {r.stderr[:200]}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Lance dans Kaggle notebook\nimport subprocess\n\n# Installe les dépendances\nsubprocess.run([\n    \"pip\", \"install\", \"timm\", \"transformers\", \n    \"huggingface_hub\", \"--quiet\"\n], capture_output=True)\n\n# Test accès UNI\nfrom huggingface_hub import hf_hub_download\nimport timm, torch\n\ntry:\n    # Télécharge UNI\n    local_dir = \"/kaggle/working/uni_model\"\n    hf_hub_download(\n        repo_id=\"MahmoodLab/UNI\",\n        filename=\"pytorch_model.bin\",\n        local_dir=local_dir,\n        force_download=False\n    )\n    print(\"OK! UNI téléchargé!\")\nexcept Exception as e:\n    print(f\"ERREUR: {e}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from huggingface_hub import login\n\nlogin(token=\"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\")\n\n# Test téléchargement UNI\nimport timm\n\nmodel = timm.create_model(\n    \"hf-hub:MahmoodLab/UNI\",\n    pretrained=True,\n    num_classes=0,\n    init_values=1e-5,\n    dynamic_img_size=True\n)\n\nprint(f\"OK! UNI chargé!\")\nprint(f\"Output dim: {model.num_features}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from huggingface_hub import login, whoami\nimport timm\n\n# Login\nlogin(token=\"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\")\n\n# Vérifie\ninfo = whoami()\nprint(f\"Connecté: {info['name']}\")\n\n# Charge UNI\nmodel = timm.create_model(\n    \"hf-hub:MahmoodLab/UNI\",\n    pretrained=True,\n    num_classes=0,\n    init_values=1e-5,\n    dynamic_img_size=True\n)\n\nprint(f\"OK! UNI chargé!\")\nprint(f\"Features dim: {model.num_features}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\n=============================================================\nStage 0 — HybridUNITransformer\nUNI (MahmoodLab) foundation model + Transformer Encoder\nTrained from scratch on Shard 0 (399 WSIs, 4,319 patches)\n=============================================================\nAttache dans Kaggle :\n  khouloudbejaoui20/dataset\nGPU : T4×2\n\"\"\"\n\nimport os, glob, json, random, shutil, subprocess, datetime\nimport numpy as np\nfrom PIL import Image\nfrom collections import Counter\nfrom sklearn.model_selection import StratifiedGroupKFold\nfrom sklearn.metrics import cohen_kappa_score\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as T\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\n\n# ─────────────────────────────────────────────\n# 1. LOGIN HUGGINGFACE\n# ─────────────────────────────────────────────\nfrom huggingface_hub import login\nimport timm\n\nHF_TOKEN = \"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\"\nlogin(token=HF_TOKEN)\nprint(\"OK HuggingFace login!\")\n\n# ─────────────────────────────────────────────\n# 2. CONFIG\n# ─────────────────────────────────────────────\nSHARD_ID    = 0\nSEED        = 42\nBATCH_SIZE  = 16   # UNI est plus grand → batch size réduit\nNUM_WORKERS = 2\nEPOCHS      = 20   # moins car UNI déjà pretrained histopathologie\nPATIENCE    = 5\nLR_UNI      = 1e-5  # très faible → préserve les features histopathologiques\nLR_TR       = 4e-5  # Transformer encoder\n\nrandom.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)\ntorch.backends.cudnn.deterministic = True\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Device  : {DEVICE}\")\nprint(f\"Stage   : {SHARD_ID} — HybridUNITransformer\")\nprint(f\"EPOCHS={EPOCHS} | PATIENCE={PATIENCE} | LR_UNI={LR_UNI} | LR_TR={LR_TR}\")\n\n# ─────────────────────────────────────────────\n# 3. CHEMINS\n# ─────────────────────────────────────────────\nBASE = \"/kaggle/input/datasets/khouloudbejaoui20\"\n\nSHARD_PATH = f\"{BASE}/dataset/patches_400-20260630T212405Z-3-001/patches_400\"\n\nSAVE_DIR     = \"/kaggle/working\"\nRESULTS_JSON = f\"{SAVE_DIR}/stage0_uni_cv_results.json\"\n\n# ─────────────────────────────────────────────\n# 4. CHARGEMENT PATCHES\n# ─────────────────────────────────────────────\nprint(f\"\\n── Loading patches ──────────────────────\")\n\ndef load_patches(folder):\n    files = glob.glob(f\"{folder}/*.png\")\n    if not files:\n        files = glob.glob(f\"{folder}/**/*.png\", recursive=True)\n    records = []\n    for fp in files:\n        fname = os.path.basename(fp)\n        try:\n            parts  = fname.split(\"_slide\")\n            label  = int(parts[0].replace(\"grade\", \"\"))\n            wsi_id = parts[1].split(\"_\")[0]\n            records.append({\"path\": fp, \"wsi_id\": wsi_id, \"label\": label})\n        except:\n            continue\n    return records\n\nall_records = load_patches(SHARD_PATH)\nprint(f\"  Shard 0: {len(all_records)} patches\")\n\nall_wsis = list(set(r[\"wsi_id\"] for r in all_records))\nprint(f\"  Total WSIs: {len(all_wsis)}\")\ncnt = Counter(r[\"label\"] for r in all_records)\nfor c in range(6):\n    print(f\"    ISUP {c}: {cnt.get(c, 0)}\")\n\n# ─────────────────────────────────────────────\n# 5. TRANSFORMS\n# ─────────────────────────────────────────────\n# UNI utilise les mêmes stats ImageNet\nMEAN = [0.485, 0.456, 0.406]\nSTD  = [0.229, 0.224, 0.225]\n\ntrain_tf = T.Compose([\n    T.Resize((224, 224)),\n    T.RandomHorizontalFlip(),\n    T.RandomVerticalFlip(),\n    T.RandomRotation(90),\n    T.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.2, hue=0.1),\n    T.RandomApply([T.GaussianBlur(5)], p=0.3),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n    T.RandomApply([T.RandomErasing(p=0.5, scale=(0.02, 0.2))], p=0.2),\n])\n\nval_tf = T.Compose([\n    T.Resize((224, 224)),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n])\n\nclass PatchDataset(Dataset):\n    def __init__(self, records, transform):\n        self.records   = records\n        self.transform = transform\n    def __len__(self): return len(self.records)\n    def __getitem__(self, idx):\n        r = self.records[idx]\n        try:\n            img = Image.open(r[\"path\"]).convert(\"RGB\")\n        except:\n            img = Image.new(\"RGB\", (224, 224), (255, 255, 255))\n        return self.transform(img), int(r[\"label\"])\n\n# ─────────────────────────────────────────────\n# 6. MODELE — HybridUNITransformer\n# ─────────────────────────────────────────────\nclass HybridUNITransformer(nn.Module):\n    \"\"\"\n    UNI foundation model (1024-d features) +\n    Transformer Encoder (2 layers, 8 heads) +\n    MLP Classification Head (6 ISUP grades)\n    \"\"\"\n    def __init__(self, num_classes=6, d_model=512, nhead=8,\n                 num_layers=2, dim_feedforward=2048, dropout=0.1):\n        super().__init__()\n\n        # UNI backbone (histopathology foundation model)\n        self.uni_backbone = timm.create_model(\n            \"hf-hub:MahmoodLab/UNI\",\n            pretrained=True,\n            num_classes=0,         # enlève la tête de classification\n            init_values=1e-5,\n            dynamic_img_size=True\n        )\n        uni_dim = self.uni_backbone.num_features  # 1024\n\n        # Projection UNI → Transformer dim\n        self.proj = nn.Sequential(\n            nn.Linear(uni_dim, d_model),\n            nn.LayerNorm(d_model)\n        )\n\n        # CLS token + positional embedding\n        self.cls_token = nn.Parameter(torch.zeros(1, 1, d_model))\n        self.pos_embed = nn.Parameter(torch.zeros(1, 2, d_model))\n\n        # Transformer Encoder\n        encoder_layer = nn.TransformerEncoderLayer(\n            d_model=d_model, nhead=nhead,\n            dim_feedforward=dim_feedforward,\n            dropout=dropout, batch_first=True\n        )\n        self.transformer = nn.TransformerEncoder(\n            encoder_layer, num_layers=num_layers,\n            norm=nn.LayerNorm(d_model)\n        )\n\n        # MLP Head\n        self.classifier = nn.Sequential(\n            nn.Linear(d_model, 256),\n            nn.ReLU(),\n            nn.Dropout(dropout),\n            nn.Linear(256, num_classes)\n        )\n\n        # Init\n        nn.init.trunc_normal_(self.cls_token, std=0.02)\n        nn.init.trunc_normal_(self.pos_embed, std=0.02)\n\n    def forward(self, x):\n        # UNI features\n        f   = self.uni_backbone(x)           # (B, 1024)\n        f   = self.proj(f).unsqueeze(1)      # (B, 1, 512)\n\n        # CLS token + sequence\n        cls = self.cls_token.expand(x.size(0), -1, -1)\n        seq = torch.cat([cls, f], dim=1) + self.pos_embed\n\n        # Transformer\n        out = self.transformer(seq)[:, 0]    # CLS token\n\n        return self.classifier(out)\n\ndef build_model():\n    model  = HybridUNITransformer().to(DEVICE)\n    n_gpus = torch.cuda.device_count()\n    if n_gpus > 1:\n        model = nn.DataParallel(model)\n        print(f\"  OK DataParallel: {n_gpus} GPUs\")\n    n_param = sum(p.numel() for p in model.parameters())\n    print(f\"  OK Parameters: {n_param:,}\")\n    return model\n\n# ─────────────────────────────────────────────\n# 7. WSI-LEVEL QWK\n# ─────────────────────────────────────────────\ndef wsi_level_qwk(model, val_data, device):\n    model.eval()\n    wsi_preds  = {}\n    wsi_labels = {}\n    all_preds  = []\n    all_labels = []\n    loader = DataLoader(\n        PatchDataset(val_data, val_tf),\n        batch_size=BATCH_SIZE, shuffle=False,\n        num_workers=NUM_WORKERS\n    )\n    with torch.no_grad():\n        for imgs, labels in loader:\n            imgs  = imgs.to(device)\n            preds = model(imgs).argmax(1).cpu().numpy()\n            all_preds.extend(preds)\n            all_labels.extend(labels.numpy())\n    for i, r in enumerate(val_data):\n        wid = r[\"wsi_id\"]\n        if wid not in wsi_preds:\n            wsi_preds[wid]  = []\n            wsi_labels[wid] = r[\"label\"]\n        wsi_preds[wid].append(all_preds[i])\n    wf = []; wl = []\n    for wid in wsi_preds:\n        votes = wsi_preds[wid]\n        wf.append(max(set(votes), key=votes.count))\n        wl.append(wsi_labels[wid])\n    wsi_qwk = cohen_kappa_score(wl, wf, weights=\"quadratic\")\n    return wsi_qwk, len(wf)\n\n# ─────────────────────────────────────────────\n# 8. PUSH KAGGLE\n# ─────────────────────────────────────────────\ndef push_to_kaggle(ckpt_path, message):\n    try:\n        tmp = f\"{SAVE_DIR}/_push_uni\"\n        os.makedirs(tmp, exist_ok=True)\n        shutil.copy(ckpt_path, tmp)\n        if os.path.exists(RESULTS_JSON):\n            shutil.copy(RESULTS_JSON, tmp)\n        json.dump({\n            \"title\": \"Stage0 UNI Results Bejaoui 2026\",\n            \"id\": \"khouloudbejaoui20/stage0-uni-results\",\n            \"licenses\": [{\"name\": \"CC0-1.0\"}]\n        }, open(f\"{tmp}/dataset-metadata.json\", \"w\"))\n        r = subprocess.run(\n            [\"kaggle\", \"datasets\", \"create\",\n             \"-p\", tmp, \"--dir-mode\", \"skip\"],\n            capture_output=True, text=True, timeout=300\n        )\n        if r.returncode != 0:\n            r = subprocess.run(\n                [\"kaggle\", \"datasets\", \"version\",\n                 \"-p\", tmp, \"-m\", message, \"--dir-mode\", \"skip\"],\n                capture_output=True, text=True, timeout=300\n            )\n        print(f\"   {'OK!' if r.returncode==0 else 'ERREUR: '+r.stderr[:100]}\")\n    except Exception as e:\n        print(f\"   ERREUR: {e}\")\n\n# ─────────────────────────────────────────────\n# 9. TRAIN ONE FOLD\n# ─────────────────────────────────────────────\ndef train_one_fold(fold_id, train_data, val_data, n_folds=5):\n    print(f\"\\n  Fold {fold_id+1}/{n_folds} ──────────────────────────\")\n    print(f\"     Train: {len(train_data)} | Val: {len(val_data)}\")\n\n    train_loader = DataLoader(\n        PatchDataset(train_data, train_tf),\n        batch_size=BATCH_SIZE, shuffle=True,\n        num_workers=NUM_WORKERS, pin_memory=True\n    )\n\n    model = build_model()\n\n    # Class weights\n    cnt_tr = Counter(r[\"label\"] for r in train_data)\n    w = torch.tensor(\n        [1.0 / max(cnt_tr.get(c, 1), 1) for c in range(6)],\n        dtype=torch.float\n    )\n    w = (w / w.sum() * 6).to(DEVICE)\n    criterion = nn.CrossEntropyLoss(weight=w, label_smoothing=0.1)\n\n    # Differential LR : UNI backbone vs Transformer\n    base_model = model.module if hasattr(model, \"module\") else model\n    optimizer  = optim.AdamW([\n        {\"params\": base_model.uni_backbone.parameters(), \"lr\": LR_UNI},\n        {\"params\": [p for n, p in base_model.named_parameters()\n                    if not n.startswith(\"uni_backbone\")], \"lr\": LR_TR},\n    ], weight_decay=1e-4)\n\n    scheduler = CosineAnnealingLR(optimizer, T_max=EPOCHS, eta_min=1e-7)\n\n    best_kappa = 0.0\n    no_improve = 0\n    best_ckpt  = f\"{SAVE_DIR}/stage{SHARD_ID}_uni_fold{fold_id+1}_best.pth\"\n\n    for epoch in range(1, EPOCHS + 1):\n        model.train()\n        t_loss, t_correct, t_total = 0.0, 0, 0\n        for imgs, labels in train_loader:\n            imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n            optimizer.zero_grad()\n            out  = model(imgs)\n            loss = criterion(out, labels)\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n            optimizer.step()\n            t_loss    += loss.item() * imgs.size(0)\n            t_correct += (out.argmax(1) == labels).sum().item()\n            t_total   += imgs.size(0)\n        scheduler.step()\n\n        wsi_kappa, n_wsis_val = wsi_level_qwk(model, val_data, DEVICE)\n        is_best = wsi_kappa > best_kappa\n\n        print(f\"     Ep[{epoch:02d}/{EPOCHS}] \"\n              f\"loss={t_loss/t_total:.4f} acc={t_correct/t_total:.3f} | \"\n              f\"WSI_QWK={wsi_kappa:.4f} ({n_wsis_val} WSIs)\"\n              + (\" *\" if is_best else \"\"))\n\n        if is_best:\n            best_kappa = wsi_kappa\n            no_improve = 0\n            torch.save({\n                \"epoch\":            epoch,\n                \"model_state_dict\": model.state_dict() if not hasattr(model, \"module\")\n                                    else model.module.state_dict(),\n                \"kappa\":            wsi_kappa,\n                \"shard_id\":         SHARD_ID,\n                \"fold_id\":          fold_id,\n                \"backbone\":         \"UNI\",\n                \"n_wsis\":           len(all_wsis),\n            }, best_ckpt)\n        else:\n            no_improve += 1\n            if no_improve >= PATIENCE:\n                print(f\"     Early stopping epoch {epoch}\")\n                break\n\n    print(f\"  OK Fold {fold_id+1} best WSI_QWK: {best_kappa:.4f}\")\n    push_to_kaggle(best_ckpt,\n        f\"Stage{SHARD_ID} UNI Fold{fold_id+1} QWK={best_kappa:.4f}\")\n    return best_kappa, best_ckpt\n\n# ─────────────────────────────────────────────\n# 10. 5-FOLD CV\n# ─────────────────────────────────────────────\nN_FOLDS = 5\nprint(f\"\\n{'='*60}\")\nprint(f\"Stage {SHARD_ID} UNI — {len(all_wsis)} WSIs — 5-Fold CV\")\nprint(f\"Backbone: UNI (MahmoodLab) — histopathology foundation model\")\nprint(f\"Comparison: ResNet-50 baseline mean=0.5764\")\nprint(f\"{'='*60}\")\n\nwsi_ids    = [r[\"wsi_id\"] for r in all_records]\nlabels_arr = [r[\"label\"]  for r in all_records]\nindices    = list(range(len(all_records)))\n\nsgkf = StratifiedGroupKFold(n_splits=N_FOLDS, shuffle=True, random_state=SEED)\n\nfold_kappas       = []\nbest_overall      = -1.0\nbest_ckpt_overall = None\n\nfor fold_id, (train_idx, val_idx) in enumerate(\n        sgkf.split(indices, labels_arr, groups=wsi_ids)):\n\n    train_data = [all_records[i] for i in train_idx]\n    val_data   = [all_records[i] for i in val_idx]\n\n    tw = set(all_records[i][\"wsi_id\"] for i in train_idx)\n    vw = set(all_records[i][\"wsi_id\"] for i in val_idx)\n    assert len(tw & vw) == 0, f\"DATA LEAKAGE fold {fold_id+1}!\"\n\n    kappa, ckpt_path = train_one_fold(fold_id, train_data, val_data, N_FOLDS)\n    fold_kappas.append(kappa)\n\n    if kappa > best_overall:\n        best_overall      = kappa\n        best_ckpt_overall = ckpt_path\n\n# ─────────────────────────────────────────────\n# 11. RESUME FINAL\n# ─────────────────────────────────────────────\nmean_kappa = np.mean(fold_kappas)\nstd_kappa  = np.std(fold_kappas)\n\nprint(f\"\\n{'='*60}\")\nprint(f\"OK Stage {SHARD_ID} UNI — 5-Fold CV Results\")\nprint(f\"{'='*60}\")\nfor i, k in enumerate(fold_kappas):\n    print(f\"   Fold {i+1}: QWK = {k:.4f}\")\nprint(f\"{'─'*40}\")\nprint(f\"   Mean QWK UNI      : {mean_kappa:.4f} +/- {std_kappa:.4f}\")\nprint(f\"   Mean QWK ResNet-50: 0.5764 +/- 0.0224 (baseline)\")\nprint(f\"   Gain UNI vs ResNet: {mean_kappa - 0.5764:+.4f}\")\nprint(f\"   Best QWK UNI      : {best_overall:.4f}\")\nprint(f\"   WSIs              : {len(all_wsis)}\")\nprint(f\"   SOTA [1]          : 0.934\")\nprint(f\"{'='*60}\")\n\nresults = {\n    \"stage\":          SHARD_ID,\n    \"backbone\":       \"UNI\",\n    \"n_wsis\":         len(all_wsis),\n    \"n_patches\":      len(all_records),\n    \"fold_kappas\":    fold_kappas,\n    \"mean_kappa\":     mean_kappa,\n    \"std_kappa\":      std_kappa,\n    \"best_kappa\":     best_overall,\n    \"resnet50_mean\":  0.5764,\n    \"gain_vs_resnet\": mean_kappa - 0.5764,\n    \"timestamp\":      datetime.datetime.now().isoformat(),\n}\nwith open(RESULTS_JSON, \"w\") as f:\n    json.dump(results, f, indent=2)\n\nif best_ckpt_overall:\n    stage_best = f\"{SAVE_DIR}/stage{SHARD_ID}_uni_best_qwk{best_overall:.4f}.pth\"\n    shutil.copy(best_ckpt_overall, stage_best)\n    print(f\"\\nOK Best checkpoint: {stage_best}\")\n    print(f\"\\nComparaison finale:\")\n    print(f\"  ResNet-50 Stage 0 : mean=0.5764, best=0.6116\")\n    print(f\"  UNI Stage 0       : mean={mean_kappa:.4f}, best={best_overall:.4f}\")\n    print(f\"  Gain              : {mean_kappa - 0.5764:+.4f}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, glob\n\nfor f in glob.glob(\"/kaggle/working/stage0_uni*.pth\"):\n    mb = os.path.getsize(f)/(1024*1024)\n    print(f\"{os.path.basename(f)} ({mb:.1f} MB)\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\n=============================================================\nStage 0 — HybridUNITransformer\nUNI (MahmoodLab) foundation model + Transformer Encoder\nTrained from scratch on Shard 0 (399 WSIs, 4,319 patches)\n=============================================================\nAttache dans Kaggle :\n  khouloudbejaoui20/dataset\nGPU : T4×2\n\"\"\"\n\nimport os, glob, json, random, shutil, subprocess, datetime\nimport numpy as np\nfrom PIL import Image\nfrom collections import Counter\nfrom sklearn.model_selection import StratifiedGroupKFold\nfrom sklearn.metrics import cohen_kappa_score\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as T\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\n\n# ─────────────────────────────────────────────\n# 1. LOGIN HUGGINGFACE\n# ─────────────────────────────────────────────\nfrom huggingface_hub import login\nimport timm\n\nHF_TOKEN = \"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\"\nlogin(token=HF_TOKEN)\nprint(\"OK HuggingFace login!\")\n\n# ─────────────────────────────────────────────\n# 2. CONFIG\n# ─────────────────────────────────────────────\nSHARD_ID    = 0\nSTART_FOLD  = 1        # Reprend depuis Fold 2\nDONE_KAPPAS = [0.7215] # Fold 1 résultat\nSEED        = 42\nBATCH_SIZE  = 16   # UNI est plus grand → batch size réduit\nNUM_WORKERS = 2\nEPOCHS      = 20   # moins car UNI déjà pretrained histopathologie\nPATIENCE    = 5\nLR_UNI      = 1e-5  # très faible → préserve les features histopathologiques\nLR_TR       = 4e-5  # Transformer encoder\n\nrandom.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)\ntorch.backends.cudnn.deterministic = True\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Device  : {DEVICE}\")\nprint(f\"Stage   : {SHARD_ID} — HybridUNITransformer\")\nprint(f\"EPOCHS={EPOCHS} | PATIENCE={PATIENCE} | LR_UNI={LR_UNI} | LR_TR={LR_TR}\")\n\n# ─────────────────────────────────────────────\n# 3. CHEMINS\n# ─────────────────────────────────────────────\nBASE = \"/kaggle/input/datasets/khouloudbejaoui20\"\n\nSHARD_PATH = f\"{BASE}/dataset/patches_400-20260630T212405Z-3-001/patches_400\"\n\nSAVE_DIR     = \"/kaggle/working\"\nRESULTS_JSON = f\"{SAVE_DIR}/stage0_uni_cv_results.json\"\n\n# ─────────────────────────────────────────────\n# 4. CHARGEMENT PATCHES\n# ─────────────────────────────────────────────\nprint(f\"\\n── Loading patches ──────────────────────\")\n\ndef load_patches(folder):\n    files = glob.glob(f\"{folder}/*.png\")\n    if not files:\n        files = glob.glob(f\"{folder}/**/*.png\", recursive=True)\n    records = []\n    for fp in files:\n        fname = os.path.basename(fp)\n        try:\n            parts  = fname.split(\"_slide\")\n            label  = int(parts[0].replace(\"grade\", \"\"))\n            wsi_id = parts[1].split(\"_\")[0]\n            records.append({\"path\": fp, \"wsi_id\": wsi_id, \"label\": label})\n        except:\n            continue\n    return records\n\nall_records = load_patches(SHARD_PATH)\nprint(f\"  Shard 0: {len(all_records)} patches\")\n\nall_wsis = list(set(r[\"wsi_id\"] for r in all_records))\nprint(f\"  Total WSIs: {len(all_wsis)}\")\ncnt = Counter(r[\"label\"] for r in all_records)\nfor c in range(6):\n    print(f\"    ISUP {c}: {cnt.get(c, 0)}\")\n\n# ─────────────────────────────────────────────\n# 5. TRANSFORMS\n# ─────────────────────────────────────────────\n# UNI utilise les mêmes stats ImageNet\nMEAN = [0.485, 0.456, 0.406]\nSTD  = [0.229, 0.224, 0.225]\n\ntrain_tf = T.Compose([\n    T.Resize((224, 224)),\n    T.RandomHorizontalFlip(),\n    T.RandomVerticalFlip(),\n    T.RandomRotation(90),\n    T.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.2, hue=0.1),\n    T.RandomApply([T.GaussianBlur(5)], p=0.3),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n    T.RandomApply([T.RandomErasing(p=0.5, scale=(0.02, 0.2))], p=0.2),\n])\n\nval_tf = T.Compose([\n    T.Resize((224, 224)),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n])\n\nclass PatchDataset(Dataset):\n    def __init__(self, records, transform):\n        self.records   = records\n        self.transform = transform\n    def __len__(self): return len(self.records)\n    def __getitem__(self, idx):\n        r = self.records[idx]\n        try:\n            img = Image.open(r[\"path\"]).convert(\"RGB\")\n        except:\n            img = Image.new(\"RGB\", (224, 224), (255, 255, 255))\n        return self.transform(img), int(r[\"label\"])\n\n# ─────────────────────────────────────────────\n# 6. MODELE — HybridUNITransformer\n# ─────────────────────────────────────────────\nclass HybridUNITransformer(nn.Module):\n    \"\"\"\n    UNI foundation model (1024-d features) +\n    Transformer Encoder (2 layers, 8 heads) +\n    MLP Classification Head (6 ISUP grades)\n    \"\"\"\n    def __init__(self, num_classes=6, d_model=512, nhead=8,\n                 num_layers=2, dim_feedforward=2048, dropout=0.1):\n        super().__init__()\n\n        # UNI backbone (histopathology foundation model)\n        self.uni_backbone = timm.create_model(\n            \"hf-hub:MahmoodLab/UNI\",\n            pretrained=True,\n            num_classes=0,         # enlève la tête de classification\n            init_values=1e-5,\n            dynamic_img_size=True\n        )\n        uni_dim = self.uni_backbone.num_features  # 1024\n\n        # Projection UNI → Transformer dim\n        self.proj = nn.Sequential(\n            nn.Linear(uni_dim, d_model),\n            nn.LayerNorm(d_model)\n        )\n\n        # CLS token + positional embedding\n        self.cls_token = nn.Parameter(torch.zeros(1, 1, d_model))\n        self.pos_embed = nn.Parameter(torch.zeros(1, 2, d_model))\n\n        # Transformer Encoder\n        encoder_layer = nn.TransformerEncoderLayer(\n            d_model=d_model, nhead=nhead,\n            dim_feedforward=dim_feedforward,\n            dropout=dropout, batch_first=True\n        )\n        self.transformer = nn.TransformerEncoder(\n            encoder_layer, num_layers=num_layers,\n            norm=nn.LayerNorm(d_model)\n        )\n\n        # MLP Head\n        self.classifier = nn.Sequential(\n            nn.Linear(d_model, 256),\n            nn.ReLU(),\n            nn.Dropout(dropout),\n            nn.Linear(256, num_classes)\n        )\n\n        # Init\n        nn.init.trunc_normal_(self.cls_token, std=0.02)\n        nn.init.trunc_normal_(self.pos_embed, std=0.02)\n\n    def forward(self, x):\n        # UNI features\n        f   = self.uni_backbone(x)           # (B, 1024)\n        f   = self.proj(f).unsqueeze(1)      # (B, 1, 512)\n\n        # CLS token + sequence\n        cls = self.cls_token.expand(x.size(0), -1, -1)\n        seq = torch.cat([cls, f], dim=1) + self.pos_embed\n\n        # Transformer\n        out = self.transformer(seq)[:, 0]    # CLS token\n\n        return self.classifier(out)\n\ndef build_model():\n    model  = HybridUNITransformer().to(DEVICE)\n    n_gpus = torch.cuda.device_count()\n    if n_gpus > 1:\n        model = nn.DataParallel(model)\n        print(f\"  OK DataParallel: {n_gpus} GPUs\")\n    n_param = sum(p.numel() for p in model.parameters())\n    print(f\"  OK Parameters: {n_param:,}\")\n    return model\n\n# ─────────────────────────────────────────────\n# 7. WSI-LEVEL QWK\n# ─────────────────────────────────────────────\ndef wsi_level_qwk(model, val_data, device):\n    model.eval()\n    wsi_preds  = {}\n    wsi_labels = {}\n    all_preds  = []\n    all_labels = []\n    loader = DataLoader(\n        PatchDataset(val_data, val_tf),\n        batch_size=BATCH_SIZE, shuffle=False,\n        num_workers=NUM_WORKERS\n    )\n    with torch.no_grad():\n        for imgs, labels in loader:\n            imgs  = imgs.to(device)\n            preds = model(imgs).argmax(1).cpu().numpy()\n            all_preds.extend(preds)\n            all_labels.extend(labels.numpy())\n    for i, r in enumerate(val_data):\n        wid = r[\"wsi_id\"]\n        if wid not in wsi_preds:\n            wsi_preds[wid]  = []\n            wsi_labels[wid] = r[\"label\"]\n        wsi_preds[wid].append(all_preds[i])\n    wf = []; wl = []\n    for wid in wsi_preds:\n        votes = wsi_preds[wid]\n        wf.append(max(set(votes), key=votes.count))\n        wl.append(wsi_labels[wid])\n    wsi_qwk = cohen_kappa_score(wl, wf, weights=\"quadratic\")\n    return wsi_qwk, len(wf)\n\n# ─────────────────────────────────────────────\n# 8. PUSH KAGGLE\n# ─────────────────────────────────────────────\ndef push_to_kaggle(ckpt_path, message):\n    try:\n        tmp = f\"{SAVE_DIR}/_push_uni\"\n        os.makedirs(tmp, exist_ok=True)\n        shutil.copy(ckpt_path, tmp)\n        if os.path.exists(RESULTS_JSON):\n            shutil.copy(RESULTS_JSON, tmp)\n        json.dump({\n            \"title\": \"Stage0 UNI Results Bejaoui 2026\",\n            \"id\": \"khouloudbejaoui20/stage0-uni-results\",\n            \"licenses\": [{\"name\": \"CC0-1.0\"}]\n        }, open(f\"{tmp}/dataset-metadata.json\", \"w\"))\n        r = subprocess.run(\n            [\"kaggle\", \"datasets\", \"create\",\n             \"-p\", tmp, \"--dir-mode\", \"skip\"],\n            capture_output=True, text=True, timeout=300\n        )\n        if r.returncode != 0:\n            r = subprocess.run(\n                [\"kaggle\", \"datasets\", \"version\",\n                 \"-p\", tmp, \"-m\", message, \"--dir-mode\", \"skip\"],\n                capture_output=True, text=True, timeout=300\n            )\n        print(f\"   {'OK!' if r.returncode==0 else 'ERREUR: '+r.stderr[:100]}\")\n    except Exception as e:\n        print(f\"   ERREUR: {e}\")\n\n# ─────────────────────────────────────────────\n# 9. TRAIN ONE FOLD\n# ─────────────────────────────────────────────\ndef train_one_fold(fold_id, train_data, val_data, n_folds=5):\n    print(f\"\\n  Fold {fold_id+1}/{n_folds} ──────────────────────────\")\n    print(f\"     Train: {len(train_data)} | Val: {len(val_data)}\")\n\n    train_loader = DataLoader(\n        PatchDataset(train_data, train_tf),\n        batch_size=BATCH_SIZE, shuffle=True,\n        num_workers=NUM_WORKERS, pin_memory=True\n    )\n\n    model = build_model()\n\n    # Class weights\n    cnt_tr = Counter(r[\"label\"] for r in train_data)\n    w = torch.tensor(\n        [1.0 / max(cnt_tr.get(c, 1), 1) for c in range(6)],\n        dtype=torch.float\n    )\n    w = (w / w.sum() * 6).to(DEVICE)\n    criterion = nn.CrossEntropyLoss(weight=w, label_smoothing=0.1)\n\n    # Differential LR : UNI backbone vs Transformer\n    base_model = model.module if hasattr(model, \"module\") else model\n    optimizer  = optim.AdamW([\n        {\"params\": base_model.uni_backbone.parameters(), \"lr\": LR_UNI},\n        {\"params\": [p for n, p in base_model.named_parameters()\n                    if not n.startswith(\"uni_backbone\")], \"lr\": LR_TR},\n    ], weight_decay=1e-4)\n\n    scheduler = CosineAnnealingLR(optimizer, T_max=EPOCHS, eta_min=1e-7)\n\n    best_kappa = 0.0\n    no_improve = 0\n    best_ckpt  = f\"{SAVE_DIR}/stage{SHARD_ID}_uni_fold{fold_id+1}_best.pth\"\n\n    for epoch in range(1, EPOCHS + 1):\n        model.train()\n        t_loss, t_correct, t_total = 0.0, 0, 0\n        for imgs, labels in train_loader:\n            imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n            optimizer.zero_grad()\n            out  = model(imgs)\n            loss = criterion(out, labels)\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n            optimizer.step()\n            t_loss    += loss.item() * imgs.size(0)\n            t_correct += (out.argmax(1) == labels).sum().item()\n            t_total   += imgs.size(0)\n        scheduler.step()\n\n        wsi_kappa, n_wsis_val = wsi_level_qwk(model, val_data, DEVICE)\n        is_best = wsi_kappa > best_kappa\n\n        print(f\"     Ep[{epoch:02d}/{EPOCHS}] \"\n              f\"loss={t_loss/t_total:.4f} acc={t_correct/t_total:.3f} | \"\n              f\"WSI_QWK={wsi_kappa:.4f} ({n_wsis_val} WSIs)\"\n              + (\" *\" if is_best else \"\"))\n\n        if is_best:\n            best_kappa = wsi_kappa\n            no_improve = 0\n            torch.save({\n                \"epoch\":            epoch,\n                \"model_state_dict\": model.state_dict() if not hasattr(model, \"module\")\n                                    else model.module.state_dict(),\n                \"kappa\":            wsi_kappa,\n                \"shard_id\":         SHARD_ID,\n                \"fold_id\":          fold_id,\n                \"backbone\":         \"UNI\",\n                \"n_wsis\":           len(all_wsis),\n            }, best_ckpt)\n        else:\n            no_improve += 1\n            if no_improve >= PATIENCE:\n                print(f\"     Early stopping epoch {epoch}\")\n                break\n\n    print(f\"  OK Fold {fold_id+1} best WSI_QWK: {best_kappa:.4f}\")\n    push_to_kaggle(best_ckpt,\n        f\"Stage{SHARD_ID} UNI Fold{fold_id+1} QWK={best_kappa:.4f}\")\n    return best_kappa, best_ckpt\n\n# ─────────────────────────────────────────────\n# 10. 5-FOLD CV\n# ─────────────────────────────────────────────\nN_FOLDS = 5\nprint(f\"\\n{'='*60}\")\nprint(f\"Stage {SHARD_ID} UNI — {len(all_wsis)} WSIs — 5-Fold CV\")\nprint(f\"Backbone: UNI (MahmoodLab) — histopathology foundation model\")\nprint(f\"Comparison: ResNet-50 baseline mean=0.5764\")\nprint(f\"{'='*60}\")\n\nwsi_ids    = [r[\"wsi_id\"] for r in all_records]\nlabels_arr = [r[\"label\"]  for r in all_records]\nindices    = list(range(len(all_records)))\n\nsgkf = StratifiedGroupKFold(n_splits=N_FOLDS, shuffle=True, random_state=SEED)\n\nfold_kappas       = DONE_KAPPAS.copy()\nbest_overall      = -1.0\nbest_ckpt_overall = None\n\nfor fold_id, (train_idx, val_idx) in enumerate(\n        sgkf.split(indices, labels_arr, groups=wsi_ids)):\n\n    train_data = [all_records[i] for i in train_idx]\n    val_data   = [all_records[i] for i in val_idx]\n\n    if fold_id < START_FOLD:\n        continue\n\n    tw = set(all_records[i][\"wsi_id\"] for i in train_idx)\n    vw = set(all_records[i][\"wsi_id\"] for i in val_idx)\n    assert len(tw & vw) == 0, f\"DATA LEAKAGE fold {fold_id+1}!\"\n\n    kappa, ckpt_path = train_one_fold(fold_id, train_data, val_data, N_FOLDS)\n    fold_kappas.append(kappa)\n\n    if kappa > best_overall:\n        best_overall      = kappa\n        best_ckpt_overall = ckpt_path\n\n# ─────────────────────────────────────────────\n# 11. RESUME FINAL\n# ─────────────────────────────────────────────\nmean_kappa = np.mean(fold_kappas)\nstd_kappa  = np.std(fold_kappas)\n\nprint(f\"\\n{'='*60}\")\nprint(f\"OK Stage {SHARD_ID} UNI — 5-Fold CV Results\")\nprint(f\"{'='*60}\")\nfor i, k in enumerate(fold_kappas):\n    print(f\"   Fold {i+1}: QWK = {k:.4f}\")\nprint(f\"{'─'*40}\")\nprint(f\"   Mean QWK UNI      : {mean_kappa:.4f} +/- {std_kappa:.4f}\")\nprint(f\"   Mean QWK ResNet-50: 0.5764 +/- 0.0224 (baseline)\")\nprint(f\"   Gain UNI vs ResNet: {mean_kappa - 0.5764:+.4f}\")\nprint(f\"   Best QWK UNI      : {best_overall:.4f}\")\nprint(f\"   WSIs              : {len(all_wsis)}\")\nprint(f\"   SOTA [1]          : 0.934\")\nprint(f\"{'='*60}\")\n\nresults = {\n    \"stage\":          SHARD_ID,\n    \"backbone\":       \"UNI\",\n    \"n_wsis\":         len(all_wsis),\n    \"n_patches\":      len(all_records),\n    \"fold_kappas\":    fold_kappas,\n    \"mean_kappa\":     mean_kappa,\n    \"std_kappa\":      std_kappa,\n    \"best_kappa\":     best_overall,\n    \"resnet50_mean\":  0.5764,\n    \"gain_vs_resnet\": mean_kappa - 0.5764,\n    \"timestamp\":      datetime.datetime.now().isoformat(),\n}\nwith open(RESULTS_JSON, \"w\") as f:\n    json.dump(results, f, indent=2)\n\nif best_ckpt_overall:\n    stage_best = f\"{SAVE_DIR}/stage{SHARD_ID}_uni_best_qwk{best_overall:.4f}.pth\"\n    shutil.copy(best_ckpt_overall, stage_best)\n    print(f\"\\nOK Best checkpoint: {stage_best}\")\n    print(f\"\\nComparaison finale:\")\n    print(f\"  ResNet-50 Stage 0 : mean=0.5764, best=0.6116\")\n    print(f\"  UNI Stage 0       : mean={mean_kappa:.4f}, best={best_overall:.4f}\")\n    print(f\"  Gain              : {mean_kappa - 0.5764:+.4f}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import subprocess, os, json\n\n# Config bejaouikhouloud\nkaggle_creds = {\n    \"username\": \"bejaouikhouloud\",\n    \"key\": \"af385e4bd0b6181b85b83134f6f1c405\"\n}\nos.makedirs(\"/root/.config/kaggle\", exist_ok=True)\nwith open(\"/root/.config/kaggle/kaggle.json\", \"w\") as f:\n    json.dump(kaggle_creds, f)\nos.chmod(\"/root/.config/kaggle/kaggle.json\", 0o600)\n\n# Télécharge dataset UNI\nsubprocess.run([\n    \"kaggle\", \"datasets\", \"download\",\n    \"bejaouikhouloud/stage0-uni-results\",\n    \"-p\", \"/kaggle/working\", \"--unzip\"\n], capture_output=True)\n\nimport glob\nfor f in glob.glob(\"/kaggle/working/stage0_uni*.pth\"):\n    print(f\"OK: {os.path.basename(f)}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\n=============================================================\nStage 0 — HybridUNITransformer\nUNI (MahmoodLab) foundation model + Transformer Encoder\nTrained from scratch on Shard 0 (399 WSIs, 4,319 patches)\n=============================================================\nAttache dans Kaggle :\n  khouloudbejaoui20/dataset\nGPU : T4×2\n\"\"\"\n\nimport os, glob, json, random, shutil, subprocess, datetime\nimport numpy as np\nfrom PIL import Image\nfrom collections import Counter\nfrom sklearn.model_selection import StratifiedGroupKFold\nfrom sklearn.metrics import cohen_kappa_score\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as T\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\n\n# ─────────────────────────────────────────────\n# 1. LOGIN HUGGINGFACE\n# ─────────────────────────────────────────────\nfrom huggingface_hub import login\nimport timm\n\nHF_TOKEN = \"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\"\nlogin(token=HF_TOKEN)\nprint(\"OK HuggingFace login!\")\n\n# ─────────────────────────────────────────────\n# 2. CONFIG\n# ─────────────────────────────────────────────\nSHARD_ID    = 0\n\nSTART_FOLD  = 4        # Fold 5 (index commence à 0)\nDONE_KAPPAS = [0.7215, 0.6920, 0.7964, 0.5231]\n\n\nSEED        = 42\nBATCH_SIZE  = 8   # UNI est plus grand → batch size réduit\nNUM_WORKERS = 2\nEPOCHS      = 20   # moins car UNI déjà pretrained histopathologie\nPATIENCE    = 5\nLR_UNI      = 1e-5  # très faible → préserve les features histopathologiques\nLR_TR       = 4e-5  # Transformer encoder\n\nrandom.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)\ntorch.backends.cudnn.deterministic = True\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Device  : {DEVICE}\")\nprint(f\"Stage   : {SHARD_ID} — HybridUNITransformer\")\nprint(f\"EPOCHS={EPOCHS} | PATIENCE={PATIENCE} | LR_UNI={LR_UNI} | LR_TR={LR_TR}\")\n\n# ─────────────────────────────────────────────\n# 3. CHEMINS\n# ─────────────────────────────────────────────\nBASE = \"/kaggle/input/datasets/khouloudbejaoui20\"\n\nSHARD_PATH = f\"{BASE}/dataset/patches_400-20260630T212405Z-3-001/patches_400\"\n\nSAVE_DIR     = \"/kaggle/working\"\nRESULTS_JSON = f\"{SAVE_DIR}/stage0_uni_cv_results.json\"\n\n# ─────────────────────────────────────────────\n# 4. CHARGEMENT PATCHES\n# ─────────────────────────────────────────────\nprint(f\"\\n── Loading patches ──────────────────────\")\n\ndef load_patches(folder):\n    files = glob.glob(f\"{folder}/*.png\")\n    if not files:\n        files = glob.glob(f\"{folder}/**/*.png\", recursive=True)\n    records = []\n    for fp in files:\n        fname = os.path.basename(fp)\n        try:\n            parts  = fname.split(\"_slide\")\n            label  = int(parts[0].replace(\"grade\", \"\"))\n            wsi_id = parts[1].split(\"_\")[0]\n            records.append({\"path\": fp, \"wsi_id\": wsi_id, \"label\": label})\n        except:\n            continue\n    return records\n\nall_records = load_patches(SHARD_PATH)\nprint(f\"  Shard 0: {len(all_records)} patches\")\n\nall_wsis = list(set(r[\"wsi_id\"] for r in all_records))\nprint(f\"  Total WSIs: {len(all_wsis)}\")\ncnt = Counter(r[\"label\"] for r in all_records)\nfor c in range(6):\n    print(f\"    ISUP {c}: {cnt.get(c, 0)}\")\n\n# ─────────────────────────────────────────────\n# 5. TRANSFORMS\n# ─────────────────────────────────────────────\n# UNI utilise les mêmes stats ImageNet\nMEAN = [0.485, 0.456, 0.406]\nSTD  = [0.229, 0.224, 0.225]\n\ntrain_tf = T.Compose([\n    T.Resize((224, 224)),\n    T.RandomHorizontalFlip(),\n    T.RandomVerticalFlip(),\n    T.RandomRotation(90),\n    T.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.2, hue=0.1),\n    T.RandomApply([T.GaussianBlur(5)], p=0.3),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n    T.RandomApply([T.RandomErasing(p=0.5, scale=(0.02, 0.2))], p=0.2),\n])\n\nval_tf = T.Compose([\n    T.Resize((224, 224)),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n])\n\nclass PatchDataset(Dataset):\n    def __init__(self, records, transform):\n        self.records   = records\n        self.transform = transform\n    def __len__(self): return len(self.records)\n    def __getitem__(self, idx):\n        r = self.records[idx]\n        try:\n            img = Image.open(r[\"path\"]).convert(\"RGB\")\n        except:\n            img = Image.new(\"RGB\", (224, 224), (255, 255, 255))\n        return self.transform(img), int(r[\"label\"])\n\n# ─────────────────────────────────────────────\n# 6. MODELE — HybridUNITransformer\n# ─────────────────────────────────────────────\nclass HybridUNITransformer(nn.Module):\n    \"\"\"\n    UNI foundation model (1024-d features) +\n    Transformer Encoder (2 layers, 8 heads) +\n    MLP Classification Head (6 ISUP grades)\n    \"\"\"\n    def __init__(self, num_classes=6, d_model=512, nhead=8,\n                 num_layers=2, dim_feedforward=2048, dropout=0.1):\n        super().__init__()\n\n        # UNI backbone (histopathology foundation model)\n        self.uni_backbone = timm.create_model(\n            \"hf-hub:MahmoodLab/UNI\",\n            pretrained=True,\n            num_classes=0,         # enlève la tête de classification\n            init_values=1e-5,\n            dynamic_img_size=True\n        )\n        uni_dim = self.uni_backbone.num_features  # 1024\n\n        # Projection UNI → Transformer dim\n        self.proj = nn.Sequential(\n            nn.Linear(uni_dim, d_model),\n            nn.LayerNorm(d_model)\n        )\n\n        # CLS token + positional embedding\n        self.cls_token = nn.Parameter(torch.zeros(1, 1, d_model))\n        self.pos_embed = nn.Parameter(torch.zeros(1, 2, d_model))\n\n        # Transformer Encoder\n        encoder_layer = nn.TransformerEncoderLayer(\n            d_model=d_model, nhead=nhead,\n            dim_feedforward=dim_feedforward,\n            dropout=dropout, batch_first=True\n        )\n        self.transformer = nn.TransformerEncoder(\n            encoder_layer, num_layers=num_layers,\n            norm=nn.LayerNorm(d_model)\n        )\n\n        # MLP Head\n        self.classifier = nn.Sequential(\n            nn.Linear(d_model, 256),\n            nn.ReLU(),\n            nn.Dropout(dropout),\n            nn.Linear(256, num_classes)\n        )\n\n        # Init\n        nn.init.trunc_normal_(self.cls_token, std=0.02)\n        nn.init.trunc_normal_(self.pos_embed, std=0.02)\n\n    def forward(self, x):\n        # UNI features\n        f   = self.uni_backbone(x)           # (B, 1024)\n        f   = self.proj(f).unsqueeze(1)      # (B, 1, 512)\n\n        # CLS token + sequence\n        cls = self.cls_token.expand(x.size(0), -1, -1)\n        seq = torch.cat([cls, f], dim=1) + self.pos_embed\n\n        # Transformer\n        out = self.transformer(seq)[:, 0]    # CLS token\n\n        return self.classifier(out)\n\ndef build_model():\n    model  = HybridUNITransformer().to(DEVICE)\n    n_gpus = torch.cuda.device_count()\n    if n_gpus > 1:\n        model = nn.DataParallel(model)\n        print(f\"  OK DataParallel: {n_gpus} GPUs\")\n    n_param = sum(p.numel() for p in model.parameters())\n    print(f\"  OK Parameters: {n_param:,}\")\n    return model\n\n# ─────────────────────────────────────────────\n# 7. WSI-LEVEL QWK\n# ─────────────────────────────────────────────\ndef wsi_level_qwk(model, val_data, device):\n    model.eval()\n    wsi_preds  = {}\n    wsi_labels = {}\n    all_preds  = []\n    all_labels = []\n    loader = DataLoader(\n        PatchDataset(val_data, val_tf),\n        batch_size=BATCH_SIZE, shuffle=False,\n        num_workers=NUM_WORKERS\n    )\n    with torch.no_grad():\n        for imgs, labels in loader:\n            imgs  = imgs.to(device)\n            preds = model(imgs).argmax(1).cpu().numpy()\n            all_preds.extend(preds)\n            all_labels.extend(labels.numpy())\n    for i, r in enumerate(val_data):\n        wid = r[\"wsi_id\"]\n        if wid not in wsi_preds:\n            wsi_preds[wid]  = []\n            wsi_labels[wid] = r[\"label\"]\n        wsi_preds[wid].append(all_preds[i])\n    wf = []; wl = []\n    for wid in wsi_preds:\n        votes = wsi_preds[wid]\n        wf.append(max(set(votes), key=votes.count))\n        wl.append(wsi_labels[wid])\n    wsi_qwk = cohen_kappa_score(wl, wf, weights=\"quadratic\")\n    return wsi_qwk, len(wf)\n\n# ─────────────────────────────────────────────\n# 8. PUSH KAGGLE\n# ─────────────────────────────────────────────\ndef push_to_kaggle(ckpt_path, message):\n    try:\n        tmp = f\"{SAVE_DIR}/_push_uni\"\n        os.makedirs(tmp, exist_ok=True)\n        shutil.copy(ckpt_path, tmp)\n        if os.path.exists(RESULTS_JSON):\n            shutil.copy(RESULTS_JSON, tmp)\n        json.dump({\n            \"title\": \"Stage0 UNI Results Bejaoui 2026\",\n            \"id\": \"khouloudbejaoui20/stage0-uni-results\",\n            \"licenses\": [{\"name\": \"CC0-1.0\"}]\n        }, open(f\"{tmp}/dataset-metadata.json\", \"w\"))\n        r = subprocess.run(\n            [\"kaggle\", \"datasets\", \"create\",\n             \"-p\", tmp, \"--dir-mode\", \"skip\"],\n            capture_output=True, text=True, timeout=300\n        )\n        if r.returncode != 0:\n            r = subprocess.run(\n                [\"kaggle\", \"datasets\", \"version\",\n                 \"-p\", tmp, \"-m\", message, \"--dir-mode\", \"skip\"],\n                capture_output=True, text=True, timeout=300\n            )\n        print(f\"   {'OK!' if r.returncode==0 else 'ERREUR: '+r.stderr[:100]}\")\n    except Exception as e:\n        print(f\"   ERREUR: {e}\")\n\n# ─────────────────────────────────────────────\n# 9. TRAIN ONE FOLD\n# ─────────────────────────────────────────────\ndef train_one_fold(fold_id, train_data, val_data, n_folds=5):\n    print(f\"\\n  Fold {fold_id+1}/{n_folds} ──────────────────────────\")\n    print(f\"     Train: {len(train_data)} | Val: {len(val_data)}\")\n\n    train_loader = DataLoader(\n        PatchDataset(train_data, train_tf),\n        batch_size=BATCH_SIZE, shuffle=True,\n        num_workers=NUM_WORKERS, pin_memory=True\n    )\n\n    model = build_model()\n\n    # Class weights\n    cnt_tr = Counter(r[\"label\"] for r in train_data)\n    w = torch.tensor(\n        [1.0 / max(cnt_tr.get(c, 1), 1) for c in range(6)],\n        dtype=torch.float\n    )\n    w = (w / w.sum() * 6).to(DEVICE)\n    criterion = nn.CrossEntropyLoss(weight=w, label_smoothing=0.1)\n\n    # Differential LR : UNI backbone vs Transformer\n    base_model = model.module if hasattr(model, \"module\") else model\n    optimizer  = optim.AdamW([\n        {\"params\": base_model.uni_backbone.parameters(), \"lr\": LR_UNI},\n        {\"params\": [p for n, p in base_model.named_parameters()\n                    if not n.startswith(\"uni_backbone\")], \"lr\": LR_TR},\n    ], weight_decay=1e-4)\n\n    scheduler = CosineAnnealingLR(optimizer, T_max=EPOCHS, eta_min=1e-7)\n\n    best_kappa = 0.0\n    no_improve = 0\n    best_ckpt  = f\"{SAVE_DIR}/stage{SHARD_ID}_uni_fold{fold_id+1}_best.pth\"\n\n    for epoch in range(1, EPOCHS + 1):\n        model.train()\n        t_loss, t_correct, t_total = 0.0, 0, 0\n        for imgs, labels in train_loader:\n            imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n            optimizer.zero_grad()\n            out  = model(imgs)\n            loss = criterion(out, labels)\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n            optimizer.step()\n            t_loss    += loss.item() * imgs.size(0)\n            t_correct += (out.argmax(1) == labels).sum().item()\n            t_total   += imgs.size(0)\n        scheduler.step()\n\n        wsi_kappa, n_wsis_val = wsi_level_qwk(model, val_data, DEVICE)\n        is_best = wsi_kappa > best_kappa\n\n        print(f\"     Ep[{epoch:02d}/{EPOCHS}] \"\n              f\"loss={t_loss/t_total:.4f} acc={t_correct/t_total:.3f} | \"\n              f\"WSI_QWK={wsi_kappa:.4f} ({n_wsis_val} WSIs)\"\n              + (\" *\" if is_best else \"\"))\n\n        if is_best:\n            best_kappa = wsi_kappa\n            no_improve = 0\n            torch.save({\n                \"epoch\":            epoch,\n                \"model_state_dict\": model.state_dict() if not hasattr(model, \"module\")\n                                    else model.module.state_dict(),\n                \"kappa\":            wsi_kappa,\n                \"shard_id\":         SHARD_ID,\n                \"fold_id\":          fold_id,\n                \"backbone\":         \"UNI\",\n                \"n_wsis\":           len(all_wsis),\n            }, best_ckpt)\n        else:\n            no_improve += 1\n            if no_improve >= PATIENCE:\n                print(f\"     Early stopping epoch {epoch}\")\n                break\n\n    print(f\"  OK Fold {fold_id+1} best WSI_QWK: {best_kappa:.4f}\")\n    push_to_kaggle(best_ckpt,\n        f\"Stage{SHARD_ID} UNI Fold{fold_id+1} QWK={best_kappa:.4f}\")\n    return best_kappa, best_ckpt\n\n# ─────────────────────────────────────────────\n# 10. 5-FOLD CV\n# ─────────────────────────────────────────────\nN_FOLDS = 5\nprint(f\"\\n{'='*60}\")\nprint(f\"Stage {SHARD_ID} UNI — {len(all_wsis)} WSIs — 5-Fold CV\")\nprint(f\"Backbone: UNI (MahmoodLab) — histopathology foundation model\")\nprint(f\"Comparison: ResNet-50 baseline mean=0.5764\")\nprint(f\"{'='*60}\")\n\nwsi_ids    = [r[\"wsi_id\"] for r in all_records]\nlabels_arr = [r[\"label\"]  for r in all_records]\nindices    = list(range(len(all_records)))\n\nsgkf = StratifiedGroupKFold(n_splits=N_FOLDS, shuffle=True, random_state=SEED)\n\nfold_kappas       = DONE_KAPPAS.copy()\nbest_overall      = -1.0\nbest_ckpt_overall = None\n\nfor fold_id, (train_idx, val_idx) in enumerate(\n        sgkf.split(indices, labels_arr, groups=wsi_ids)):\n\n    train_data = [all_records[i] for i in train_idx]\n    val_data   = [all_records[i] for i in val_idx]\n\n    if fold_id < START_FOLD:\n        continue\n\n    tw = set(all_records[i][\"wsi_id\"] for i in train_idx)\n    vw = set(all_records[i][\"wsi_id\"] for i in val_idx)\n    assert len(tw & vw) == 0, f\"DATA LEAKAGE fold {fold_id+1}!\"\n\n    kappa, ckpt_path = train_one_fold(fold_id, train_data, val_data, N_FOLDS)\n    fold_kappas.append(kappa)\n\n    if kappa > best_overall:\n        best_overall      = kappa\n        best_ckpt_overall = ckpt_path\n\n# ─────────────────────────────────────────────\n# 11. RESUME FINAL\n# ─────────────────────────────────────────────\nmean_kappa = np.mean(fold_kappas)\nstd_kappa  = np.std(fold_kappas)\n\nprint(f\"\\n{'='*60}\")\nprint(f\"OK Stage {SHARD_ID} UNI — 5-Fold CV Results\")\nprint(f\"{'='*60}\")\nfor i, k in enumerate(fold_kappas):\n    print(f\"   Fold {i+1}: QWK = {k:.4f}\")\nprint(f\"{'─'*40}\")\nprint(f\"   Mean QWK UNI      : {mean_kappa:.4f} +/- {std_kappa:.4f}\")\nprint(f\"   Mean QWK ResNet-50: 0.5764 +/- 0.0224 (baseline)\")\nprint(f\"   Gain UNI vs ResNet: {mean_kappa - 0.5764:+.4f}\")\nprint(f\"   Best QWK UNI      : {best_overall:.4f}\")\nprint(f\"   WSIs              : {len(all_wsis)}\")\nprint(f\"   SOTA [1]          : 0.934\")\nprint(f\"{'='*60}\")\n\nresults = {\n    \"stage\":          SHARD_ID,\n    \"backbone\":       \"UNI\",\n    \"n_wsis\":         len(all_wsis),\n    \"n_patches\":      len(all_records),\n    \"fold_kappas\":    fold_kappas,\n    \"mean_kappa\":     mean_kappa,\n    \"std_kappa\":      std_kappa,\n    \"best_kappa\":     best_overall,\n    \"resnet50_mean\":  0.5764,\n    \"gain_vs_resnet\": mean_kappa - 0.5764,\n    \"timestamp\":      datetime.datetime.now().isoformat(),\n}\nwith open(RESULTS_JSON, \"w\") as f:\n    json.dump(results, f, indent=2)\n\nif best_ckpt_overall:\n    stage_best = f\"{SAVE_DIR}/stage{SHARD_ID}_uni_best_qwk{best_overall:.4f}.pth\"\n    shutil.copy(best_ckpt_overall, stage_best)\n    print(f\"\\nOK Best checkpoint: {stage_best}\")\n    print(f\"\\nComparaison finale:\")\n    print(f\"  ResNet-50 Stage 0 : mean=0.5764, best=0.6116\")\n    print(f\"  UNI Stage 0       : mean={mean_kappa:.4f}, best={best_overall:.4f}\")\n    print(f\"  Gain              : {mean_kappa - 0.5764:+.4f}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\n=============================================================\nStage 3 — HybridUNITransformer — FROM SCRATCH\nUNI (MahmoodLab) foundation model + Transformer Encoder\nDonnées cumulatives :\n  Shard 0 : patches_400        (399  WSIs,  4,319 patches)\n  Shard 1 : panda-patches-run1 (1000 WSIs, 14,842 patches)\n  Shard 2 : panda-patches-run2 (999  WSIs, 14,867 patches)\n  Shard 3 : panda-patches-run3 (999  WSIs, 14,835 patches)\nTotal     : ~3,398 WSIs, ~48,863 patches\n\nComparaison :\n  ResNet-50 Stage 3 : mean=0.7280, best=0.7845\n  UNI Stage 3       : ???\n=============================================================\nAttache dans Kaggle :\n  khouloudbejaoui20/dataset\n  khouloudbejaoui20/panda-patches-run1\n  khouloudbejaoui20/panda-patches-run2\n  khouloudbejaoui20/panda-patches-run3\nGPU : T4×2\n\"\"\"\n\nimport os, glob, json, random, shutil, subprocess, datetime\nimport numpy as np\nfrom PIL import Image\nfrom collections import Counter\nfrom sklearn.model_selection import StratifiedGroupKFold\nfrom sklearn.metrics import cohen_kappa_score\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as T\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\n\n# ─────────────────────────────────────────────\n# 1. LOGIN HUGGINGFACE\n# ─────────────────────────────────────────────\nfrom huggingface_hub import login, whoami\nimport timm\n\nHF_TOKEN = \"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\"\nlogin(token=HF_TOKEN)\ninfo = whoami()\nprint(f\"OK HuggingFace: {info['name']}\")\n\n# ─────────────────────────────────────────────\n# 2. CONFIG\n# ─────────────────────────────────────────────\nSHARD_ID    = 3\nSTART_FOLD  = 0\nDONE_KAPPAS = []\n\n# UNI Stage 3 = FROM SCRATCH (pas de PREV_CKPT)\nPREV_CKPT   = None\n\n# FREEZE UNI backbone pour éviter overfitting !\nFREEZE_UNI  = True   # ← clé pour réduire overfitting\n\nN_FOLDS     = 5\nSEED        = 42\nBATCH_SIZE  = 8      # réduit pour éviter OOM avec UNI\nNUM_WORKERS = 2\nEPOCHS      = 20\nPATIENCE    = 5\nLR_UNI      = 1e-5   # utilisé seulement si FREEZE_UNI=False\nLR_TR       = 4e-5   # Transformer + MLP head\n\nrandom.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)\ntorch.backends.cudnn.deterministic = True\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Device     : {DEVICE}\")\nprint(f\"Stage      : {SHARD_ID} UNI — FROM SCRATCH\")\nprint(f\"Freeze UNI : {FREEZE_UNI}\")\nprint(f\"BATCH={BATCH_SIZE} | EPOCHS={EPOCHS} | PATIENCE={PATIENCE}\")\nprint(f\"LR_TR={LR_TR}\")\n\n# ─────────────────────────────────────────────\n# 3. CHEMINS PATCHES\n# ─────────────────────────────────────────────\nBASE = \"/kaggle/input/datasets/khouloudbejaoui20\"\n\nSHARD_PATHS = {\n    0: f\"{BASE}/dataset/patches_400-20260630T212405Z-3-001/patches_400\",\n    1: f\"{BASE}/panda-patches-run1/kaggle/working/patches_panda_new\",\n    2: f\"{BASE}/panda-patches-run2/kaggle/working/patches_panda_run2\",\n    3: f\"{BASE}/panda-patches-run3/kaggle/working/patches_panda_run3_correct\",\n}\n\nSAVE_DIR     = \"/kaggle/working\"\nKAGGLE_DS    = \"bejaouikhouloud/stage3-uni-results\"\nRESULTS_JSON = f\"{SAVE_DIR}/stage3_uni_cv_results.json\"\n\n# ─────────────────────────────────────────────\n# 4. VERIFICATION SHARDS\n# ─────────────────────────────────────────────\nprint(f\"\\n── Verification ────────────────────────\")\nfor shard_id, path in SHARD_PATHS.items():\n    if os.path.exists(path):\n        files = glob.glob(f\"{path}/*.png\")\n        if not files:\n            files = glob.glob(f\"{path}/**/*.png\", recursive=True)\n        print(f\"  OK Shard {shard_id}: {len(files)} patches\")\n    else:\n        print(f\"  MANQUANT Shard {shard_id}: {path}\")\n\n# ─────────────────────────────────────────────\n# 5. CHARGEMENT PATCHES\n# ─────────────────────────────────────────────\nprint(f\"\\n── Loading patches ──────────────────────\")\n\ndef load_patches(folder):\n    files = glob.glob(f\"{folder}/*.png\")\n    if not files:\n        files = glob.glob(f\"{folder}/**/*.png\", recursive=True)\n    records = []\n    for fp in files:\n        fname = os.path.basename(fp)\n        try:\n            parts  = fname.split(\"_slide\")\n            label  = int(parts[0].replace(\"grade\", \"\"))\n            wsi_id = parts[1].split(\"_\")[0]\n            records.append({\"path\": fp, \"wsi_id\": wsi_id, \"label\": label})\n        except:\n            continue\n    return records\n\nall_records = []\nfor shard_id, path in SHARD_PATHS.items():\n    if path and os.path.exists(path):\n        r = load_patches(path)\n        all_records.extend(r)\n        print(f\"  Shard {shard_id}: {len(r)} patches\")\n    else:\n        print(f\"  Shard {shard_id}: NON TROUVE\")\n\n# Deduplicate\nseen = set()\nunique_records = []\nfor r in all_records:\n    fname = os.path.basename(r[\"path\"])\n    if fname not in seen:\n        seen.add(fname)\n        unique_records.append(r)\n\nall_wsis = list(set(r[\"wsi_id\"] for r in unique_records))\nprint(f\"\\n  Total unique : {len(unique_records)} patches\")\nprint(f\"  Total WSIs  : {len(all_wsis)}\")\ncnt = Counter(r[\"label\"] for r in unique_records)\nfor c in range(6):\n    print(f\"    ISUP {c}: {cnt.get(c, 0)}\")\n\nassert len(all_wsis) > 3000, f\"Trop peu de WSIs ({len(all_wsis)})\"\nprint(f\"\\n  OK {len(all_wsis)} WSIs confirmes\")\n\n# ─────────────────────────────────────────────\n# 6. TRANSFORMS\n# ─────────────────────────────────────────────\nMEAN = [0.485, 0.456, 0.406]\nSTD  = [0.229, 0.224, 0.225]\n\ntrain_tf = T.Compose([\n    T.Resize((224, 224)),\n    T.RandomHorizontalFlip(),\n    T.RandomVerticalFlip(),\n    T.RandomRotation(90),\n    T.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.2, hue=0.1),\n    T.RandomApply([T.GaussianBlur(5)], p=0.3),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n    T.RandomApply([T.RandomErasing(p=0.5, scale=(0.02, 0.2))], p=0.2),\n])\n\nval_tf = T.Compose([\n    T.Resize((224, 224)),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n])\n\nclass PatchDataset(Dataset):\n    def __init__(self, records, transform):\n        self.records   = records\n        self.transform = transform\n    def __len__(self): return len(self.records)\n    def __getitem__(self, idx):\n        r = self.records[idx]\n        try:\n            img = Image.open(r[\"path\"]).convert(\"RGB\")\n        except:\n            img = Image.new(\"RGB\", (224, 224), (255, 255, 255))\n        return self.transform(img), int(r[\"label\"])\n\n# ─────────────────────────────────────────────\n# 7. MODELE — HybridUNITransformer\n# ─────────────────────────────────────────────\nclass HybridUNITransformer(nn.Module):\n    def __init__(self, num_classes=6, d_model=512, nhead=8,\n                 num_layers=2, dim_feedforward=2048, dropout=0.1,\n                 freeze_uni=True):\n        super().__init__()\n\n        # UNI backbone\n        self.uni_backbone = timm.create_model(\n            \"hf-hub:MahmoodLab/UNI\",\n            pretrained=True,\n            num_classes=0,\n            init_values=1e-5,\n            dynamic_img_size=True\n        )\n        uni_dim = self.uni_backbone.num_features  # 1024\n\n        # Freeze UNI si demandé\n        if freeze_uni:\n            for param in self.uni_backbone.parameters():\n                param.requires_grad = False\n            print(f\"  UNI backbone FREEZÉ\")\n        else:\n            print(f\"  UNI backbone trainable\")\n\n        # Projection\n        self.proj = nn.Sequential(\n            nn.Linear(uni_dim, d_model),\n            nn.LayerNorm(d_model)\n        )\n\n        # CLS token + positional embedding\n        self.cls_token = nn.Parameter(torch.zeros(1, 1, d_model))\n        self.pos_embed = nn.Parameter(torch.zeros(1, 2, d_model))\n\n        # Transformer Encoder\n        encoder_layer = nn.TransformerEncoderLayer(\n            d_model=d_model, nhead=nhead,\n            dim_feedforward=dim_feedforward,\n            dropout=dropout, batch_first=True\n        )\n        self.transformer = nn.TransformerEncoder(\n            encoder_layer, num_layers=num_layers,\n            norm=nn.LayerNorm(d_model)\n        )\n\n        # MLP Head\n        self.classifier = nn.Sequential(\n            nn.Linear(d_model, 256),\n            nn.ReLU(),\n            nn.Dropout(dropout),\n            nn.Linear(256, num_classes)\n        )\n\n        nn.init.trunc_normal_(self.cls_token, std=0.02)\n        nn.init.trunc_normal_(self.pos_embed, std=0.02)\n\n    def forward(self, x):\n        with torch.no_grad() if not self.uni_backbone.training else torch.enable_grad():\n            f = self.uni_backbone(x)\n        f   = self.proj(f).unsqueeze(1)\n        cls = self.cls_token.expand(x.size(0), -1, -1)\n        seq = torch.cat([cls, f], dim=1) + self.pos_embed\n        return self.classifier(self.transformer(seq)[:, 0])\n\ndef build_model():\n    # Libère mémoire GPU\n    torch.cuda.empty_cache()\n    import gc; gc.collect()\n\n    model  = HybridUNITransformer(freeze_uni=FREEZE_UNI).to(DEVICE)\n    n_gpus = torch.cuda.device_count()\n    if n_gpus > 1:\n        model = nn.DataParallel(model)\n        print(f\"  OK DataParallel: {n_gpus} GPUs\")\n\n    # Paramètres trainables\n    n_total    = sum(p.numel() for p in model.parameters())\n    n_trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    print(f\"  Total params    : {n_total:,}\")\n    print(f\"  Trainable params: {n_trainable:,}\")\n    return model\n\n# ─────────────────────────────────────────────\n# 8. WSI-LEVEL QWK\n# ─────────────────────────────────────────────\ndef wsi_level_qwk(model, val_data, device):\n    model.eval()\n    wsi_preds  = {}\n    wsi_labels = {}\n    all_preds  = []\n    all_labels = []\n    loader = DataLoader(\n        PatchDataset(val_data, val_tf),\n        batch_size=BATCH_SIZE, shuffle=False,\n        num_workers=NUM_WORKERS\n    )\n    with torch.no_grad():\n        for imgs, labels in loader:\n            imgs  = imgs.to(device)\n            preds = model(imgs).argmax(1).cpu().numpy()\n            all_preds.extend(preds)\n            all_labels.extend(labels.numpy())\n    for i, r in enumerate(val_data):\n        wid = r[\"wsi_id\"]\n        if wid not in wsi_preds:\n            wsi_preds[wid]  = []\n            wsi_labels[wid] = r[\"label\"]\n        wsi_preds[wid].append(all_preds[i])\n    wf = []; wl = []\n    for wid in wsi_preds:\n        votes = wsi_preds[wid]\n        wf.append(max(set(votes), key=votes.count))\n        wl.append(wsi_labels[wid])\n    wsi_qwk = cohen_kappa_score(wl, wf, weights=\"quadratic\")\n    return wsi_qwk, len(wf)\n\n# ─────────────────────────────────────────────\n# 9. PUSH KAGGLE\n# ─────────────────────────────────────────────\ndef push_to_kaggle(ckpt_path, message):\n    try:\n        tmp = f\"{SAVE_DIR}/_push_uni3\"\n        os.makedirs(tmp, exist_ok=True)\n        for c in glob.glob(f\"{SAVE_DIR}/stage{SHARD_ID}_uni*.pth\"):\n            shutil.copy(c, tmp)\n        shutil.copy(ckpt_path, tmp)\n        if os.path.exists(RESULTS_JSON):\n            shutil.copy(RESULTS_JSON, tmp)\n        json.dump({\n            \"title\": \"Stage3 UNI Results Bejaoui 2026\",\n            \"id\": KAGGLE_DS,\n            \"licenses\": [{\"name\": \"CC0-1.0\"}]\n        }, open(f\"{tmp}/dataset-metadata.json\", \"w\"))\n        r = subprocess.run(\n            [\"kaggle\", \"datasets\", \"create\",\n             \"-p\", tmp, \"--dir-mode\", \"skip\"],\n            capture_output=True, text=True, timeout=600\n        )\n        if r.returncode != 0:\n            r = subprocess.run(\n                [\"kaggle\", \"datasets\", \"version\",\n                 \"-p\", tmp, \"-m\", message, \"--dir-mode\", \"skip\"],\n                capture_output=True, text=True, timeout=600\n            )\n        print(f\"   {'OK!' if r.returncode==0 else 'ERREUR: '+r.stderr[:100]}\")\n    except Exception as e:\n        print(f\"   ERREUR: {e}\")\n\n# ─────────────────────────────────────────────\n# 10. TRAIN ONE FOLD\n# ─────────────────────────────────────────────\ndef train_one_fold(fold_id, train_data, val_data, n_folds=5):\n    print(f\"\\n  Fold {fold_id+1}/{n_folds} ──────────────────────────\")\n    print(f\"     Train: {len(train_data)} | Val: {len(val_data)}\")\n    print(f\"     Freeze UNI: {FREEZE_UNI}\")\n\n    train_loader = DataLoader(\n        PatchDataset(train_data, train_tf),\n        batch_size=BATCH_SIZE, shuffle=True,\n        num_workers=NUM_WORKERS, pin_memory=True\n    )\n\n    model = build_model()\n\n    # Class weights\n    cnt_tr = Counter(r[\"label\"] for r in train_data)\n    w = torch.tensor(\n        [1.0 / max(cnt_tr.get(c, 1), 1) for c in range(6)],\n        dtype=torch.float\n    )\n    w = (w / w.sum() * 6).to(DEVICE)\n    criterion = nn.CrossEntropyLoss(weight=w, label_smoothing=0.1)\n\n    # Optimizer — seulement les paramètres trainables\n    base_model = model.module if hasattr(model, \"module\") else model\n\n    if FREEZE_UNI:\n        # Entraîne seulement proj + transformer + classifier\n        trainable_params = [p for p in model.parameters() if p.requires_grad]\n        optimizer = optim.AdamW(trainable_params, lr=LR_TR, weight_decay=1e-4)\n        print(f\"     Optimizer: AdamW lr={LR_TR} (trainable only)\")\n    else:\n        optimizer = optim.AdamW([\n            {\"params\": base_model.uni_backbone.parameters(), \"lr\": LR_UNI},\n            {\"params\": [p for n, p in base_model.named_parameters()\n                        if not n.startswith(\"uni_backbone\")], \"lr\": LR_TR},\n        ], weight_decay=1e-4)\n        print(f\"     Optimizer: AdamW LR_UNI={LR_UNI} LR_TR={LR_TR}\")\n\n    scheduler = CosineAnnealingLR(optimizer, T_max=EPOCHS, eta_min=1e-7)\n\n    best_kappa = 0.0\n    no_improve = 0\n    best_ckpt  = f\"{SAVE_DIR}/stage{SHARD_ID}_uni_fold{fold_id+1}_best.pth\"\n\n    for epoch in range(1, EPOCHS + 1):\n        model.train()\n        if FREEZE_UNI:\n            # S'assure que UNI reste en eval mode\n            base_model = model.module if hasattr(model, \"module\") else model\n            base_model.uni_backbone.eval()\n\n        t_loss, t_correct, t_total = 0.0, 0, 0\n        for imgs, labels in train_loader:\n            imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n            optimizer.zero_grad()\n            out  = model(imgs)\n            loss = criterion(out, labels)\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n            optimizer.step()\n            t_loss    += loss.item() * imgs.size(0)\n            t_correct += (out.argmax(1) == labels).sum().item()\n            t_total   += imgs.size(0)\n        scheduler.step()\n\n        wsi_kappa, n_wsis_val = wsi_level_qwk(model, val_data, DEVICE)\n        is_best = wsi_kappa > best_kappa\n\n        print(f\"     Ep[{epoch:02d}/{EPOCHS}] \"\n              f\"loss={t_loss/t_total:.4f} acc={t_correct/t_total:.3f} | \"\n              f\"WSI_QWK={wsi_kappa:.4f} ({n_wsis_val} WSIs)\"\n              + (\" *\" if is_best else \"\"))\n\n        if is_best:\n            best_kappa = wsi_kappa\n            no_improve = 0\n            torch.save({\n                \"epoch\":            epoch,\n                \"model_state_dict\": model.state_dict() if not hasattr(model, \"module\")\n                                    else model.module.state_dict(),\n                \"kappa\":            wsi_kappa,\n                \"shard_id\":         SHARD_ID,\n                \"fold_id\":          fold_id,\n                \"backbone\":         \"UNI\",\n                \"freeze_uni\":       FREEZE_UNI,\n                \"n_wsis\":           len(all_wsis),\n            }, best_ckpt)\n        else:\n            no_improve += 1\n            if no_improve >= PATIENCE:\n                print(f\"     Early stopping epoch {epoch}\")\n                break\n\n    print(f\"  OK Fold {fold_id+1} best WSI_QWK: {best_kappa:.4f}\")\n    push_to_kaggle(best_ckpt,\n        f\"Stage{SHARD_ID} UNI Fold{fold_id+1} QWK={best_kappa:.4f}\")\n    return best_kappa, best_ckpt\n\n# ─────────────────────────────────────────────\n# 11. 5-FOLD CV\n# ─────────────────────────────────────────────\nprint(f\"\\n{'='*60}\")\nprint(f\"Stage {SHARD_ID} UNI — {len(all_wsis)} WSIs — 5-Fold CV\")\nprint(f\"Backbone: UNI (MahmoodLab) | Freeze: {FREEZE_UNI}\")\nprint(f\"ResNet-50 Stage 3 baseline: mean=0.7280, best=0.7845\")\nprint(f\"{'='*60}\")\n\nwsi_ids    = [r[\"wsi_id\"] for r in unique_records]\nlabels_arr = [r[\"label\"]  for r in unique_records]\nindices    = list(range(len(unique_records)))\n\nsgkf = StratifiedGroupKFold(n_splits=N_FOLDS, shuffle=True, random_state=SEED)\n\nfold_kappas       = DONE_KAPPAS.copy()\nbest_overall      = max(DONE_KAPPAS) if DONE_KAPPAS else -1.0\nbest_ckpt_overall = None\n\nfor fold_id, (train_idx, val_idx) in enumerate(\n        sgkf.split(indices, labels_arr, groups=wsi_ids)):\n\n    if fold_id < START_FOLD:\n        continue\n\n    train_data = [unique_records[i] for i in train_idx]\n    val_data   = [unique_records[i] for i in val_idx]\n\n    tw = set(unique_records[i][\"wsi_id\"] for i in train_idx)\n    vw = set(unique_records[i][\"wsi_id\"] for i in val_idx)\n    assert len(tw & vw) == 0, f\"DATA LEAKAGE fold {fold_id+1}!\"\n\n    kappa, ckpt_path = train_one_fold(fold_id, train_data, val_data, N_FOLDS)\n    fold_kappas.append(kappa)\n\n    if kappa > best_overall:\n        best_overall      = kappa\n        best_ckpt_overall = ckpt_path\n\n# ─────────────────────────────────────────────\n# 12. RESUME FINAL\n# ─────────────────────────────────────────────\nmean_kappa = np.mean(fold_kappas)\nstd_kappa  = np.std(fold_kappas)\n\nprint(f\"\\n{'='*60}\")\nprint(f\"OK Stage {SHARD_ID} UNI — 5-Fold CV Results\")\nprint(f\"{'='*60}\")\nfor i, k in enumerate(fold_kappas):\n    print(f\"   Fold {i+1}: QWK = {k:.4f}\")\nprint(f\"{'─'*40}\")\nprint(f\"   Mean QWK UNI        : {mean_kappa:.4f} +/- {std_kappa:.4f}\")\nprint(f\"   Best QWK UNI        : {best_overall:.4f}\")\nprint(f\"   ResNet-50 Stage 3   : mean=0.7280, best=0.7845\")\nprint(f\"   Gain vs ResNet-50   : {mean_kappa - 0.7280:+.4f}\")\nprint(f\"   WSIs                : {len(all_wsis)}\")\nprint(f\"   SOTA [1]            : 0.934\")\nprint(f\"{'='*60}\")\n\nresults = {\n    \"stage\":           SHARD_ID,\n    \"backbone\":        \"UNI\",\n    \"freeze_uni\":      FREEZE_UNI,\n    \"n_wsis\":          len(all_wsis),\n    \"n_patches\":       len(unique_records),\n    \"fold_kappas\":     fold_kappas,\n    \"mean_kappa\":      mean_kappa,\n    \"std_kappa\":       std_kappa,\n    \"best_kappa\":      best_overall,\n    \"resnet50_mean\":   0.7280,\n    \"resnet50_best\":   0.7845,\n    \"gain_vs_resnet\":  mean_kappa - 0.7280,\n    \"timestamp\":       datetime.datetime.now().isoformat(),\n}\nwith open(RESULTS_JSON, \"w\") as f:\n    json.dump(results, f, indent=2)\n\nif best_ckpt_overall:\n    stage_best = f\"{SAVE_DIR}/stage{SHARD_ID}_uni_best_qwk{best_overall:.4f}.pth\"\n    shutil.copy(best_ckpt_overall, stage_best)\n    print(f\"\\nOK Best checkpoint: {stage_best}\")\n    print(f\"\\nComparaison finale:\")\n    print(f\"  ResNet-50 Stage 3 : mean=0.7280, best=0.7845\")\n    print(f\"  UNI Stage 3       : mean={mean_kappa:.4f}, best={best_overall:.4f}\")\n    print(f\"  Gain mean         : {mean_kappa - 0.7280:+.4f}\")\n    print(f\"  Gain best         : {best_overall - 0.7845:+.4f}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\n=============================================================\nStage 3 — HybridUNITransformer — FROM SCRATCH\nUNI (MahmoodLab) foundation model + Transformer Encoder\nDonnées cumulatives :\n  Shard 0 : patches_400        (399  WSIs,  4,319 patches)\n  Shard 1 : panda-patches-run1 (1000 WSIs, 14,842 patches)\n  Shard 2 : panda-patches-run2 (999  WSIs, 14,867 patches)\n  Shard 3 : panda-patches-run3 (999  WSIs, 14,835 patches)\nTotal     : ~3,398 WSIs, ~48,863 patches\n\nComparaison :\n  ResNet-50 Stage 3 : mean=0.7280, best=0.7845\n  UNI Stage 3       : ???\n=============================================================\nAttache dans Kaggle :\n  khouloudbejaoui20/dataset\n  khouloudbejaoui20/panda-patches-run1\n  khouloudbejaoui20/panda-patches-run2\n  khouloudbejaoui20/panda-patches-run3\nGPU : T4×2\n\"\"\"\n\nimport os, glob, json, random, shutil, subprocess, datetime\nimport numpy as np\nfrom PIL import Image\nfrom collections import Counter\nfrom sklearn.model_selection import StratifiedGroupKFold\nfrom sklearn.metrics import cohen_kappa_score\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as T\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\n\n# ─────────────────────────────────────────────\n# 1. LOGIN HUGGINGFACE\n# ─────────────────────────────────────────────\nfrom huggingface_hub import login, whoami\nimport timm\n\nHF_TOKEN = \"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\"\nlogin(token=HF_TOKEN)\ninfo = whoami()\nprint(f\"OK HuggingFace: {info['name']}\")\n\n# ─────────────────────────────────────────────\n# 2. CONFIG\n# ─────────────────────────────────────────────\nSHARD_ID    = 3\nSTART_FOLD  = 0\nDONE_KAPPAS = []\n\n# UNI Stage 3 = FROM SCRATCH (pas de PREV_CKPT)\nPREV_CKPT   = None\n\n# FREEZE UNI backbone pour éviter overfitting !\nFREEZE_UNI  = False   # ← clé pour réduire overfitting\n\nN_FOLDS     = 5\nSEED        = 42\nBATCH_SIZE  = 4      # réduit pour éviter OOM avec UNI\nNUM_WORKERS = 2\nEPOCHS      = 20\nPATIENCE    = 5\nLR_UNI      = 1e-6   # utilisé seulement si FREEZE_UNI=False\nLR_TR       = 4e-5   # Transformer + MLP head\n\nrandom.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)\ntorch.backends.cudnn.deterministic = True\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Device     : {DEVICE}\")\nprint(f\"Stage      : {SHARD_ID} UNI — FROM SCRATCH\")\nprint(f\"Freeze UNI : {FREEZE_UNI}\")\nprint(f\"BATCH={BATCH_SIZE} | EPOCHS={EPOCHS} | PATIENCE={PATIENCE}\")\nprint(f\"LR_TR={LR_TR}\")\n\n# ─────────────────────────────────────────────\n# 3. CHEMINS PATCHES\n# ─────────────────────────────────────────────\nBASE = \"/kaggle/input/datasets/khouloudbejaoui20\"\n\nSHARD_PATHS = {\n    0: f\"{BASE}/dataset/patches_400-20260630T212405Z-3-001/patches_400\",\n    1: f\"{BASE}/panda-patches-run1/kaggle/working/patches_panda_new\",\n    2: f\"{BASE}/panda-patches-run2/kaggle/working/patches_panda_run2\",\n    3: f\"{BASE}/panda-patches-run3/kaggle/working/patches_panda_run3_correct\",\n}\n\nSAVE_DIR     = \"/kaggle/working\"\nKAGGLE_DS    = \"bejaouikhouloud/stage3-uni-results\"\nRESULTS_JSON = f\"{SAVE_DIR}/stage3_uni_cv_results.json\"\n\n# ─────────────────────────────────────────────\n# 4. VERIFICATION SHARDS\n# ─────────────────────────────────────────────\nprint(f\"\\n── Verification ────────────────────────\")\nfor shard_id, path in SHARD_PATHS.items():\n    if os.path.exists(path):\n        files = glob.glob(f\"{path}/*.png\")\n        if not files:\n            files = glob.glob(f\"{path}/**/*.png\", recursive=True)\n        print(f\"  OK Shard {shard_id}: {len(files)} patches\")\n    else:\n        print(f\"  MANQUANT Shard {shard_id}: {path}\")\n\n# ─────────────────────────────────────────────\n# 5. CHARGEMENT PATCHES\n# ─────────────────────────────────────────────\nprint(f\"\\n── Loading patches ──────────────────────\")\n\ndef load_patches(folder):\n    files = glob.glob(f\"{folder}/*.png\")\n    if not files:\n        files = glob.glob(f\"{folder}/**/*.png\", recursive=True)\n    records = []\n    for fp in files:\n        fname = os.path.basename(fp)\n        try:\n            parts  = fname.split(\"_slide\")\n            label  = int(parts[0].replace(\"grade\", \"\"))\n            wsi_id = parts[1].split(\"_\")[0]\n            records.append({\"path\": fp, \"wsi_id\": wsi_id, \"label\": label})\n        except:\n            continue\n    return records\n\nall_records = []\nfor shard_id, path in SHARD_PATHS.items():\n    if path and os.path.exists(path):\n        r = load_patches(path)\n        all_records.extend(r)\n        print(f\"  Shard {shard_id}: {len(r)} patches\")\n    else:\n        print(f\"  Shard {shard_id}: NON TROUVE\")\n\n# Deduplicate\nseen = set()\nunique_records = []\nfor r in all_records:\n    fname = os.path.basename(r[\"path\"])\n    if fname not in seen:\n        seen.add(fname)\n        unique_records.append(r)\n\nall_wsis = list(set(r[\"wsi_id\"] for r in unique_records))\nprint(f\"\\n  Total unique : {len(unique_records)} patches\")\nprint(f\"  Total WSIs  : {len(all_wsis)}\")\ncnt = Counter(r[\"label\"] for r in unique_records)\nfor c in range(6):\n    print(f\"    ISUP {c}: {cnt.get(c, 0)}\")\n\nassert len(all_wsis) > 3000, f\"Trop peu de WSIs ({len(all_wsis)})\"\nprint(f\"\\n  OK {len(all_wsis)} WSIs confirmes\")\n\n# ─────────────────────────────────────────────\n# 6. TRANSFORMS\n# ─────────────────────────────────────────────\nMEAN = [0.485, 0.456, 0.406]\nSTD  = [0.229, 0.224, 0.225]\n\ntrain_tf = T.Compose([\n    T.Resize((224, 224)),\n    T.RandomHorizontalFlip(),\n    T.RandomVerticalFlip(),\n    T.RandomRotation(90),\n    T.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.2, hue=0.1),\n    T.RandomApply([T.GaussianBlur(5)], p=0.3),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n    T.RandomApply([T.RandomErasing(p=0.5, scale=(0.02, 0.2))], p=0.2),\n])\n\nval_tf = T.Compose([\n    T.Resize((224, 224)),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n])\n\nclass PatchDataset(Dataset):\n    def __init__(self, records, transform):\n        self.records   = records\n        self.transform = transform\n    def __len__(self): return len(self.records)\n    def __getitem__(self, idx):\n        r = self.records[idx]\n        try:\n            img = Image.open(r[\"path\"]).convert(\"RGB\")\n        except:\n            img = Image.new(\"RGB\", (224, 224), (255, 255, 255))\n        return self.transform(img), int(r[\"label\"])\n\n# ─────────────────────────────────────────────\n# 7. MODELE — HybridUNITransformer\n# ─────────────────────────────────────────────\nclass HybridUNITransformer(nn.Module):\n    def __init__(self, num_classes=6, d_model=512, nhead=8,\n                 num_layers=2, dim_feedforward=2048, dropout=0.1,\n                 freeze_uni=True):\n        super().__init__()\n\n        # UNI backbone\n        self.uni_backbone = timm.create_model(\n            \"hf-hub:MahmoodLab/UNI\",\n            pretrained=True,\n            num_classes=0,\n            init_values=1e-5,\n            dynamic_img_size=True\n        )\n        uni_dim = self.uni_backbone.num_features  # 1024\n\n        # Freeze UNI si demandé\n        if freeze_uni:\n            for param in self.uni_backbone.parameters():\n                param.requires_grad = False\n            print(f\"  UNI backbone FREEZÉ\")\n        else:\n            print(f\"  UNI backbone trainable\")\n\n        # Projection\n        self.proj = nn.Sequential(\n            nn.Linear(uni_dim, d_model),\n            nn.LayerNorm(d_model)\n        )\n\n        # CLS token + positional embedding\n        self.cls_token = nn.Parameter(torch.zeros(1, 1, d_model))\n        self.pos_embed = nn.Parameter(torch.zeros(1, 2, d_model))\n\n        # Transformer Encoder\n        encoder_layer = nn.TransformerEncoderLayer(\n            d_model=d_model, nhead=nhead,\n            dim_feedforward=dim_feedforward,\n            dropout=dropout, batch_first=True\n        )\n        self.transformer = nn.TransformerEncoder(\n            encoder_layer, num_layers=num_layers,\n            norm=nn.LayerNorm(d_model)\n        )\n\n        # MLP Head\n        self.classifier = nn.Sequential(\n            nn.Linear(d_model, 256),\n            nn.ReLU(),\n            nn.Dropout(dropout),\n            nn.Linear(256, num_classes)\n        )\n\n        nn.init.trunc_normal_(self.cls_token, std=0.02)\n        nn.init.trunc_normal_(self.pos_embed, std=0.02)\n\n    def forward(self, x):\n        with torch.no_grad() if not self.uni_backbone.training else torch.enable_grad():\n            f = self.uni_backbone(x)\n        f   = self.proj(f).unsqueeze(1)\n        cls = self.cls_token.expand(x.size(0), -1, -1)\n        seq = torch.cat([cls, f], dim=1) + self.pos_embed\n        return self.classifier(self.transformer(seq)[:, 0])\n\ndef build_model():\n    # Libère mémoire GPU\n    torch.cuda.empty_cache()\n    import gc; gc.collect()\n\n    model  = HybridUNITransformer(freeze_uni=FREEZE_UNI).to(DEVICE)\n    n_gpus = torch.cuda.device_count()\n    if n_gpus > 1:\n        model = nn.DataParallel(model)\n        print(f\"  OK DataParallel: {n_gpus} GPUs\")\n\n    # Paramètres trainables\n    n_total    = sum(p.numel() for p in model.parameters())\n    n_trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    print(f\"  Total params    : {n_total:,}\")\n    print(f\"  Trainable params: {n_trainable:,}\")\n    return model\n\n# ─────────────────────────────────────────────\n# 8. WSI-LEVEL QWK\n# ─────────────────────────────────────────────\ndef wsi_level_qwk(model, val_data, device):\n    model.eval()\n    wsi_preds  = {}\n    wsi_labels = {}\n    all_preds  = []\n    all_labels = []\n    loader = DataLoader(\n        PatchDataset(val_data, val_tf),\n        batch_size=BATCH_SIZE, shuffle=False,\n        num_workers=NUM_WORKERS\n    )\n    with torch.no_grad():\n        for imgs, labels in loader:\n            imgs  = imgs.to(device)\n            preds = model(imgs).argmax(1).cpu().numpy()\n            all_preds.extend(preds)\n            all_labels.extend(labels.numpy())\n    for i, r in enumerate(val_data):\n        wid = r[\"wsi_id\"]\n        if wid not in wsi_preds:\n            wsi_preds[wid]  = []\n            wsi_labels[wid] = r[\"label\"]\n        wsi_preds[wid].append(all_preds[i])\n    wf = []; wl = []\n    for wid in wsi_preds:\n        votes = wsi_preds[wid]\n        wf.append(max(set(votes), key=votes.count))\n        wl.append(wsi_labels[wid])\n    wsi_qwk = cohen_kappa_score(wl, wf, weights=\"quadratic\")\n    return wsi_qwk, len(wf)\n\n# ─────────────────────────────────────────────\n# 9. PUSH KAGGLE\n# ─────────────────────────────────────────────\ndef push_to_kaggle(ckpt_path, message):\n    try:\n        tmp = f\"{SAVE_DIR}/_push_uni3\"\n        os.makedirs(tmp, exist_ok=True)\n        for c in glob.glob(f\"{SAVE_DIR}/stage{SHARD_ID}_uni*.pth\"):\n            shutil.copy(c, tmp)\n        shutil.copy(ckpt_path, tmp)\n        if os.path.exists(RESULTS_JSON):\n            shutil.copy(RESULTS_JSON, tmp)\n        json.dump({\n            \"title\": \"Stage3 UNI Results Bejaoui 2026\",\n            \"id\": KAGGLE_DS,\n            \"licenses\": [{\"name\": \"CC0-1.0\"}]\n        }, open(f\"{tmp}/dataset-metadata.json\", \"w\"))\n        r = subprocess.run(\n            [\"kaggle\", \"datasets\", \"create\",\n             \"-p\", tmp, \"--dir-mode\", \"skip\"],\n            capture_output=True, text=True, timeout=600\n        )\n        if r.returncode != 0:\n            r = subprocess.run(\n                [\"kaggle\", \"datasets\", \"version\",\n                 \"-p\", tmp, \"-m\", message, \"--dir-mode\", \"skip\"],\n                capture_output=True, text=True, timeout=600\n            )\n        print(f\"   {'OK!' if r.returncode==0 else 'ERREUR: '+r.stderr[:100]}\")\n    except Exception as e:\n        print(f\"   ERREUR: {e}\")\n\n# ─────────────────────────────────────────────\n# 10. TRAIN ONE FOLD\n# ─────────────────────────────────────────────\ndef train_one_fold(fold_id, train_data, val_data, n_folds=5):\n    print(f\"\\n  Fold {fold_id+1}/{n_folds} ──────────────────────────\")\n    print(f\"     Train: {len(train_data)} | Val: {len(val_data)}\")\n    print(f\"     Freeze UNI: {FREEZE_UNI}\")\n\n    train_loader = DataLoader(\n        PatchDataset(train_data, train_tf),\n        batch_size=BATCH_SIZE, shuffle=True,\n        num_workers=NUM_WORKERS, pin_memory=True\n    )\n\n    model = build_model()\n\n    # Class weights\n    cnt_tr = Counter(r[\"label\"] for r in train_data)\n    w = torch.tensor(\n        [1.0 / max(cnt_tr.get(c, 1), 1) for c in range(6)],\n        dtype=torch.float\n    )\n    w = (w / w.sum() * 6).to(DEVICE)\n    criterion = nn.CrossEntropyLoss(weight=w, label_smoothing=0.1)\n\n    # Optimizer — seulement les paramètres trainables\n    base_model = model.module if hasattr(model, \"module\") else model\n\n    if FREEZE_UNI:\n        # Entraîne seulement proj + transformer + classifier\n        trainable_params = [p for p in model.parameters() if p.requires_grad]\n        optimizer = optim.AdamW(trainable_params, lr=LR_TR, weight_decay=1e-4)\n        print(f\"     Optimizer: AdamW lr={LR_TR} (trainable only)\")\n    else:\n        optimizer = optim.AdamW([\n            {\"params\": base_model.uni_backbone.parameters(), \"lr\": LR_UNI},\n            {\"params\": [p for n, p in base_model.named_parameters()\n                        if not n.startswith(\"uni_backbone\")], \"lr\": LR_TR},\n        ], weight_decay=1e-4)\n        print(f\"     Optimizer: AdamW LR_UNI={LR_UNI} LR_TR={LR_TR}\")\n\n    scheduler = CosineAnnealingLR(optimizer, T_max=EPOCHS, eta_min=1e-7)\n\n    best_kappa = 0.0\n    no_improve = 0\n    best_ckpt  = f\"{SAVE_DIR}/stage{SHARD_ID}_uni_fold{fold_id+1}_best.pth\"\n\n    for epoch in range(1, EPOCHS + 1):\n        model.train()\n        if FREEZE_UNI:\n            # S'assure que UNI reste en eval mode\n            base_model = model.module if hasattr(model, \"module\") else model\n            base_model.uni_backbone.eval()\n\n        t_loss, t_correct, t_total = 0.0, 0, 0\n        for imgs, labels in train_loader:\n            imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n            optimizer.zero_grad()\n            out  = model(imgs)\n            loss = criterion(out, labels)\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n            optimizer.step()\n            t_loss    += loss.item() * imgs.size(0)\n            t_correct += (out.argmax(1) == labels).sum().item()\n            t_total   += imgs.size(0)\n        scheduler.step()\n\n        wsi_kappa, n_wsis_val = wsi_level_qwk(model, val_data, DEVICE)\n        is_best = wsi_kappa > best_kappa\n\n        print(f\"     Ep[{epoch:02d}/{EPOCHS}] \"\n              f\"loss={t_loss/t_total:.4f} acc={t_correct/t_total:.3f} | \"\n              f\"WSI_QWK={wsi_kappa:.4f} ({n_wsis_val} WSIs)\"\n              + (\" *\" if is_best else \"\"))\n\n        if is_best:\n            best_kappa = wsi_kappa\n            no_improve = 0\n            torch.save({\n                \"epoch\":            epoch,\n                \"model_state_dict\": model.state_dict() if not hasattr(model, \"module\")\n                                    else model.module.state_dict(),\n                \"kappa\":            wsi_kappa,\n                \"shard_id\":         SHARD_ID,\n                \"fold_id\":          fold_id,\n                \"backbone\":         \"UNI\",\n                \"freeze_uni\":       FREEZE_UNI,\n                \"n_wsis\":           len(all_wsis),\n            }, best_ckpt)\n        else:\n            no_improve += 1\n            if no_improve >= PATIENCE:\n                print(f\"     Early stopping epoch {epoch}\")\n                break\n\n    print(f\"  OK Fold {fold_id+1} best WSI_QWK: {best_kappa:.4f}\")\n    push_to_kaggle(best_ckpt,\n        f\"Stage{SHARD_ID} UNI Fold{fold_id+1} QWK={best_kappa:.4f}\")\n    return best_kappa, best_ckpt\n\n# ─────────────────────────────────────────────\n# 11. 5-FOLD CV\n# ─────────────────────────────────────────────\nprint(f\"\\n{'='*60}\")\nprint(f\"Stage {SHARD_ID} UNI — {len(all_wsis)} WSIs — 5-Fold CV\")\nprint(f\"Backbone: UNI (MahmoodLab) | Freeze: {FREEZE_UNI}\")\nprint(f\"ResNet-50 Stage 3 baseline: mean=0.7280, best=0.7845\")\nprint(f\"{'='*60}\")\n\nwsi_ids    = [r[\"wsi_id\"] for r in unique_records]\nlabels_arr = [r[\"label\"]  for r in unique_records]\nindices    = list(range(len(unique_records)))\n\nsgkf = StratifiedGroupKFold(n_splits=N_FOLDS, shuffle=True, random_state=SEED)\n\nfold_kappas       = DONE_KAPPAS.copy()\nbest_overall      = max(DONE_KAPPAS) if DONE_KAPPAS else -1.0\nbest_ckpt_overall = None\n\nfor fold_id, (train_idx, val_idx) in enumerate(\n        sgkf.split(indices, labels_arr, groups=wsi_ids)):\n\n    if fold_id < START_FOLD:\n        continue\n\n    train_data = [unique_records[i] for i in train_idx]\n    val_data   = [unique_records[i] for i in val_idx]\n\n    tw = set(unique_records[i][\"wsi_id\"] for i in train_idx)\n    vw = set(unique_records[i][\"wsi_id\"] for i in val_idx)\n    assert len(tw & vw) == 0, f\"DATA LEAKAGE fold {fold_id+1}!\"\n\n    kappa, ckpt_path = train_one_fold(fold_id, train_data, val_data, N_FOLDS)\n    fold_kappas.append(kappa)\n\n    if kappa > best_overall:\n        best_overall      = kappa\n        best_ckpt_overall = ckpt_path\n\n# ─────────────────────────────────────────────\n# 12. RESUME FINAL\n# ─────────────────────────────────────────────\nmean_kappa = np.mean(fold_kappas)\nstd_kappa  = np.std(fold_kappas)\n\nprint(f\"\\n{'='*60}\")\nprint(f\"OK Stage {SHARD_ID} UNI — 5-Fold CV Results\")\nprint(f\"{'='*60}\")\nfor i, k in enumerate(fold_kappas):\n    print(f\"   Fold {i+1}: QWK = {k:.4f}\")\nprint(f\"{'─'*40}\")\nprint(f\"   Mean QWK UNI        : {mean_kappa:.4f} +/- {std_kappa:.4f}\")\nprint(f\"   Best QWK UNI        : {best_overall:.4f}\")\nprint(f\"   ResNet-50 Stage 3   : mean=0.7280, best=0.7845\")\nprint(f\"   Gain vs ResNet-50   : {mean_kappa - 0.7280:+.4f}\")\nprint(f\"   WSIs                : {len(all_wsis)}\")\nprint(f\"   SOTA [1]            : 0.934\")\nprint(f\"{'='*60}\")\n\nresults = {\n    \"stage\":           SHARD_ID,\n    \"backbone\":        \"UNI\",\n    \"freeze_uni\":      FREEZE_UNI,\n    \"n_wsis\":          len(all_wsis),\n    \"n_patches\":       len(unique_records),\n    \"fold_kappas\":     fold_kappas,\n    \"mean_kappa\":      mean_kappa,\n    \"std_kappa\":       std_kappa,\n    \"best_kappa\":      best_overall,\n    \"resnet50_mean\":   0.7280,\n    \"resnet50_best\":   0.7845,\n    \"gain_vs_resnet\":  mean_kappa - 0.7280,\n    \"timestamp\":       datetime.datetime.now().isoformat(),\n}\nwith open(RESULTS_JSON, \"w\") as f:\n    json.dump(results, f, indent=2)\n\nif best_ckpt_overall:\n    stage_best = f\"{SAVE_DIR}/stage{SHARD_ID}_uni_best_qwk{best_overall:.4f}.pth\"\n    shutil.copy(best_ckpt_overall, stage_best)\n    print(f\"\\nOK Best checkpoint: {stage_best}\")\n    print(f\"\\nComparaison finale:\")\n    print(f\"  ResNet-50 Stage 3 : mean=0.7280, best=0.7845\")\n    print(f\"  UNI Stage 3       : mean={mean_kappa:.4f}, best={best_overall:.4f}\")\n    print(f\"  Gain mean         : {mean_kappa - 0.7280:+.4f}\")\n    print(f\"  Gain best         : {best_overall - 0.7845:+.4f}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\n=============================================================\nStage 3 — HybridUNITransformer — FROM SCRATCH\nUNI (MahmoodLab) foundation model + Transformer Encoder\nDonnées cumulatives :\n  Shard 0 : patches_400        (399  WSIs,  4,319 patches)\n  Shard 1 : panda-patches-run1 (1000 WSIs, 14,842 patches)\n  Shard 2 : panda-patches-run2 (999  WSIs, 14,867 patches)\n  Shard 3 : panda-patches-run3 (999  WSIs, 14,835 patches)\nTotal     : ~3,398 WSIs, ~48,863 patches\n\nComparaison :\n  ResNet-50 Stage 5 : mean=0.7215, best=0.7415\n  UNI Stage 3       : ???\n=============================================================\nAttache dans Kaggle :\n  khouloudbejaoui20/dataset\n  khouloudbejaoui20/panda-patches-run1\n  khouloudbejaoui20/panda-patches-run2\n  khouloudbejaoui20/panda-patches-run3\nGPU : T4×2\n\"\"\"\n\nimport os, glob, json, random, shutil, subprocess, datetime\nimport numpy as np\nfrom PIL import Image\nfrom collections import Counter\nfrom sklearn.model_selection import StratifiedGroupKFold\nfrom sklearn.metrics import cohen_kappa_score\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as T\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\n\n# ─────────────────────────────────────────────\n# 1. LOGIN HUGGINGFACE\n# ─────────────────────────────────────────────\nfrom huggingface_hub import login, whoami\nimport timm\n\nHF_TOKEN = \"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\"\nlogin(token=HF_TOKEN)\ninfo = whoami()\nprint(f\"OK HuggingFace: {info['name']}\")\n\n# ─────────────────────────────────────────────\n# 2. CONFIG\n# ─────────────────────────────────────────────\nSHARD_ID    = 5\nSTART_FOLD  = 0\nDONE_KAPPAS = []\n\n# UNI Stage 3 = FROM SCRATCH (pas de PREV_CKPT)\nPREV_CKPT   = None\n\n# FREEZE UNI backbone pour éviter overfitting !\nFREEZE_UNI  = True   # ← clé pour réduire overfitting\n\nN_FOLDS     = 5\nSEED        = 42\nBATCH_SIZE  = 8      # réduit pour éviter OOM avec UNI\nNUM_WORKERS = 2\nEPOCHS      = 20\nPATIENCE    = 5\nLR_UNI      = 1e-5   # utilisé seulement si FREEZE_UNI=False\nLR_TR       = 4e-5   # Transformer + MLP head\n\nrandom.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)\ntorch.backends.cudnn.deterministic = True\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Device     : {DEVICE}\")\nprint(f\"Stage      : {SHARD_ID} UNI — FROM SCRATCH\")\nprint(f\"Freeze UNI : {FREEZE_UNI}\")\nprint(f\"BATCH={BATCH_SIZE} | EPOCHS={EPOCHS} | PATIENCE={PATIENCE}\")\nprint(f\"LR_TR={LR_TR}\")\n\n# ─────────────────────────────────────────────\n# 3. CHEMINS PATCHES\n# ─────────────────────────────────────────────\nBASE = \"/kaggle/input/datasets/khouloudbejaoui20\"\n\nRUN5_BASE = \"/kaggle/input/datasets/bejaouikhouloud/panda-patches-run5-parts\"\n\nSHARD_PATHS = {\n    0: f\"{BASE}/dataset/patches_400-20260630T212405Z-3-001/patches_400\",\n    1: f\"{BASE}/panda-patches-run1/kaggle/working/patches_panda_new\",\n    2: f\"{BASE}/panda-patches-run2/kaggle/working/patches_panda_run2\",\n    3: f\"{BASE}/panda-patches-run3/kaggle/working/patches_panda_run3_correct\",\n    4: f\"{BASE}/panda-patches-run4\",\n}\n\nSAVE_DIR     = \"/kaggle/working\"\nKAGGLE_DS    = \"bejaouikhouloud/stage5-uni-results\"\nRESULTS_JSON = f\"{SAVE_DIR}/stage5_uni_cv_results.json\"\n\n# ─────────────────────────────────────────────\n# 4. VERIFICATION SHARDS\n# ─────────────────────────────────────────────\nprint(f\"\\n── Verification ────────────────────────\")\nfor shard_id, path in SHARD_PATHS.items():\n    if os.path.exists(path):\n        files = glob.glob(f\"{path}/*.png\")\n        if not files:\n            files = glob.glob(f\"{path}/**/*.png\", recursive=True)\n        print(f\"  OK Shard {shard_id}: {len(files)} patches\")\n    else:\n        print(f\"  MANQUANT Shard {shard_id}: {path}\")\n\n# ─────────────────────────────────────────────\n# 5. CHARGEMENT PATCHES\n# ─────────────────────────────────────────────\nprint(f\"\\n── Loading patches ──────────────────────\")\n\ndef load_patches(folder):\n    files = glob.glob(f\"{folder}/*.png\")\n    if not files:\n        files = glob.glob(f\"{folder}/**/*.png\", recursive=True)\n    records = []\n    for fp in files:\n        fname = os.path.basename(fp)\n        try:\n            parts  = fname.split(\"_slide\")\n            label  = int(parts[0].replace(\"grade\", \"\"))\n            wsi_id = parts[1].split(\"_\")[0]\n            records.append({\"path\": fp, \"wsi_id\": wsi_id, \"label\": label})\n        except:\n            continue\n    return records\n\nall_records = []\nfor shard_id, path in SHARD_PATHS.items():\n    if path and os.path.exists(path):\n        r = load_patches(path)\n        all_records.extend(r)\n        print(f\"  Shard {shard_id}: {len(r)} patches\")\n    else:\n        print(f\"  Shard {shard_id}: NON TROUVE\")\n\n# Shard 5 — 15 parties\nrun5_total = 0\nfor part in range(1, 16):\n    path = f\"{RUN5_BASE}/run5_part{part}/kaggle/working/patches_panda_run5\"\n    if os.path.exists(path):\n        r = load_patches(path)\n        all_records.extend(r)\n        run5_total += len(r)\nprint(f\"  Shard 5 (15 parts): {run5_total} patches\")\n\n# Deduplicate\nseen = set()\nunique_records = []\nfor r in all_records:\n    fname = os.path.basename(r[\"path\"])\n    if fname not in seen:\n        seen.add(fname)\n        unique_records.append(r)\n\nall_wsis = list(set(r[\"wsi_id\"] for r in unique_records))\nprint(f\"\\n  Total unique : {len(unique_records)} patches\")\nprint(f\"  Total WSIs  : {len(all_wsis)}\")\ncnt = Counter(r[\"label\"] for r in unique_records)\nfor c in range(6):\n    print(f\"    ISUP {c}: {cnt.get(c, 0)}\")\n\nassert len(all_wsis) > 4500, f\"Trop peu de WSIs ({len(all_wsis)})\"\nprint(f\"\\n  OK {len(all_wsis)} WSIs confirmes\")\n\n# ─────────────────────────────────────────────\n# 6. TRANSFORMS\n# ─────────────────────────────────────────────\nMEAN = [0.485, 0.456, 0.406]\nSTD  = [0.229, 0.224, 0.225]\n\ntrain_tf = T.Compose([\n    T.Resize((224, 224)),\n    T.RandomHorizontalFlip(),\n    T.RandomVerticalFlip(),\n    T.RandomRotation(90),\n    T.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.2, hue=0.1),\n    T.RandomApply([T.GaussianBlur(5)], p=0.3),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n    T.RandomApply([T.RandomErasing(p=0.5, scale=(0.02, 0.2))], p=0.2),\n])\n\nval_tf = T.Compose([\n    T.Resize((224, 224)),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n])\n\nclass PatchDataset(Dataset):\n    def __init__(self, records, transform):\n        self.records   = records\n        self.transform = transform\n    def __len__(self): return len(self.records)\n    def __getitem__(self, idx):\n        r = self.records[idx]\n        try:\n            img = Image.open(r[\"path\"]).convert(\"RGB\")\n        except:\n            img = Image.new(\"RGB\", (224, 224), (255, 255, 255))\n        return self.transform(img), int(r[\"label\"])\n\n# ─────────────────────────────────────────────\n# 7. MODELE — HybridUNITransformer\n# ─────────────────────────────────────────────\nclass HybridUNITransformer(nn.Module):\n    def __init__(self, num_classes=6, d_model=512, nhead=8,\n                 num_layers=2, dim_feedforward=2048, dropout=0.1,\n                 freeze_uni=True):\n        super().__init__()\n\n        # UNI backbone\n        self.uni_backbone = timm.create_model(\n            \"hf-hub:MahmoodLab/UNI\",\n            pretrained=True,\n            num_classes=0,\n            init_values=1e-5,\n            dynamic_img_size=True\n        )\n        uni_dim = self.uni_backbone.num_features  # 1024\n\n        # Freeze UNI si demandé\n        if freeze_uni:\n            for param in self.uni_backbone.parameters():\n                param.requires_grad = False\n            print(f\"  UNI backbone FREEZÉ\")\n        else:\n            print(f\"  UNI backbone trainable\")\n\n        # Projection\n        self.proj = nn.Sequential(\n            nn.Linear(uni_dim, d_model),\n            nn.LayerNorm(d_model)\n        )\n\n        # CLS token + positional embedding\n        self.cls_token = nn.Parameter(torch.zeros(1, 1, d_model))\n        self.pos_embed = nn.Parameter(torch.zeros(1, 2, d_model))\n\n        # Transformer Encoder\n        encoder_layer = nn.TransformerEncoderLayer(\n            d_model=d_model, nhead=nhead,\n            dim_feedforward=dim_feedforward,\n            dropout=dropout, batch_first=True\n        )\n        self.transformer = nn.TransformerEncoder(\n            encoder_layer, num_layers=num_layers,\n            norm=nn.LayerNorm(d_model)\n        )\n\n        # MLP Head\n        self.classifier = nn.Sequential(\n            nn.Linear(d_model, 256),\n            nn.ReLU(),\n            nn.Dropout(dropout),\n            nn.Linear(256, num_classes)\n        )\n\n        nn.init.trunc_normal_(self.cls_token, std=0.02)\n        nn.init.trunc_normal_(self.pos_embed, std=0.02)\n\n    def forward(self, x):\n        with torch.no_grad() if not self.uni_backbone.training else torch.enable_grad():\n            f = self.uni_backbone(x)\n        f   = self.proj(f).unsqueeze(1)\n        cls = self.cls_token.expand(x.size(0), -1, -1)\n        seq = torch.cat([cls, f], dim=1) + self.pos_embed\n        return self.classifier(self.transformer(seq)[:, 0])\n\ndef build_model():\n    # Libère mémoire GPU\n    torch.cuda.empty_cache()\n    import gc; gc.collect()\n\n    model  = HybridUNITransformer(freeze_uni=FREEZE_UNI).to(DEVICE)\n    n_gpus = torch.cuda.device_count()\n    if n_gpus > 1:\n        model = nn.DataParallel(model)\n        print(f\"  OK DataParallel: {n_gpus} GPUs\")\n\n    # Paramètres trainables\n    n_total    = sum(p.numel() for p in model.parameters())\n    n_trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    print(f\"  Total params    : {n_total:,}\")\n    print(f\"  Trainable params: {n_trainable:,}\")\n    return model\n\n# ─────────────────────────────────────────────\n# 8. WSI-LEVEL QWK\n# ─────────────────────────────────────────────\ndef wsi_level_qwk(model, val_data, device):\n    model.eval()\n    wsi_preds  = {}\n    wsi_labels = {}\n    all_preds  = []\n    all_labels = []\n    loader = DataLoader(\n        PatchDataset(val_data, val_tf),\n        batch_size=BATCH_SIZE, shuffle=False,\n        num_workers=NUM_WORKERS\n    )\n    with torch.no_grad():\n        for imgs, labels in loader:\n            imgs  = imgs.to(device)\n            preds = model(imgs).argmax(1).cpu().numpy()\n            all_preds.extend(preds)\n            all_labels.extend(labels.numpy())\n    for i, r in enumerate(val_data):\n        wid = r[\"wsi_id\"]\n        if wid not in wsi_preds:\n            wsi_preds[wid]  = []\n            wsi_labels[wid] = r[\"label\"]\n        wsi_preds[wid].append(all_preds[i])\n    wf = []; wl = []\n    for wid in wsi_preds:\n        votes = wsi_preds[wid]\n        wf.append(max(set(votes), key=votes.count))\n        wl.append(wsi_labels[wid])\n    wsi_qwk = cohen_kappa_score(wl, wf, weights=\"quadratic\")\n    return wsi_qwk, len(wf)\n\n# ─────────────────────────────────────────────\n# 9. PUSH KAGGLE\n# ─────────────────────────────────────────────\ndef push_to_kaggle(ckpt_path, message):\n    try:\n        tmp = f\"{SAVE_DIR}/_push_uni3\"\n        os.makedirs(tmp, exist_ok=True)\n        for c in glob.glob(f\"{SAVE_DIR}/stage5_uni*.pth\"):\n            shutil.copy(c, tmp)\n        shutil.copy(ckpt_path, tmp)\n        if os.path.exists(RESULTS_JSON):\n            shutil.copy(RESULTS_JSON, tmp)\n        json.dump({\n            \"title\": \"Stage3 UNI Results Bejaoui 2026\",\n            \"id\": KAGGLE_DS,\n            \"licenses\": [{\"name\": \"CC0-1.0\"}]\n        }, open(f\"{tmp}/dataset-metadata.json\", \"w\"))\n        r = subprocess.run(\n            [\"kaggle\", \"datasets\", \"create\",\n             \"-p\", tmp, \"--dir-mode\", \"skip\"],\n            capture_output=True, text=True, timeout=600\n        )\n        if r.returncode != 0:\n            r = subprocess.run(\n                [\"kaggle\", \"datasets\", \"version\",\n                 \"-p\", tmp, \"-m\", message, \"--dir-mode\", \"skip\"],\n                capture_output=True, text=True, timeout=600\n            )\n        print(f\"   {'OK!' if r.returncode==0 else 'ERREUR: '+r.stderr[:100]}\")\n    except Exception as e:\n        print(f\"   ERREUR: {e}\")\n\n# ─────────────────────────────────────────────\n# 10. TRAIN ONE FOLD\n# ─────────────────────────────────────────────\ndef train_one_fold(fold_id, train_data, val_data, n_folds=5):\n    print(f\"\\n  Fold {fold_id+1}/{n_folds} ──────────────────────────\")\n    print(f\"     Train: {len(train_data)} | Val: {len(val_data)}\")\n    print(f\"     Freeze UNI: {FREEZE_UNI}\")\n\n    train_loader = DataLoader(\n        PatchDataset(train_data, train_tf),\n        batch_size=BATCH_SIZE, shuffle=True,\n        num_workers=NUM_WORKERS, pin_memory=True\n    )\n\n    model = build_model()\n\n    # Class weights\n    cnt_tr = Counter(r[\"label\"] for r in train_data)\n    w = torch.tensor(\n        [1.0 / max(cnt_tr.get(c, 1), 1) for c in range(6)],\n        dtype=torch.float\n    )\n    w = (w / w.sum() * 6).to(DEVICE)\n    criterion = nn.CrossEntropyLoss(weight=w, label_smoothing=0.1)\n\n    # Optimizer — seulement les paramètres trainables\n    base_model = model.module if hasattr(model, \"module\") else model\n\n    if FREEZE_UNI:\n        # Entraîne seulement proj + transformer + classifier\n        trainable_params = [p for p in model.parameters() if p.requires_grad]\n        optimizer = optim.AdamW(trainable_params, lr=LR_TR, weight_decay=1e-4)\n        print(f\"     Optimizer: AdamW lr={LR_TR} (trainable only)\")\n    else:\n        optimizer = optim.AdamW([\n            {\"params\": base_model.uni_backbone.parameters(), \"lr\": LR_UNI},\n            {\"params\": [p for n, p in base_model.named_parameters()\n                        if not n.startswith(\"uni_backbone\")], \"lr\": LR_TR},\n        ], weight_decay=1e-4)\n        print(f\"     Optimizer: AdamW LR_UNI={LR_UNI} LR_TR={LR_TR}\")\n\n    scheduler = CosineAnnealingLR(optimizer, T_max=EPOCHS, eta_min=1e-7)\n\n    best_kappa = 0.0\n    no_improve = 0\n    best_ckpt  = f\"{SAVE_DIR}/stage{SHARD_ID}_uni_fold{fold_id+1}_best.pth\"\n\n    for epoch in range(1, EPOCHS + 1):\n        model.train()\n        if FREEZE_UNI:\n            # S'assure que UNI reste en eval mode\n            base_model = model.module if hasattr(model, \"module\") else model\n            base_model.uni_backbone.eval()\n\n        t_loss, t_correct, t_total = 0.0, 0, 0\n        for imgs, labels in train_loader:\n            imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n            optimizer.zero_grad()\n            out  = model(imgs)\n            loss = criterion(out, labels)\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n            optimizer.step()\n            t_loss    += loss.item() * imgs.size(0)\n            t_correct += (out.argmax(1) == labels).sum().item()\n            t_total   += imgs.size(0)\n        scheduler.step()\n\n        wsi_kappa, n_wsis_val = wsi_level_qwk(model, val_data, DEVICE)\n        is_best = wsi_kappa > best_kappa\n\n        print(f\"     Ep[{epoch:02d}/{EPOCHS}] \"\n              f\"loss={t_loss/t_total:.4f} acc={t_correct/t_total:.3f} | \"\n              f\"WSI_QWK={wsi_kappa:.4f} ({n_wsis_val} WSIs)\"\n              + (\" *\" if is_best else \"\"))\n\n        if is_best:\n            best_kappa = wsi_kappa\n            no_improve = 0\n            torch.save({\n                \"epoch\":            epoch,\n                \"model_state_dict\": model.state_dict() if not hasattr(model, \"module\")\n                                    else model.module.state_dict(),\n                \"kappa\":            wsi_kappa,\n                \"shard_id\":         SHARD_ID,\n                \"fold_id\":          fold_id,\n                \"backbone\":         \"UNI\",\n                \"freeze_uni\":       FREEZE_UNI,\n                \"n_wsis\":           len(all_wsis),\n            }, best_ckpt)\n        else:\n            no_improve += 1\n            if no_improve >= PATIENCE:\n                print(f\"     Early stopping epoch {epoch}\")\n                break\n\n    print(f\"  OK Fold {fold_id+1} best WSI_QWK: {best_kappa:.4f}\")\n    push_to_kaggle(best_ckpt,\n        f\"Stage{SHARD_ID} UNI Fold{fold_id+1} QWK={best_kappa:.4f}\")\n    return best_kappa, best_ckpt\n\n# ─────────────────────────────────────────────\n# 11. 5-FOLD CV\n# ─────────────────────────────────────────────\nprint(f\"\\n{'='*60}\")\nprint(f\"Stage {SHARD_ID} UNI — {len(all_wsis)} WSIs — 5-Fold CV\")\nprint(f\"Backbone: UNI (MahmoodLab) | Freeze: {FREEZE_UNI}\")\nprint(f\"ResNet-50 Stage 5 baseline: mean=0.7215, best=0.7415\")\nprint(f\"{'='*60}\")\n\nwsi_ids    = [r[\"wsi_id\"] for r in unique_records]\nlabels_arr = [r[\"label\"]  for r in unique_records]\nindices    = list(range(len(unique_records)))\n\nsgkf = StratifiedGroupKFold(n_splits=N_FOLDS, shuffle=True, random_state=SEED)\n\nfold_kappas       = DONE_KAPPAS.copy()\nbest_overall      = max(DONE_KAPPAS) if DONE_KAPPAS else -1.0\nbest_ckpt_overall = None\n\nfor fold_id, (train_idx, val_idx) in enumerate(\n        sgkf.split(indices, labels_arr, groups=wsi_ids)):\n\n    if fold_id < START_FOLD:\n        continue\n\n    train_data = [unique_records[i] for i in train_idx]\n    val_data   = [unique_records[i] for i in val_idx]\n\n    tw = set(unique_records[i][\"wsi_id\"] for i in train_idx)\n    vw = set(unique_records[i][\"wsi_id\"] for i in val_idx)\n    assert len(tw & vw) == 0, f\"DATA LEAKAGE fold {fold_id+1}!\"\n\n    kappa, ckpt_path = train_one_fold(fold_id, train_data, val_data, N_FOLDS)\n    fold_kappas.append(kappa)\n\n    if kappa > best_overall:\n        best_overall      = kappa\n        best_ckpt_overall = ckpt_path\n\n# ─────────────────────────────────────────────\n# 12. RESUME FINAL\n# ─────────────────────────────────────────────\nmean_kappa = np.mean(fold_kappas)\nstd_kappa  = np.std(fold_kappas)\n\nprint(f\"\\n{'='*60}\")\nprint(f\"OK Stage {SHARD_ID} UNI — 5-Fold CV Results\")\nprint(f\"{'='*60}\")\nfor i, k in enumerate(fold_kappas):\n    print(f\"   Fold {i+1}: QWK = {k:.4f}\")\nprint(f\"{'─'*40}\")\nprint(f\"   Mean QWK UNI        : {mean_kappa:.4f} +/- {std_kappa:.4f}\")\nprint(f\"   Best QWK UNI        : {best_overall:.4f}\")\nprint(f\"   ResNet-50 Stage 3   : mean=0.7280, best=0.7845\")\nprint(f\"   Gain vs ResNet-50   : {mean_kappa - 0.7215:+.4f}\")\nprint(f\"   WSIs                : {len(all_wsis)}\")\nprint(f\"   SOTA [1]            : 0.934\")\nprint(f\"{'='*60}\")\n\nresults = {\n    \"stage\":           SHARD_ID,\n    \"backbone\":        \"UNI\",\n    \"freeze_uni\":      FREEZE_UNI,\n    \"n_wsis\":          len(all_wsis),\n    \"n_patches\":       len(unique_records),\n    \"fold_kappas\":     fold_kappas,\n    \"mean_kappa\":      mean_kappa,\n    \"std_kappa\":       std_kappa,\n    \"best_kappa\":      best_overall,\n    \"resnet50_mean\":   0.7215,\n    \"resnet50_best\":   0.7415,\n    \"gain_vs_resnet\":  mean_kappa - 0.7215,\n    \"timestamp\":       datetime.datetime.now().isoformat(),\n}\nwith open(RESULTS_JSON, \"w\") as f:\n    json.dump(results, f, indent=2)\n\nif best_ckpt_overall:\n    stage_best = f\"{SAVE_DIR}/stage{SHARD_ID}_uni_best_qwk{best_overall:.4f}.pth\"\n    shutil.copy(best_ckpt_overall, stage_best)\n    print(f\"\\nOK Best checkpoint: {stage_best}\")\n    print(f\"\\nComparaison finale:\")\n    print(f\"  ResNet-50 Stage 5 : mean=0.7215, best=0.7415\")\n    print(f\"  UNI Stage 3       : mean={mean_kappa:.4f}, best={best_overall:.4f}\")\n    print(f\"  Gain mean         : {mean_kappa - 0.7215:+.4f}\")\n    print(f\"  Gain best         : {best_overall - 0.7415:+.4f}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\n=============================================================\nStage 0 — HybridFusion (ResNet-101 + UNI) — Parallel Fusion\nArchitecture proposée par Dr. Mouelhi :\n  Branch 1 : ResNet-101 → features locales (CNN)\n  Branch 2 : UNI → features globales (Transformer)\n  Fusion    : 0.6 × CNN + 0.4 × UNI → MLP → ISUP\n\nDataset: patches_400 (399 WSIs, 4,319 patches)\nComparaison :\n  ResNet-50 Stage 0  : mean=0.5764, best=0.6116\n  UNI Stage 0        : mean=0.6705, best=0.7964\n  ResNet+UNI Stage 0 : ???\n=============================================================\nAttache dans Kaggle :\n  khouloudbejaoui20/dataset\nGPU : T4x2\n\"\"\"\n\nimport os, glob, json, random, shutil, subprocess, datetime\nimport numpy as np\nfrom PIL import Image\nfrom collections import Counter\nfrom sklearn.model_selection import StratifiedGroupKFold\nfrom sklearn.metrics import cohen_kappa_score\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as T\nimport torchvision.models as models\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\n\n# ─────────────────────────────────────────────\n# 1. LOGIN HUGGINGFACE\n# ─────────────────────────────────────────────\nfrom huggingface_hub import login, whoami\nimport timm\n\nHF_TOKEN = \"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\"\nlogin(token=HF_TOKEN)\ninfo = whoami()\nprint(f\"OK HuggingFace: {info['name']}\")\n\n# ─────────────────────────────────────────────\n# 2. CONFIG\n# ─────────────────────────────────────────────\nSHARD_ID    = 0\nSTART_FOLD  = 0\nDONE_KAPPAS = []\n\nCNN_BACKBONE = \"resnet101\"  # resnet50 or resnet101\nFREEZE_UNI   = True         # UNI freezé pour éviter overfitting\nALPHA        = 0.6          # poids CNN\nBETA         = 0.4          # poids UNI\n\nN_FOLDS     = 5\nSEED        = 42\nBATCH_SIZE  = 8    # réduit car 2 modèles en parallèle\nNUM_WORKERS = 2\nEPOCHS      = 20\nPATIENCE    = 5\nLR_CNN      = 5e-5\nLR_TR       = 4e-5  # pour les couches non-UNI\n\nrandom.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)\ntorch.backends.cudnn.deterministic = True\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Device      : {DEVICE}\")\nprint(f\"Architecture: HybridFusion ({CNN_BACKBONE} + UNI)\")\nprint(f\"Fusion      : {ALPHA} x CNN + {BETA} x UNI\")\nprint(f\"Freeze UNI  : {FREEZE_UNI}\")\nprint(f\"BATCH={BATCH_SIZE} | EPOCHS={EPOCHS} | PATIENCE={PATIENCE}\")\n\n# ─────────────────────────────────────────────\n# 3. CHEMINS\n# ─────────────────────────────────────────────\nBASE = \"/kaggle/input/datasets/khouloudbejaoui20\"\nSHARD_PATH = f\"{BASE}/dataset/patches_400-20260630T212405Z-3-001/patches_400\"\n\nSAVE_DIR     = \"/kaggle/working\"\nKAGGLE_DS    = \"bejaouikhouloud/stage0-fusion-results\"\nRESULTS_JSON = f\"{SAVE_DIR}/stage0_fusion_cv_results.json\"\n\n# ─────────────────────────────────────────────\n# 4. LOAD PATCHES\n# ─────────────────────────────────────────────\nprint(f\"\\n── Loading patches ──────────────────────\")\n\ndef load_patches(folder):\n    files = glob.glob(f\"{folder}/*.png\")\n    if not files:\n        files = glob.glob(f\"{folder}/**/*.png\", recursive=True)\n    records = []\n    for fp in files:\n        fname = os.path.basename(fp)\n        try:\n            parts  = fname.split(\"_slide\")\n            label  = int(parts[0].replace(\"grade\", \"\"))\n            wsi_id = parts[1].split(\"_\")[0]\n            records.append({\"path\": fp, \"wsi_id\": wsi_id, \"label\": label})\n        except:\n            continue\n    return records\n\nall_records = load_patches(SHARD_PATH)\nprint(f\"  Shard 0: {len(all_records)} patches\")\n\nall_wsis = list(set(r[\"wsi_id\"] for r in all_records))\nprint(f\"  Total WSIs: {len(all_wsis)}\")\ncnt = Counter(r[\"label\"] for r in all_records)\nfor c in range(6):\n    print(f\"    ISUP {c}: {cnt.get(c, 0)}\")\n\n# ─────────────────────────────────────────────\n# 5. TRANSFORMS\n# ─────────────────────────────────────────────\nMEAN = [0.485, 0.456, 0.406]\nSTD  = [0.229, 0.224, 0.225]\n\ntrain_tf = T.Compose([\n    T.Resize((224, 224)),\n    T.RandomHorizontalFlip(),\n    T.RandomVerticalFlip(),\n    T.RandomRotation(90),\n    T.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.2, hue=0.1),\n    T.RandomApply([T.GaussianBlur(5)], p=0.3),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n    T.RandomApply([T.RandomErasing(p=0.5, scale=(0.02, 0.2))], p=0.2),\n])\n\nval_tf = T.Compose([\n    T.Resize((224, 224)),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n])\n\nclass PatchDataset(Dataset):\n    def __init__(self, records, transform):\n        self.records   = records\n        self.transform = transform\n    def __len__(self): return len(self.records)\n    def __getitem__(self, idx):\n        r = self.records[idx]\n        try:\n            img = Image.open(r[\"path\"]).convert(\"RGB\")\n        except:\n            img = Image.new(\"RGB\", (224, 224), (255, 255, 255))\n        return self.transform(img), int(r[\"label\"])\n\n# ─────────────────────────────────────────────\n# 6. MODELE — HybridFusion\n# ─────────────────────────────────────────────\nclass HybridFusion(nn.Module):\n    \"\"\"\n    Architecture fusion parallèle :\n    Branch 1 : ResNet-101 → features locales (CNN)\n    Branch 2 : UNI (freezé) → features globales\n    Fusion   : alpha × CNN + beta × UNI → MLP → ISUP\n\n    Les deux branches reçoivent la MÊME image en entrée.\n    \"\"\"\n    def __init__(self, cnn_backbone=\"resnet101\",\n                 num_classes=6, d_model=512,\n                 alpha=0.6, beta=0.4,\n                 freeze_uni=True, dropout=0.1):\n        super().__init__()\n\n        self.alpha = alpha\n        self.beta  = beta\n\n        # ── Branch 1 : CNN ──────────────────────\n        if cnn_backbone == \"resnet50\":\n            net = models.resnet50(weights=models.ResNet50_Weights.DEFAULT)\n            cnn_dim = 2048\n        elif cnn_backbone == \"resnet101\":\n            net = models.resnet101(weights=models.ResNet101_Weights.DEFAULT)\n            cnn_dim = 2048\n        elif cnn_backbone == \"resnet152\":\n            net = models.resnet152(weights=models.ResNet152_Weights.DEFAULT)\n            cnn_dim = 2048\n        else:\n            raise ValueError(f\"Unknown backbone: {cnn_backbone}\")\n\n        self.cnn_backbone = nn.Sequential(*list(net.children())[:-2])\n        self.pool         = nn.AdaptiveAvgPool2d((1, 1))\n        self.proj_cnn     = nn.Sequential(\n            nn.Linear(cnn_dim, d_model),\n            nn.LayerNorm(d_model),\n            nn.ReLU(),\n            nn.Dropout(dropout)\n        )\n        print(f\"  CNN Branch : {cnn_backbone} → {cnn_dim}-d → {d_model}-d\")\n\n        # ── Branch 2 : UNI ──────────────────────\n        self.uni_backbone = timm.create_model(\n            \"hf-hub:MahmoodLab/UNI\",\n            pretrained=True,\n            num_classes=0,\n            init_values=1e-5,\n            dynamic_img_size=True\n        )\n        uni_dim = self.uni_backbone.num_features  # 1024\n\n        if freeze_uni:\n            for param in self.uni_backbone.parameters():\n                param.requires_grad = False\n            print(f\"  UNI Branch : UNI FREEZÉ → {uni_dim}-d → {d_model}-d\")\n        else:\n            print(f\"  UNI Branch : UNI trainable → {uni_dim}-d → {d_model}-d\")\n\n        self.proj_uni = nn.Sequential(\n            nn.Linear(uni_dim, d_model),\n            nn.LayerNorm(d_model),\n            nn.ReLU(),\n            nn.Dropout(dropout)\n        )\n\n        # ── Fusion + Classifier ─────────────────\n        print(f\"  Fusion     : {alpha} x CNN + {beta} x UNI\")\n        self.classifier = nn.Sequential(\n            nn.Linear(d_model, 256),\n            nn.ReLU(),\n            nn.Dropout(dropout),\n            nn.Linear(256, num_classes)\n        )\n\n    def forward(self, x):\n        # Branch 1 : CNN features\n        cnn_feat = self.pool(\n            self.cnn_backbone(x)\n        ).flatten(1)                      # (B, 2048)\n        cnn_feat = self.proj_cnn(cnn_feat) # (B, 512)\n\n        # Branch 2 : UNI features\n        if not self.uni_backbone.training:\n            with torch.no_grad():\n                uni_feat = self.uni_backbone(x)  # (B, 1024)\n        else:\n            uni_feat = self.uni_backbone(x)\n        uni_feat = self.proj_uni(uni_feat)  # (B, 512)\n\n        # Fusion pondérée\n        fused = self.alpha * cnn_feat + self.beta * uni_feat  # (B, 512)\n\n        return self.classifier(fused)\n\ndef build_model():\n    torch.cuda.empty_cache()\n    import gc; gc.collect()\n\n    model  = HybridFusion(\n        cnn_backbone=CNN_BACKBONE,\n        alpha=ALPHA,\n        beta=BETA,\n        freeze_uni=FREEZE_UNI\n    ).to(DEVICE)\n\n    n_gpus = torch.cuda.device_count()\n    if n_gpus > 1:\n        model = nn.DataParallel(model)\n        print(f\"  OK DataParallel: {n_gpus} GPUs\")\n\n    n_total     = sum(p.numel() for p in model.parameters())\n    n_trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    print(f\"  Total params    : {n_total:,}\")\n    print(f\"  Trainable params: {n_trainable:,}\")\n    return model\n\n# ─────────────────────────────────────────────\n# 7. WSI-LEVEL QWK\n# ─────────────────────────────────────────────\ndef wsi_level_qwk(model, val_data, device):\n    model.eval()\n    wsi_preds  = {}\n    wsi_labels = {}\n    all_preds  = []\n    all_labels = []\n    loader = DataLoader(\n        PatchDataset(val_data, val_tf),\n        batch_size=BATCH_SIZE, shuffle=False,\n        num_workers=NUM_WORKERS\n    )\n    with torch.no_grad():\n        for imgs, labels in loader:\n            imgs  = imgs.to(device)\n            preds = model(imgs).argmax(1).cpu().numpy()\n            all_preds.extend(preds)\n            all_labels.extend(labels.numpy())\n    for i, r in enumerate(val_data):\n        wid = r[\"wsi_id\"]\n        if wid not in wsi_preds:\n            wsi_preds[wid]  = []\n            wsi_labels[wid] = r[\"label\"]\n        wsi_preds[wid].append(all_preds[i])\n    wf = []; wl = []\n    for wid in wsi_preds:\n        votes = wsi_preds[wid]\n        wf.append(max(set(votes), key=votes.count))\n        wl.append(wsi_labels[wid])\n    wsi_qwk = cohen_kappa_score(wl, wf, weights=\"quadratic\")\n    return wsi_qwk, len(wf)\n\n# ─────────────────────────────────────────────\n# 8. PUSH KAGGLE\n# ─────────────────────────────────────────────\ndef push_to_kaggle(ckpt_path, message):\n    try:\n        tmp = f\"{SAVE_DIR}/_push_fusion\"\n        os.makedirs(tmp, exist_ok=True)\n        for c in glob.glob(f\"{SAVE_DIR}/stage{SHARD_ID}_fusion*.pth\"):\n            shutil.copy(c, tmp)\n        shutil.copy(ckpt_path, tmp)\n        if os.path.exists(RESULTS_JSON):\n            shutil.copy(RESULTS_JSON, tmp)\n        json.dump({\n            \"title\": \"Stage0 Fusion Results Bejaoui 2026\",\n            \"id\": KAGGLE_DS,\n            \"licenses\": [{\"name\": \"CC0-1.0\"}]\n        }, open(f\"{tmp}/dataset-metadata.json\", \"w\"))\n        r = subprocess.run(\n            [\"kaggle\", \"datasets\", \"create\",\n             \"-p\", tmp, \"--dir-mode\", \"skip\"],\n            capture_output=True, text=True, timeout=600\n        )\n        if r.returncode != 0:\n            r = subprocess.run(\n                [\"kaggle\", \"datasets\", \"version\",\n                 \"-p\", tmp, \"-m\", message, \"--dir-mode\", \"skip\"],\n                capture_output=True, text=True, timeout=600\n            )\n        print(f\"   {'OK!' if r.returncode==0 else 'ERREUR: '+r.stderr[:100]}\")\n    except Exception as e:\n        print(f\"   ERREUR: {e}\")\n\n# ─────────────────────────────────────────────\n# 9. TRAIN ONE FOLD\n# ─────────────────────────────────────────────\ndef train_one_fold(fold_id, train_data, val_data, n_folds=5):\n    print(f\"\\n  Fold {fold_id+1}/{n_folds} ──────────────────────────\")\n    print(f\"     Train: {len(train_data)} | Val: {len(val_data)}\")\n\n    train_loader = DataLoader(\n        PatchDataset(train_data, train_tf),\n        batch_size=BATCH_SIZE, shuffle=True,\n        num_workers=NUM_WORKERS, pin_memory=True\n    )\n\n    model = build_model()\n\n    # Class weights\n    cnt_tr = Counter(r[\"label\"] for r in train_data)\n    w = torch.tensor(\n        [1.0 / max(cnt_tr.get(c, 1), 1) for c in range(6)],\n        dtype=torch.float\n    )\n    w = (w / w.sum() * 6).to(DEVICE)\n    criterion = nn.CrossEntropyLoss(weight=w, label_smoothing=0.1)\n\n    # Optimizer\n    base_model = model.module if hasattr(model, \"module\") else model\n    trainable  = [p for p in model.parameters() if p.requires_grad]\n    optimizer  = optim.AdamW([\n        {\"params\": base_model.cnn_backbone.parameters(), \"lr\": LR_CNN},\n        {\"params\": [p for n, p in base_model.named_parameters()\n                    if not n.startswith(\"cnn_backbone\") and p.requires_grad],\n         \"lr\": LR_TR},\n    ], weight_decay=1e-4)\n\n    scheduler = CosineAnnealingLR(optimizer, T_max=EPOCHS, eta_min=1e-7)\n\n    best_kappa = 0.0\n    no_improve = 0\n    best_ckpt  = f\"{SAVE_DIR}/stage{SHARD_ID}_fusion_fold{fold_id+1}_best.pth\"\n\n    for epoch in range(1, EPOCHS + 1):\n        model.train()\n        # UNI reste en eval mode si freezé\n        if FREEZE_UNI:\n            base = model.module if hasattr(model, \"module\") else model\n            base.uni_backbone.eval()\n\n        t_loss, t_correct, t_total = 0.0, 0, 0\n        for imgs, labels in train_loader:\n            imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n            optimizer.zero_grad()\n            out  = model(imgs)\n            loss = criterion(out, labels)\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n            optimizer.step()\n            t_loss    += loss.item() * imgs.size(0)\n            t_correct += (out.argmax(1) == labels).sum().item()\n            t_total   += imgs.size(0)\n        scheduler.step()\n\n        wsi_kappa, n_wsis_val = wsi_level_qwk(model, val_data, DEVICE)\n        is_best = wsi_kappa > best_kappa\n\n        print(f\"     Ep[{epoch:02d}/{EPOCHS}] \"\n              f\"loss={t_loss/t_total:.4f} acc={t_correct/t_total:.3f} | \"\n              f\"WSI_QWK={wsi_kappa:.4f} ({n_wsis_val} WSIs)\"\n              + (\" *\" if is_best else \"\"))\n\n        if is_best:\n            best_kappa = wsi_kappa\n            no_improve = 0\n            torch.save({\n                \"epoch\":            epoch,\n                \"model_state_dict\": model.state_dict() if not hasattr(model, \"module\")\n                                    else model.module.state_dict(),\n                \"kappa\":            wsi_kappa,\n                \"cnn_backbone\":     CNN_BACKBONE,\n                \"alpha\":            ALPHA,\n                \"beta\":             BETA,\n                \"freeze_uni\":       FREEZE_UNI,\n                \"shard_id\":         SHARD_ID,\n                \"fold_id\":          fold_id,\n                \"n_wsis\":           len(all_wsis),\n            }, best_ckpt)\n        else:\n            no_improve += 1\n            if no_improve >= PATIENCE:\n                print(f\"     Early stopping epoch {epoch}\")\n                break\n\n    print(f\"  OK Fold {fold_id+1} best WSI_QWK: {best_kappa:.4f}\")\n    push_to_kaggle(best_ckpt,\n        f\"Stage{SHARD_ID} Fusion Fold{fold_id+1} QWK={best_kappa:.4f}\")\n    return best_kappa, best_ckpt\n\n# ─────────────────────────────────────────────\n# 10. 5-FOLD CV\n# ─────────────────────────────────────────────\nprint(f\"\\n{'='*60}\")\nprint(f\"Stage {SHARD_ID} HybridFusion — {len(all_wsis)} WSIs — 5-Fold CV\")\nprint(f\"Architecture : {CNN_BACKBONE} + UNI (freeze={FREEZE_UNI})\")\nprint(f\"Fusion       : {ALPHA} x CNN + {BETA} x UNI\")\nprint(f\"Baseline     : ResNet-50=0.5764 | UNI=0.6705\")\nprint(f\"{'='*60}\")\n\nwsi_ids    = [r[\"wsi_id\"] for r in all_records]\nlabels_arr = [r[\"label\"]  for r in all_records]\nindices    = list(range(len(all_records)))\n\nsgkf = StratifiedGroupKFold(n_splits=N_FOLDS, shuffle=True, random_state=SEED)\n\nfold_kappas       = DONE_KAPPAS.copy()\nbest_overall      = max(DONE_KAPPAS) if DONE_KAPPAS else -1.0\nbest_ckpt_overall = None\n\nfor fold_id, (train_idx, val_idx) in enumerate(\n        sgkf.split(indices, labels_arr, groups=wsi_ids)):\n\n    if fold_id < START_FOLD:\n        continue\n\n    train_data = [all_records[i] for i in train_idx]\n    val_data   = [all_records[i] for i in val_idx]\n\n    tw = set(all_records[i][\"wsi_id\"] for i in train_idx)\n    vw = set(all_records[i][\"wsi_id\"] for i in val_idx)\n    assert len(tw & vw) == 0, f\"DATA LEAKAGE fold {fold_id+1}!\"\n\n    kappa, ckpt_path = train_one_fold(fold_id, train_data, val_data, N_FOLDS)\n    fold_kappas.append(kappa)\n\n    if kappa > best_overall:\n        best_overall      = kappa\n        best_ckpt_overall = ckpt_path\n\n# ─────────────────────────────────────────────\n# 11. RESUME FINAL\n# ─────────────────────────────────────────────\nmean_kappa = np.mean(fold_kappas)\nstd_kappa  = np.std(fold_kappas)\n\nprint(f\"\\n{'='*60}\")\nprint(f\"OK Stage {SHARD_ID} HybridFusion — 5-Fold CV Results\")\nprint(f\"{'='*60}\")\nfor i, k in enumerate(fold_kappas):\n    print(f\"   Fold {i+1}: QWK = {k:.4f}\")\nprint(f\"{'─'*40}\")\nprint(f\"   Mean QWK Fusion   : {mean_kappa:.4f} +/- {std_kappa:.4f}\")\nprint(f\"   Best QWK Fusion   : {best_overall:.4f}\")\nprint(f\"   ResNet-50 Stage 0 : mean=0.5764, best=0.6116\")\nprint(f\"   UNI Stage 0       : mean=0.6705, best=0.7964\")\nprint(f\"   Gain vs ResNet-50 : {mean_kappa - 0.5764:+.4f}\")\nprint(f\"   Gain vs UNI       : {mean_kappa - 0.6705:+.4f}\")\nprint(f\"   WSIs              : {len(all_wsis)}\")\nprint(f\"   SOTA [1]          : 0.934\")\nprint(f\"{'='*60}\")\n\nresults = {\n    \"stage\":            SHARD_ID,\n    \"architecture\":     f\"HybridFusion_{CNN_BACKBONE}_UNI\",\n    \"cnn_backbone\":     CNN_BACKBONE,\n    \"freeze_uni\":       FREEZE_UNI,\n    \"alpha\":            ALPHA,\n    \"beta\":             BETA,\n    \"n_wsis\":           len(all_wsis),\n    \"n_patches\":        len(all_records),\n    \"fold_kappas\":      fold_kappas,\n    \"mean_kappa\":       mean_kappa,\n    \"std_kappa\":        std_kappa,\n    \"best_kappa\":       best_overall,\n    \"resnet50_mean\":    0.5764,\n    \"uni_mean\":         0.6705,\n    \"gain_vs_resnet50\": mean_kappa - 0.5764,\n    \"gain_vs_uni\":      mean_kappa - 0.6705,\n    \"timestamp\":        datetime.datetime.now().isoformat(),\n}\nwith open(RESULTS_JSON, \"w\") as f:\n    json.dump(results, f, indent=2)\n\nif best_ckpt_overall:\n    stage_best = f\"{SAVE_DIR}/stage{SHARD_ID}_fusion_best_qwk{best_overall:.4f}.pth\"\n    shutil.copy(best_ckpt_overall, stage_best)\n    print(f\"\\nOK Best checkpoint: {stage_best}\")\n    print(f\"\\nComparaison finale:\")\n    print(f\"  ResNet-50 Stage 0 : mean=0.5764, best=0.6116\")\n    print(f\"  UNI Stage 0       : mean=0.6705, best=0.7964\")\n    print(f\"  Fusion Stage 0    : mean={mean_kappa:.4f}, best={best_overall:.4f}\")\n    print(f\"  Gain vs ResNet-50 : {mean_kappa - 0.5764:+.4f}\")\n    print(f\"  Gain vs UNI       : {mean_kappa - 0.6705:+.4f}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, glob\n\nfor f in glob.glob(\"/kaggle/working/stage0_fusion*.pth\"):\n    mb = os.path.getsize(f)/(1024*1024)\n    print(f\"{os.path.basename(f)} ({mb:.0f} MB)\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import subprocess, json, os, shutil, glob\n\ntmp = \"/kaggle/working/_push_fusion\"\nos.makedirs(tmp, exist_ok=True)\n\nfor f in glob.glob(\"/kaggle/working/stage0_fusion*.pth\"):\n    shutil.copy(f, tmp)\n    print(f\"Copied: {os.path.basename(f)}\")\n\njson.dump({\n    \"title\": \"Stage0 Fusion Results Bejaoui 2026\",\n    \"id\": \"bejaouikhouloud/stage0-fusion-results\",\n    \"licenses\": [{\"name\": \"CC0-1.0\"}]\n}, open(f\"{tmp}/dataset-metadata.json\", \"w\"))\n\nr = subprocess.run([\n    \"kaggle\", \"datasets\", \"create\",\n    \"-p\", tmp, \"--dir-mode\", \"skip\"\n], capture_output=True, text=True, timeout=600)\n\nprint(\"OK!\" if r.returncode == 0 else f\"ERREUR: {r.stderr[:200]}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"tmp = \"/kaggle/working/_push_fusion\"\nos.makedirs(tmp, exist_ok=True)\n\n# Copie seulement fold 3 (le meilleur !)\nf = \"/kaggle/working/stage0_fusion_fold3_best.pth\"\nif os.path.exists(f):\n    shutil.copy(f, tmp)\n    print(\"Copied fold3!\")\n\njson.dump({\n    \"title\": \"Stage0 Fusion Results Bejaoui 2026\",\n    \"id\": \"bejaouikhouloud/stage0-fusion-results\",\n    \"licenses\": [{\"name\": \"CC0-1.0\"}]\n}, open(f\"{tmp}/dataset-metadata.json\", \"w\"))\n\nr = subprocess.run([\n    \"kaggle\", \"datasets\", \"create\",\n    \"-p\", tmp, \"--dir-mode\", \"skip\"\n], capture_output=True, text=True, timeout=600)\n\nprint(\"OK!\" if r.returncode == 0 else f\"ERREUR: {r.stderr[:200]}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, shutil, subprocess, json, glob\n\ntmp = \"/kaggle/working/_push_fusion\"\nos.makedirs(tmp, exist_ok=True)\n\nf = \"/kaggle/working/stage0_fusion_fold3_best.pth\"\nif os.path.exists(f):\n    shutil.copy(f, tmp)\n    print(\"Copied fold3!\")\nelse:\n    print(\"Fold3 not found!\")\n    # Cherche tous les checkpoints\n    for ck in glob.glob(\"/kaggle/working/stage0_fusion*.pth\"):\n        print(f\"Found: {os.path.basename(ck)}\")\n\njson.dump({\n    \"title\": \"Stage0 Fusion Results Bejaoui 2026\",\n    \"id\": \"bejaouikhouloud/stage0-fusion-results\",\n    \"licenses\": [{\"name\": \"CC0-1.0\"}]\n}, open(f\"{tmp}/dataset-metadata.json\", \"w\"))\n\nr = subprocess.run([\n    \"kaggle\", \"datasets\", \"create\",\n    \"-p\", tmp, \"--dir-mode\", \"skip\"\n], capture_output=True, text=True, timeout=600)\n\nprint(\"OK!\" if r.returncode == 0 else f\"ERREUR: {r.stderr[:200]}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import glob, os\n\n# Cherche tous les .pth\nfor f in glob.glob(\"/kaggle/working/**/*.pth\", recursive=True):\n    mb = os.path.getsize(f)/(1024*1024)\n    print(f\"{f} ({mb:.0f} MB)\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\n=============================================================\nStage 0 — HybridFusion (ResNet-101 + UNI) — Parallel Fusion\nArchitecture proposée par Dr. Mouelhi :\n  Branch 1 : ResNet-101 → features locales (CNN)\n  Branch 2 : UNI → features globales (Transformer)\n  Fusion    : 0.6 × CNN + 0.4 × UNI → MLP → ISUP\n\nDataset: patches_400 (399 WSIs, 4,319 patches)\nComparaison :\n  ResNet-50 Stage 0  : mean=0.5764, best=0.6116\n  UNI Stage 0        : mean=0.6705, best=0.7964\n  ResNet+UNI Stage 0 : ???\n=============================================================\nAttache dans Kaggle :\n  khouloudbejaoui20/dataset\nGPU : T4x2\n\"\"\"\n\nimport os, glob, json, random, shutil, subprocess, datetime\nimport numpy as np\nfrom PIL import Image\nfrom collections import Counter\nfrom sklearn.model_selection import StratifiedGroupKFold\nfrom sklearn.metrics import cohen_kappa_score\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as T\nimport torchvision.models as models\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\n\n# ─────────────────────────────────────────────\n# 1. LOGIN HUGGINGFACE\n# ─────────────────────────────────────────────\nfrom huggingface_hub import login, whoami\nimport timm\n\nHF_TOKEN = \"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\"\nlogin(token=HF_TOKEN)\ninfo = whoami()\nprint(f\"OK HuggingFace: {info['name']}\")\n\n# ─────────────────────────────────────────────\n# 2. CONFIG\n# ─────────────────────────────────────────────\nSHARD_ID    = 0\nSTART_FOLD  = 3        # Fold 4 (index commence à 0)\n\nDONE_KAPPAS = [0.6789, 0.7256, 0.7835]  # Folds 1+2+3\n\nCNN_BACKBONE = \"resnet101\"  # resnet50 or resnet101\nFREEZE_UNI   = True         # UNI freezé pour éviter overfitting\nALPHA        = 0.6          # poids CNN\nBETA         = 0.4          # poids UNI\n\nN_FOLDS     = 5\nSEED        = 42\nBATCH_SIZE  = 8    # réduit car 2 modèles en parallèle\nNUM_WORKERS = 2\nEPOCHS      = 20\nPATIENCE    = 5\nLR_CNN      = 5e-5\nLR_TR       = 4e-5  # pour les couches non-UNI\n\nrandom.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)\ntorch.backends.cudnn.deterministic = True\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Device      : {DEVICE}\")\nprint(f\"Architecture: HybridFusion ({CNN_BACKBONE} + UNI)\")\nprint(f\"Fusion      : {ALPHA} x CNN + {BETA} x UNI\")\nprint(f\"Freeze UNI  : {FREEZE_UNI}\")\nprint(f\"BATCH={BATCH_SIZE} | EPOCHS={EPOCHS} | PATIENCE={PATIENCE}\")\n\n# ─────────────────────────────────────────────\n# 3. CHEMINS\n# ─────────────────────────────────────────────\nBASE = \"/kaggle/input/datasets/khouloudbejaoui20\"\nSHARD_PATH = f\"{BASE}/dataset/patches_400-20260630T212405Z-3-001/patches_400\"\n\nSAVE_DIR     = \"/kaggle/working\"\nKAGGLE_DS    = \"bejaouikhouloud/stage0-fusion-results\"\nRESULTS_JSON = f\"{SAVE_DIR}/stage0_fusion_cv_results.json\"\n\n# ─────────────────────────────────────────────\n# 4. LOAD PATCHES\n# ─────────────────────────────────────────────\nprint(f\"\\n── Loading patches ──────────────────────\")\n\ndef load_patches(folder):\n    files = glob.glob(f\"{folder}/*.png\")\n    if not files:\n        files = glob.glob(f\"{folder}/**/*.png\", recursive=True)\n    records = []\n    for fp in files:\n        fname = os.path.basename(fp)\n        try:\n            parts  = fname.split(\"_slide\")\n            label  = int(parts[0].replace(\"grade\", \"\"))\n            wsi_id = parts[1].split(\"_\")[0]\n            records.append({\"path\": fp, \"wsi_id\": wsi_id, \"label\": label})\n        except:\n            continue\n    return records\n\nall_records = load_patches(SHARD_PATH)\nprint(f\"  Shard 0: {len(all_records)} patches\")\n\nall_wsis = list(set(r[\"wsi_id\"] for r in all_records))\nprint(f\"  Total WSIs: {len(all_wsis)}\")\ncnt = Counter(r[\"label\"] for r in all_records)\nfor c in range(6):\n    print(f\"    ISUP {c}: {cnt.get(c, 0)}\")\n\n# ─────────────────────────────────────────────\n# 5. TRANSFORMS\n# ─────────────────────────────────────────────\nMEAN = [0.485, 0.456, 0.406]\nSTD  = [0.229, 0.224, 0.225]\n\ntrain_tf = T.Compose([\n    T.Resize((224, 224)),\n    T.RandomHorizontalFlip(),\n    T.RandomVerticalFlip(),\n    T.RandomRotation(90),\n    T.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.2, hue=0.1),\n    T.RandomApply([T.GaussianBlur(5)], p=0.3),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n    T.RandomApply([T.RandomErasing(p=0.5, scale=(0.02, 0.2))], p=0.2),\n])\n\nval_tf = T.Compose([\n    T.Resize((224, 224)),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n])\n\nclass PatchDataset(Dataset):\n    def __init__(self, records, transform):\n        self.records   = records\n        self.transform = transform\n    def __len__(self): return len(self.records)\n    def __getitem__(self, idx):\n        r = self.records[idx]\n        try:\n            img = Image.open(r[\"path\"]).convert(\"RGB\")\n        except:\n            img = Image.new(\"RGB\", (224, 224), (255, 255, 255))\n        return self.transform(img), int(r[\"label\"])\n\n# ─────────────────────────────────────────────\n# 6. MODELE — HybridFusion\n# ─────────────────────────────────────────────\nclass HybridFusion(nn.Module):\n    \"\"\"\n    Architecture fusion parallèle :\n    Branch 1 : ResNet-101 → features locales (CNN)\n    Branch 2 : UNI (freezé) → features globales\n    Fusion   : alpha × CNN + beta × UNI → MLP → ISUP\n\n    Les deux branches reçoivent la MÊME image en entrée.\n    \"\"\"\n    def __init__(self, cnn_backbone=\"resnet101\",\n                 num_classes=6, d_model=512,\n                 alpha=0.6, beta=0.4,\n                 freeze_uni=True, dropout=0.1):\n        super().__init__()\n\n        self.alpha = alpha\n        self.beta  = beta\n\n        # ── Branch 1 : CNN ──────────────────────\n        if cnn_backbone == \"resnet50\":\n            net = models.resnet50(weights=models.ResNet50_Weights.DEFAULT)\n            cnn_dim = 2048\n        elif cnn_backbone == \"resnet101\":\n            net = models.resnet101(weights=models.ResNet101_Weights.DEFAULT)\n            cnn_dim = 2048\n        elif cnn_backbone == \"resnet152\":\n            net = models.resnet152(weights=models.ResNet152_Weights.DEFAULT)\n            cnn_dim = 2048\n        else:\n            raise ValueError(f\"Unknown backbone: {cnn_backbone}\")\n\n        self.cnn_backbone = nn.Sequential(*list(net.children())[:-2])\n        self.pool         = nn.AdaptiveAvgPool2d((1, 1))\n        self.proj_cnn     = nn.Sequential(\n            nn.Linear(cnn_dim, d_model),\n            nn.LayerNorm(d_model),\n            nn.ReLU(),\n            nn.Dropout(dropout)\n        )\n        print(f\"  CNN Branch : {cnn_backbone} → {cnn_dim}-d → {d_model}-d\")\n\n        # ── Branch 2 : UNI ──────────────────────\n        self.uni_backbone = timm.create_model(\n            \"hf-hub:MahmoodLab/UNI\",\n            pretrained=True,\n            num_classes=0,\n            init_values=1e-5,\n            dynamic_img_size=True\n        )\n        uni_dim = self.uni_backbone.num_features  # 1024\n\n        if freeze_uni:\n            for param in self.uni_backbone.parameters():\n                param.requires_grad = False\n            print(f\"  UNI Branch : UNI FREEZÉ → {uni_dim}-d → {d_model}-d\")\n        else:\n            print(f\"  UNI Branch : UNI trainable → {uni_dim}-d → {d_model}-d\")\n\n        self.proj_uni = nn.Sequential(\n            nn.Linear(uni_dim, d_model),\n            nn.LayerNorm(d_model),\n            nn.ReLU(),\n            nn.Dropout(dropout)\n        )\n\n        # ── Attention Mechanism (patch-par-patch) ──\n        # Pour chaque patch, calcule ses propres poids !\n        # Input : concat(CNN, UNI) = 1024-d\n        # Output : (alpha, beta) par patch\n        self.attention = nn.Sequential(\n            nn.Linear(d_model * 2, 128),\n            nn.Tanh(),\n            nn.Linear(128, 2),\n            nn.Softmax(dim=1)\n        )\n        print(f\"  Fusion     : Attention patch-par-patch (CNN+UNI)\")\n\n        # ── Classifier ──────────────────────────\n        self.classifier = nn.Sequential(\n            nn.Linear(d_model, 256),\n            nn.ReLU(),\n            nn.Dropout(dropout),\n            nn.Linear(256, num_classes)\n        )\n\n    def forward(self, x):\n        # Branch 1 : CNN features\n        cnn_feat = self.pool(\n            self.cnn_backbone(x)\n        ).flatten(1)                      # (B, 2048)\n        cnn_feat = self.proj_cnn(cnn_feat) # (B, 512)\n\n        # Branch 2 : UNI features\n        if not self.uni_backbone.training:\n            with torch.no_grad():\n                uni_feat = self.uni_backbone(x)  # (B, 1024)\n        else:\n            uni_feat = self.uni_backbone(x)\n        uni_feat = self.proj_uni(uni_feat)  # (B, 512)\n\n        # Attention patch-par-patch !\n        combined = torch.cat([cnn_feat, uni_feat], dim=1)  # (B, 1024)\n        weights  = self.attention(combined)                 # (B, 2)\n        alpha    = weights[:, 0:1]  # (B, 1) CNN weight par patch\n        beta     = weights[:, 1:2]  # (B, 1) UNI weight par patch\n        fused    = alpha * cnn_feat + beta * uni_feat       # (B, 512)\n\n        return self.classifier(fused)\n\ndef build_model():\n    torch.cuda.empty_cache()\n    import gc; gc.collect()\n\n    model  = HybridFusion(\n        cnn_backbone=CNN_BACKBONE,\n        alpha=ALPHA,\n        beta=BETA,\n        freeze_uni=FREEZE_UNI\n    ).to(DEVICE)\n\n    n_gpus = torch.cuda.device_count()\n    if n_gpus > 1:\n        model = nn.DataParallel(model)\n        print(f\"  OK DataParallel: {n_gpus} GPUs\")\n\n    n_total     = sum(p.numel() for p in model.parameters())\n    n_trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    print(f\"  Total params    : {n_total:,}\")\n    print(f\"  Trainable params: {n_trainable:,}\")\n    return model\n\n# ─────────────────────────────────────────────\n# 7. WSI-LEVEL QWK\n# ─────────────────────────────────────────────\ndef wsi_level_qwk(model, val_data, device):\n    model.eval()\n    wsi_preds  = {}\n    wsi_labels = {}\n    all_preds  = []\n    all_labels = []\n    loader = DataLoader(\n        PatchDataset(val_data, val_tf),\n        batch_size=BATCH_SIZE, shuffle=False,\n        num_workers=NUM_WORKERS\n    )\n    with torch.no_grad():\n        for imgs, labels in loader:\n            imgs  = imgs.to(device)\n            preds = model(imgs).argmax(1).cpu().numpy()\n            all_preds.extend(preds)\n            all_labels.extend(labels.numpy())\n    for i, r in enumerate(val_data):\n        wid = r[\"wsi_id\"]\n        if wid not in wsi_preds:\n            wsi_preds[wid]  = []\n            wsi_labels[wid] = r[\"label\"]\n        wsi_preds[wid].append(all_preds[i])\n    wf = []; wl = []\n    for wid in wsi_preds:\n        votes = wsi_preds[wid]\n        wf.append(max(set(votes), key=votes.count))\n        wl.append(wsi_labels[wid])\n    wsi_qwk = cohen_kappa_score(wl, wf, weights=\"quadratic\")\n    return wsi_qwk, len(wf)\n\n# ─────────────────────────────────────────────\n# 8. PUSH KAGGLE\n# ─────────────────────────────────────────────\ndef push_to_kaggle(ckpt_path, message):\n    try:\n        tmp = f\"{SAVE_DIR}/_push_fusion\"\n        os.makedirs(tmp, exist_ok=True)\n        for c in glob.glob(f\"{SAVE_DIR}/stage{SHARD_ID}_fusion*.pth\"):\n            shutil.copy(c, tmp)\n        shutil.copy(ckpt_path, tmp)\n        if os.path.exists(RESULTS_JSON):\n            shutil.copy(RESULTS_JSON, tmp)\n        json.dump({\n            \"title\": \"Stage0 Fusion Results Bejaoui 2026\",\n            \"id\": KAGGLE_DS,\n            \"licenses\": [{\"name\": \"CC0-1.0\"}]\n        }, open(f\"{tmp}/dataset-metadata.json\", \"w\"))\n        r = subprocess.run(\n            [\"kaggle\", \"datasets\", \"create\",\n             \"-p\", tmp, \"--dir-mode\", \"skip\"],\n            capture_output=True, text=True, timeout=600\n        )\n        if r.returncode != 0:\n            r = subprocess.run(\n                [\"kaggle\", \"datasets\", \"version\",\n                 \"-p\", tmp, \"-m\", message, \"--dir-mode\", \"skip\"],\n                capture_output=True, text=True, timeout=600\n            )\n        print(f\"   {'OK!' if r.returncode==0 else 'ERREUR: '+r.stderr[:100]}\")\n    except Exception as e:\n        print(f\"   ERREUR: {e}\")\n\n# ─────────────────────────────────────────────\n# 9. TRAIN ONE FOLD\n# ─────────────────────────────────────────────\ndef train_one_fold(fold_id, train_data, val_data, n_folds=5):\n    print(f\"\\n  Fold {fold_id+1}/{n_folds} ──────────────────────────\")\n    print(f\"     Train: {len(train_data)} | Val: {len(val_data)}\")\n\n    train_loader = DataLoader(\n        PatchDataset(train_data, train_tf),\n        batch_size=BATCH_SIZE, shuffle=True,\n        num_workers=NUM_WORKERS, pin_memory=True\n    )\n\n    model = build_model()\n\n    # Class weights\n    cnt_tr = Counter(r[\"label\"] for r in train_data)\n    w = torch.tensor(\n        [1.0 / max(cnt_tr.get(c, 1), 1) for c in range(6)],\n        dtype=torch.float\n    )\n    w = (w / w.sum() * 6).to(DEVICE)\n    criterion = nn.CrossEntropyLoss(weight=w, label_smoothing=0.1)\n\n    # Optimizer\n    base_model = model.module if hasattr(model, \"module\") else model\n    trainable  = [p for p in model.parameters() if p.requires_grad]\n    optimizer  = optim.AdamW([\n        {\"params\": base_model.cnn_backbone.parameters(), \"lr\": LR_CNN},\n        {\"params\": [p for n, p in base_model.named_parameters()\n                    if not n.startswith(\"cnn_backbone\") and p.requires_grad],\n         \"lr\": LR_TR},\n    ], weight_decay=1e-4)\n\n    scheduler = CosineAnnealingLR(optimizer, T_max=EPOCHS, eta_min=1e-7)\n\n    best_kappa = 0.0\n    no_improve = 0\n    best_ckpt  = f\"{SAVE_DIR}/stage{SHARD_ID}_fusion_fold{fold_id+1}_best.pth\"\n\n    for epoch in range(1, EPOCHS + 1):\n        model.train()\n        # UNI reste en eval mode si freezé\n        if FREEZE_UNI:\n            base = model.module if hasattr(model, \"module\") else model\n            base.uni_backbone.eval()\n\n        t_loss, t_correct, t_total = 0.0, 0, 0\n        for imgs, labels in train_loader:\n            imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n            optimizer.zero_grad()\n            out  = model(imgs)\n            loss = criterion(out, labels)\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n            optimizer.step()\n            t_loss    += loss.item() * imgs.size(0)\n            t_correct += (out.argmax(1) == labels).sum().item()\n            t_total   += imgs.size(0)\n        scheduler.step()\n\n        wsi_kappa, n_wsis_val = wsi_level_qwk(model, val_data, DEVICE)\n        is_best = wsi_kappa > best_kappa\n\n        print(f\"     Ep[{epoch:02d}/{EPOCHS}] \"\n              f\"loss={t_loss/t_total:.4f} acc={t_correct/t_total:.3f} | \"\n              f\"WSI_QWK={wsi_kappa:.4f} ({n_wsis_val} WSIs)\"\n              + (\" *\" if is_best else \"\"))\n\n        if is_best:\n            best_kappa = wsi_kappa\n            no_improve = 0\n            # Pas de poids globaux à afficher\n            # (attention patch-par-patch)\n            torch.save({\n                \"epoch\":            epoch,\n                \"model_state_dict\": model.state_dict() if not hasattr(model, \"module\")\n                                    else model.module.state_dict(),\n                \"kappa\":            wsi_kappa,\n                \"cnn_backbone\":     CNN_BACKBONE,\n                \"alpha\":            ALPHA,\n                \"beta\":             BETA,\n                \"freeze_uni\":       FREEZE_UNI,\n                \"fusion\":           \"attention_patch_level\",\n                \"shard_id\":         SHARD_ID,\n                \"fold_id\":          fold_id,\n                \"n_wsis\":           len(all_wsis),\n            }, best_ckpt)\n        else:\n            no_improve += 1\n            if no_improve >= PATIENCE:\n                print(f\"     Early stopping epoch {epoch}\")\n                break\n\n    print(f\"  OK Fold {fold_id+1} best WSI_QWK: {best_kappa:.4f}\")\n    push_to_kaggle(best_ckpt,\n        f\"Stage{SHARD_ID} Fusion Fold{fold_id+1} QWK={best_kappa:.4f}\")\n    return best_kappa, best_ckpt\n\n# ─────────────────────────────────────────────\n# 10. 5-FOLD CV\n# ─────────────────────────────────────────────\nprint(f\"\\n{'='*60}\")\nprint(f\"Stage {SHARD_ID} HybridFusion — {len(all_wsis)} WSIs — 5-Fold CV\")\nprint(f\"Architecture : {CNN_BACKBONE} + UNI + Attention\")\nprint(f\"Fusion       : {ALPHA} x CNN + {BETA} x UNI\")\nprint(f\"Baseline     : ResNet-50=0.5764 | UNI=0.6705\")\nprint(f\"{'='*60}\")\n\nwsi_ids    = [r[\"wsi_id\"] for r in all_records]\nlabels_arr = [r[\"label\"]  for r in all_records]\nindices    = list(range(len(all_records)))\n\nsgkf = StratifiedGroupKFold(n_splits=N_FOLDS, shuffle=True, random_state=SEED)\n\nfold_kappas       = DONE_KAPPAS.copy()\nbest_overall      = max(DONE_KAPPAS) if DONE_KAPPAS else -1.0\nbest_ckpt_overall = None\n\nfor fold_id, (train_idx, val_idx) in enumerate(\n        sgkf.split(indices, labels_arr, groups=wsi_ids)):\n\n    if fold_id < START_FOLD:\n        continue\n\n    train_data = [all_records[i] for i in train_idx]\n    val_data   = [all_records[i] for i in val_idx]\n\n    tw = set(all_records[i][\"wsi_id\"] for i in train_idx)\n    vw = set(all_records[i][\"wsi_id\"] for i in val_idx)\n    assert len(tw & vw) == 0, f\"DATA LEAKAGE fold {fold_id+1}!\"\n\n    kappa, ckpt_path = train_one_fold(fold_id, train_data, val_data, N_FOLDS)\n    fold_kappas.append(kappa)\n\n    if kappa > best_overall:\n        best_overall      = kappa\n        best_ckpt_overall = ckpt_path\n\n# ─────────────────────────────────────────────\n# 11. RESUME FINAL\n# ─────────────────────────────────────────────\nmean_kappa = np.mean(fold_kappas)\nstd_kappa  = np.std(fold_kappas)\n\nprint(f\"\\n{'='*60}\")\nprint(f\"OK Stage {SHARD_ID} HybridFusion — 5-Fold CV Results\")\nprint(f\"{'='*60}\")\nfor i, k in enumerate(fold_kappas):\n    print(f\"   Fold {i+1}: QWK = {k:.4f}\")\nprint(f\"{'─'*40}\")\nprint(f\"   Mean QWK Fusion   : {mean_kappa:.4f} +/- {std_kappa:.4f}\")\nprint(f\"   Best QWK Fusion   : {best_overall:.4f}\")\nprint(f\"   ResNet-50 Stage 0 : mean=0.5764, best=0.6116\")\nprint(f\"   UNI Stage 0       : mean=0.6705, best=0.7964\")\nprint(f\"   Gain vs ResNet-50 : {mean_kappa - 0.5764:+.4f}\")\nprint(f\"   Gain vs UNI       : {mean_kappa - 0.6705:+.4f}\")\nprint(f\"   WSIs              : {len(all_wsis)}\")\nprint(f\"   SOTA [1]          : 0.934\")\nprint(f\"{'='*60}\")\n\nresults = {\n    \"stage\":            SHARD_ID,\n    \"architecture\":     f\"HybridFusion_{CNN_BACKBONE}_UNI\",\n    \"cnn_backbone\":     CNN_BACKBONE,\n    \"freeze_uni\":       FREEZE_UNI,\n    \"alpha\":            ALPHA,\n    \"beta\":             BETA,\n    \"n_wsis\":           len(all_wsis),\n    \"n_patches\":        len(all_records),\n    \"fold_kappas\":      fold_kappas,\n    \"mean_kappa\":       mean_kappa,\n    \"std_kappa\":        std_kappa,\n    \"best_kappa\":       best_overall,\n    \"resnet50_mean\":    0.5764,\n    \"uni_mean\":         0.6705,\n    \"gain_vs_resnet50\": mean_kappa - 0.5764,\n    \"gain_vs_uni\":      mean_kappa - 0.6705,\n    \"timestamp\":        datetime.datetime.now().isoformat(),\n}\nwith open(RESULTS_JSON, \"w\") as f:\n    json.dump(results, f, indent=2)\n\nif best_ckpt_overall:\n    stage_best = f\"{SAVE_DIR}/stage{SHARD_ID}_fusion_best_qwk{best_overall:.4f}.pth\"\n    shutil.copy(best_ckpt_overall, stage_best)\n    print(f\"\\nOK Best checkpoint: {stage_best}\")\n    print(f\"\\nComparaison finale:\")\n    print(f\"  ResNet-50 Stage 0 : mean=0.5764, best=0.6116\")\n    print(f\"  UNI Stage 0       : mean=0.6705, best=0.7964\")\n    print(f\"  Fusion Stage 0    : mean={mean_kappa:.4f}, best={best_overall:.4f}\")\n    print(f\"  Gain vs ResNet-50 : {mean_kappa - 0.5764:+.4f}\")\n    print(f\"  Gain vs UNI       : {mean_kappa - 0.6705:+.4f}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\n=============================================================\nStage 0 — HybridFusion (ResNet-101 + UNI non-freezé)\n1000 WSIs ALÉATOIRES depuis PANDA\nLR_UNI = 1e-6 (très faible pour éviter catastrophic forgetting)\nBATCH_SIZE = 2 (évite OOM)\n\nComparaison :\n  ResNet-50 Stage 0       : mean=0.5764\n  UNI non-freezé Stage 0  : mean=0.6705\n  Fusion freezé Stage 0   : mean=0.6422\n  Fusion non-freezé 1000  : ???\n=============================================================\nAttache dans Kaggle :\n  prostate-cancer-grade-assessment (competition)\nGPU : T4x2\n\"\"\"\n\nimport os, glob, random, json, shutil, subprocess, datetime\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\nfrom collections import Counter\nfrom concurrent.futures import ThreadPoolExecutor, as_completed\nfrom sklearn.model_selection import StratifiedGroupKFold\nfrom sklearn.metrics import cohen_kappa_score\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as T\nimport torchvision.models as models\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\n\n# ─────────────────────────────────────────────\n# 1. LOGIN HUGGINGFACE\n# ─────────────────────────────────────────────\nfrom huggingface_hub import login, whoami\nimport timm\n\nHF_TOKEN = \"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\"\nlogin(token=HF_TOKEN)\ninfo = whoami()\nprint(f\"OK HuggingFace: {info['name']}\")\n\n# ─────────────────────────────────────────────\n# 2. CONFIG\n# ─────────────────────────────────────────────\nSEED        = 42\nN_WSI       = 1000   # 1000 WSIs aléatoires !\nPATCH_SIZE  = 512\nTISSUE_THR  = 0.65\nMAX_PATCHES = 15\n\n# Training\nFREEZE_UNI  = False  # NON-FREEZÉ !\nBATCH_SIZE  = 2      # très petit pour éviter OOM !\nNUM_WORKERS = 2\nEPOCHS      = 15\nPATIENCE    = 5\nLR_CNN      = 2e-5\nLR_UNI      = 1e-6   # très faible ! préserve features UNI\nLR_TR       = 4e-5\nN_FOLDS     = 5\n\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\ntorch.backends.cudnn.deterministic = True\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Device      : {DEVICE}\")\nprint(f\"N_WSI       : {N_WSI} aléatoires\")\nprint(f\"Freeze UNI  : {FREEZE_UNI}\")\nprint(f\"BATCH_SIZE  : {BATCH_SIZE}\")\nprint(f\"LR_UNI      : {LR_UNI}\")\n\n# ─────────────────────────────────────────────\n# 3. CHEMINS\n# ─────────────────────────────────────────────\nPANDA_IMG_DIR  = \"/kaggle/input/competitions/prostate-cancer-grade-assessment/train_images\"\nPANDA_CSV_PATH = \"/kaggle/input/competitions/prostate-cancer-grade-assessment/train.csv\"\nOUT_DIR        = \"/kaggle/working/patches_random_1000\"\nSAVE_DIR       = \"/kaggle/working\"\nRESULTS_JSON   = f\"{SAVE_DIR}/stage0_fusion_nonfrozen_1000_results.json\"\nKAGGLE_DS      = \"bejaouikhouloud/stage0-fusion-nonfrozen-results\"\n\nos.makedirs(OUT_DIR, exist_ok=True)\n\n# ─────────────────────────────────────────────\n# 4. SELECTION ALEATOIRE 1000 WSIs\n# ─────────────────────────────────────────────\nprint(f\"\\n── Sélection aléatoire {N_WSI} WSIs ──────────\")\n\ndf = pd.read_csv(PANDA_CSV_PATH)\nprint(f\"Total PANDA : {len(df)} WSIs\")\n\ndf_shuffled = df.sample(frac=1, random_state=SEED).reset_index(drop=True)\ndf_random   = df_shuffled.iloc[:N_WSI].reset_index(drop=True)\n\nprint(f\"\\nDistribution {N_WSI} WSIs aléatoires :\")\nfor g in range(6):\n    n = (df_random['isup_grade'] == g).sum()\n    print(f\"  ISUP {g}: {n}\")\n\n# ─────────────────────────────────────────────\n# 5. EXTRACTION PATCHES\n# ─────────────────────────────────────────────\nprint(f\"\\n── Extraction patches ───────────────────\")\n\nexisting  = glob.glob(f\"{OUT_DIR}/*.png\")\ndone_wsis = set()\nfor f in existing:\n    try:\n        wsi_id = os.path.basename(f).split(\"_slide\")[1].split(\"_\")[0]\n        done_wsis.add(wsi_id)\n    except:\n        pass\nprint(f\"Déjà extraits: {len(existing)} patches ({len(done_wsis)} WSIs)\")\n\ndef has_tissue(patch_arr):\n    gray = np.mean(patch_arr, axis=2)\n    return (gray < 220).sum() / gray.size >= TISSUE_THR\n\ndef extract_wsi(row):\n    import openslide\n    wsi_id = str(row['image_id'])\n    label  = int(row['isup_grade'])\n    if wsi_id in done_wsis:\n        return 0, label, wsi_id, \"skipped\"\n    wsi_path = os.path.join(PANDA_IMG_DIR, f\"{wsi_id}.tiff\")\n    if not os.path.exists(wsi_path):\n        return 0, label, wsi_id, \"not_found\"\n    try:\n        slide  = openslide.OpenSlide(wsi_path)\n        w, h   = slide.dimensions\n        saved  = 0\n        positions = [(x, y)\n                     for y in range(0, h - PATCH_SIZE, PATCH_SIZE)\n                     for x in range(0, w - PATCH_SIZE, PATCH_SIZE)]\n        random.shuffle(positions)\n        for x, y in positions:\n            if saved >= MAX_PATCHES:\n                break\n            try:\n                patch     = slide.read_region((x, y), 0, (PATCH_SIZE, PATCH_SIZE))\n                patch_rgb = np.array(patch.convert(\"RGB\"))\n                if has_tissue(patch_rgb):\n                    fname = f\"grade{label}_slide{wsi_id}_{saved:03d}.png\"\n                    Image.fromarray(patch_rgb).save(\n                        os.path.join(OUT_DIR, fname), \"PNG\")\n                    saved += 1\n            except:\n                continue\n        slide.close()\n        return saved, label, wsi_id, \"ok\"\n    except Exception as e:\n        return 0, label, wsi_id, \"error\"\n\nrows = [row for _, row in df_random.iterrows()\n        if str(row['image_id']) not in done_wsis]\n\nif rows:\n    print(f\"WSIs à extraire : {len(rows)}\")\n    total_patches = len(existing)\n    total_ok = len(done_wsis)\n    total_failed = 0\n    with ThreadPoolExecutor(max_workers=2) as executor:\n        futures = {executor.submit(extract_wsi, row): row for row in rows}\n        for i, future in enumerate(as_completed(futures)):\n            try:\n                n, grade, wsi_id, status = future.result()\n                if status in [\"ok\", \"skipped\"]:\n                    total_patches += n\n                    total_ok += 1\n                else:\n                    total_failed += 1\n                if (i+1) % 100 == 0 or (i+1) == len(rows):\n                    print(f\"  [{i+1:4d}/{len(rows)}] OK={total_ok} | \"\n                          f\"Patches={total_patches}\")\n            except:\n                total_failed += 1\nelse:\n    print(\"Tous les WSIs déjà extraits !\")\n\nfiles = glob.glob(f\"{OUT_DIR}/*.png\")\nprint(f\"\\nExtraction terminée: {len(files)} patches\")\n\n# ─────────────────────────────────────────────\n# 6. LOAD PATCHES\n# ─────────────────────────────────────────────\nprint(f\"\\n── Chargement patches ───────────────────\")\n\ndef load_patches(folder):\n    files = glob.glob(f\"{folder}/*.png\")\n    records = []\n    for fp in files:\n        fname = os.path.basename(fp)\n        try:\n            parts  = fname.split(\"_slide\")\n            label  = int(parts[0].replace(\"grade\", \"\"))\n            wsi_id = parts[1].split(\"_\")[0]\n            records.append({\"path\": fp, \"wsi_id\": wsi_id, \"label\": label})\n        except:\n            continue\n    return records\n\nall_records = load_patches(OUT_DIR)\nall_wsis    = list(set(r[\"wsi_id\"] for r in all_records))\nprint(f\"Total patches : {len(all_records)}\")\nprint(f\"Total WSIs    : {len(all_wsis)}\")\ncnt = Counter(r[\"label\"] for r in all_records)\nfor c in range(6):\n    print(f\"  ISUP {c}: {cnt.get(c, 0)}\")\n\n# ─────────────────────────────────────────────\n# 7. TRANSFORMS\n# ─────────────────────────────────────────────\nMEAN = [0.485, 0.456, 0.406]\nSTD  = [0.229, 0.224, 0.225]\n\ntrain_tf = T.Compose([\n    T.Resize((224, 224)),\n    T.RandomHorizontalFlip(),\n    T.RandomVerticalFlip(),\n    T.RandomRotation(90),\n    T.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.2, hue=0.1),\n    T.RandomApply([T.GaussianBlur(5)], p=0.3),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n    T.RandomApply([T.RandomErasing(p=0.5, scale=(0.02, 0.2))], p=0.2),\n])\n\nval_tf = T.Compose([\n    T.Resize((224, 224)),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n])\n\nclass PatchDataset(Dataset):\n    def __init__(self, records, transform):\n        self.records   = records\n        self.transform = transform\n    def __len__(self): return len(self.records)\n    def __getitem__(self, idx):\n        r = self.records[idx]\n        try:\n            img = Image.open(r[\"path\"]).convert(\"RGB\")\n        except:\n            img = Image.new(\"RGB\", (224, 224), (255, 255, 255))\n        return self.transform(img), int(r[\"label\"])\n\n# ─────────────────────────────────────────────\n# 8. MODELE — HybridFusion non-freezé\n# ─────────────────────────────────────────────\nclass HybridFusionNonFrozen(nn.Module):\n    def __init__(self, num_classes=6, d_model=512,\n                 freeze_uni=False, dropout=0.1):\n        super().__init__()\n\n        # ResNet-101\n        net = models.resnet101(weights=models.ResNet101_Weights.DEFAULT)\n        self.cnn_backbone = nn.Sequential(*list(net.children())[:-2])\n        self.pool         = nn.AdaptiveAvgPool2d((1, 1))\n        self.proj_cnn     = nn.Sequential(\n            nn.Linear(2048, d_model),\n            nn.LayerNorm(d_model),\n            nn.ReLU(),\n            nn.Dropout(dropout)\n        )\n\n        # UNI\n        self.uni_backbone = timm.create_model(\n            \"hf-hub:MahmoodLab/UNI\",\n            pretrained=True,\n            num_classes=0,\n            init_values=1e-5,\n            dynamic_img_size=True\n        )\n        uni_dim = self.uni_backbone.num_features  # 1024\n\n        if freeze_uni:\n            for param in self.uni_backbone.parameters():\n                param.requires_grad = False\n            print(f\"  UNI : FREEZÉ\")\n        else:\n            print(f\"  UNI : NON-FREEZÉ (LR={LR_UNI})\")\n\n        self.proj_uni = nn.Sequential(\n            nn.Linear(uni_dim, d_model),\n            nn.LayerNorm(d_model),\n            nn.ReLU(),\n            nn.Dropout(dropout)\n        )\n\n        # Attention patch-par-patch\n        self.attention = nn.Sequential(\n            nn.Linear(d_model * 2, 128),\n            nn.Tanh(),\n            nn.Linear(128, 2),\n            nn.Softmax(dim=1)\n        )\n\n        # Classifier\n        self.classifier = nn.Sequential(\n            nn.Linear(d_model, 256),\n            nn.ReLU(),\n            nn.Dropout(dropout),\n            nn.Linear(256, num_classes)\n        )\n\n    def forward(self, x):\n        # CNN branch\n        cnn_feat = self.pool(\n            self.cnn_backbone(x)\n        ).flatten(1)\n        cnn_feat = self.proj_cnn(cnn_feat)  # (B, 512)\n\n        # UNI branch\n        uni_feat = self.uni_backbone(x)     # (B, 1024)\n        uni_feat = self.proj_uni(uni_feat)  # (B, 512)\n\n        # Attention\n        combined = torch.cat([cnn_feat, uni_feat], dim=1)  # (B, 1024)\n        weights  = self.attention(combined)                 # (B, 2)\n        alpha    = weights[:, 0:1]\n        beta     = weights[:, 1:2]\n        fused    = alpha * cnn_feat + beta * uni_feat       # (B, 512)\n\n        return self.classifier(fused)\n\ndef build_model():\n    torch.cuda.empty_cache()\n    import gc; gc.collect()\n\n    model  = HybridFusionNonFrozen(freeze_uni=FREEZE_UNI).to(DEVICE)\n    n_gpus = torch.cuda.device_count()\n    if n_gpus > 1:\n        model = nn.DataParallel(model)\n        print(f\"  OK DataParallel: {n_gpus} GPUs\")\n\n    n_total     = sum(p.numel() for p in model.parameters())\n    n_trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    print(f\"  Total params    : {n_total:,}\")\n    print(f\"  Trainable params: {n_trainable:,}\")\n    return model\n\n# ─────────────────────────────────────────────\n# 9. WSI-LEVEL QWK\n# ─────────────────────────────────────────────\ndef wsi_level_qwk(model, val_data, device):\n    model.eval()\n    wsi_preds  = {}\n    wsi_labels = {}\n    all_preds  = []\n    all_labels = []\n    loader = DataLoader(\n        PatchDataset(val_data, val_tf),\n        batch_size=BATCH_SIZE, shuffle=False,\n        num_workers=NUM_WORKERS\n    )\n    with torch.no_grad():\n        for imgs, labels in loader:\n            imgs  = imgs.to(device)\n            preds = model(imgs).argmax(1).cpu().numpy()\n            all_preds.extend(preds)\n            all_labels.extend(labels.numpy())\n    for i, r in enumerate(val_data):\n        wid = r[\"wsi_id\"]\n        if wid not in wsi_preds:\n            wsi_preds[wid]  = []\n            wsi_labels[wid] = r[\"label\"]\n        wsi_preds[wid].append(all_preds[i])\n    wf = []; wl = []\n    for wid in wsi_preds:\n        votes = wsi_preds[wid]\n        wf.append(max(set(votes), key=votes.count))\n        wl.append(wsi_labels[wid])\n    wsi_qwk = cohen_kappa_score(wl, wf, weights=\"quadratic\")\n    return wsi_qwk, len(wf)\n\n# ─────────────────────────────────────────────\n# 10. PUSH KAGGLE\n# ─────────────────────────────────────────────\ndef push_to_kaggle(ckpt_path, message):\n    try:\n        tmp = f\"{SAVE_DIR}/_push_nonfrozen\"\n        os.makedirs(tmp, exist_ok=True)\n        for c in glob.glob(f\"{SAVE_DIR}/fusion_nonfrozen*.pth\"):\n            shutil.copy(c, tmp)\n        shutil.copy(ckpt_path, tmp)\n        if os.path.exists(RESULTS_JSON):\n            shutil.copy(RESULTS_JSON, tmp)\n        json.dump({\n            \"title\": \"Stage0 Fusion NonFrozen Results 2026\",\n            \"id\": KAGGLE_DS,\n            \"licenses\": [{\"name\": \"CC0-1.0\"}]\n        }, open(f\"{tmp}/dataset-metadata.json\", \"w\"))\n        r = subprocess.run(\n            [\"kaggle\", \"datasets\", \"create\",\n             \"-p\", tmp, \"--dir-mode\", \"skip\"],\n            capture_output=True, text=True, timeout=600\n        )\n        if r.returncode != 0:\n            r = subprocess.run(\n                [\"kaggle\", \"datasets\", \"version\",\n                 \"-p\", tmp, \"-m\", message, \"--dir-mode\", \"skip\"],\n                capture_output=True, text=True, timeout=600\n            )\n        print(f\"   {'OK!' if r.returncode==0 else 'ERREUR: '+r.stderr[:100]}\")\n    except Exception as e:\n        print(f\"   ERREUR: {e}\")\n\n# ─────────────────────────────────────────────\n# 11. TRAIN ONE FOLD\n# ─────────────────────────────────────────────\ndef train_one_fold(fold_id, train_data, val_data, n_folds=5):\n    print(f\"\\n  Fold {fold_id+1}/{n_folds} ──────────────────────────\")\n    print(f\"     Train: {len(train_data)} | Val: {len(val_data)}\")\n\n    train_loader = DataLoader(\n        PatchDataset(train_data, train_tf),\n        batch_size=BATCH_SIZE, shuffle=True,\n        num_workers=NUM_WORKERS, pin_memory=True\n    )\n\n    model = build_model()\n\n    cnt_tr = Counter(r[\"label\"] for r in train_data)\n    w = torch.tensor(\n        [1.0 / max(cnt_tr.get(c, 1), 1) for c in range(6)],\n        dtype=torch.float\n    )\n    w = (w / w.sum() * 6).to(DEVICE)\n    criterion = nn.CrossEntropyLoss(weight=w, label_smoothing=0.1)\n\n    base_model = model.module if hasattr(model, \"module\") else model\n\n    if FREEZE_UNI:\n        optimizer = optim.AdamW([\n            {\"params\": base_model.cnn_backbone.parameters(), \"lr\": LR_CNN},\n            {\"params\": [p for n, p in base_model.named_parameters()\n                        if not n.startswith(\"cnn_backbone\") and p.requires_grad],\n             \"lr\": LR_TR},\n        ], weight_decay=1e-4)\n    else:\n        optimizer = optim.AdamW([\n            {\"params\": base_model.cnn_backbone.parameters(), \"lr\": LR_CNN},\n            {\"params\": base_model.uni_backbone.parameters(), \"lr\": LR_UNI},\n            {\"params\": [p for n, p in base_model.named_parameters()\n                        if not n.startswith(\"cnn_backbone\")\n                        and not n.startswith(\"uni_backbone\")],\n             \"lr\": LR_TR},\n        ], weight_decay=1e-4)\n\n    scheduler = CosineAnnealingLR(optimizer, T_max=EPOCHS, eta_min=1e-7)\n\n    best_kappa = 0.0\n    no_improve = 0\n    best_ckpt  = f\"{SAVE_DIR}/fusion_nonfrozen_fold{fold_id+1}_best.pth\"\n\n    for epoch in range(1, EPOCHS + 1):\n        model.train()\n        t_loss, t_correct, t_total = 0.0, 0, 0\n        for imgs, labels in train_loader:\n            imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n            optimizer.zero_grad()\n            out  = model(imgs)\n            loss = criterion(out, labels)\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n            optimizer.step()\n            t_loss    += loss.item() * imgs.size(0)\n            t_correct += (out.argmax(1) == labels).sum().item()\n            t_total   += imgs.size(0)\n        scheduler.step()\n\n        wsi_kappa, n_wsis_val = wsi_level_qwk(model, val_data, DEVICE)\n        is_best = wsi_kappa > best_kappa\n\n        print(f\"     Ep[{epoch:02d}/{EPOCHS}] \"\n              f\"loss={t_loss/t_total:.4f} acc={t_correct/t_total:.3f} | \"\n              f\"WSI_QWK={wsi_kappa:.4f} ({n_wsis_val} WSIs)\"\n              + (\" *\" if is_best else \"\"))\n\n        if is_best:\n            best_kappa = wsi_kappa\n            no_improve = 0\n            torch.save({\n                \"epoch\":        epoch,\n                \"model_state_dict\": model.state_dict() if not hasattr(model, \"module\")\n                                    else model.module.state_dict(),\n                \"kappa\":        wsi_kappa,\n                \"freeze_uni\":   FREEZE_UNI,\n                \"n_wsis\":       len(all_wsis),\n                \"lr_uni\":       LR_UNI,\n            }, best_ckpt)\n        else:\n            no_improve += 1\n            if no_improve >= PATIENCE:\n                print(f\"     Early stopping epoch {epoch}\")\n                break\n\n    print(f\"  OK Fold {fold_id+1} best WSI_QWK: {best_kappa:.4f}\")\n    push_to_kaggle(best_ckpt,\n        f\"Fusion NonFrozen Fold{fold_id+1} QWK={best_kappa:.4f}\")\n    return best_kappa, best_ckpt\n\n# ─────────────────────────────────────────────\n# 12. 5-FOLD CV\n# ─────────────────────────────────────────────\nprint(f\"\\n{'='*60}\")\nprint(f\"HybridFusion NON-FREEZÉ — {len(all_wsis)} WSIs — 5-Fold CV\")\nprint(f\"ResNet-101 + UNI (LR_UNI={LR_UNI})\")\nprint(f\"Fusion : Attention patch-par-patch\")\nprint(f\"Baseline : Fusion freezé mean=0.6422\")\nprint(f\"{'='*60}\")\n\nwsi_ids    = [r[\"wsi_id\"] for r in all_records]\nlabels_arr = [r[\"label\"]  for r in all_records]\nindices    = list(range(len(all_records)))\n\nsgkf = StratifiedGroupKFold(n_splits=N_FOLDS, shuffle=True, random_state=SEED)\n\nfold_kappas       = []\nbest_overall      = -1.0\nbest_ckpt_overall = None\n\nfor fold_id, (train_idx, val_idx) in enumerate(\n        sgkf.split(indices, labels_arr, groups=wsi_ids)):\n\n    train_data = [all_records[i] for i in train_idx]\n    val_data   = [all_records[i] for i in val_idx]\n\n    tw = set(all_records[i][\"wsi_id\"] for i in train_idx)\n    vw = set(all_records[i][\"wsi_id\"] for i in val_idx)\n    assert len(tw & vw) == 0, f\"DATA LEAKAGE fold {fold_id+1}!\"\n\n    kappa, ckpt_path = train_one_fold(fold_id, train_data, val_data, N_FOLDS)\n    fold_kappas.append(kappa)\n\n    if kappa > best_overall:\n        best_overall      = kappa\n        best_ckpt_overall = ckpt_path\n\n# ─────────────────────────────────────────────\n# 13. RESUME FINAL\n# ─────────────────────────────────────────────\nmean_kappa = np.mean(fold_kappas)\nstd_kappa  = np.std(fold_kappas)\n\nprint(f\"\\n{'='*60}\")\nprint(f\"OK HybridFusion NON-FREEZÉ — 5-Fold CV Results\")\nprint(f\"{'='*60}\")\nfor i, k in enumerate(fold_kappas):\n    print(f\"   Fold {i+1}: QWK = {k:.4f}\")\nprint(f\"{'─'*40}\")\nprint(f\"   Mean QWK non-freezé : {mean_kappa:.4f} +/- {std_kappa:.4f}\")\nprint(f\"   Best QWK non-freezé : {best_overall:.4f}\")\nprint(f\"   Fusion freezé Stage0: mean=0.6422, best=0.7794\")\nprint(f\"   UNI non-freezé S0   : mean=0.6705, best=0.7964\")\nprint(f\"   WSIs                : {len(all_wsis)}\")\nprint(f\"   SOTA [1]            : 0.934\")\nprint(f\"{'='*60}\")\n\nresults = {\n    \"architecture\":  \"HybridFusion_ResNet101_UNI_NonFrozen\",\n    \"freeze_uni\":    FREEZE_UNI,\n    \"n_wsis\":        len(all_wsis),\n    \"n_patches\":     len(all_records),\n    \"lr_uni\":        LR_UNI,\n    \"fold_kappas\":   fold_kappas,\n    \"mean_kappa\":    mean_kappa,\n    \"std_kappa\":     std_kappa,\n    \"best_kappa\":    best_overall,\n    \"timestamp\":     datetime.datetime.now().isoformat(),\n}\nwith open(RESULTS_JSON, \"w\") as f:\n    json.dump(results, f, indent=2)\n\nif best_ckpt_overall:\n    stage_best = f\"{SAVE_DIR}/fusion_nonfrozen_best_qwk{best_overall:.4f}.pth\"\n    shutil.copy(best_ckpt_overall, stage_best)\n    print(f\"\\nOK Best checkpoint: {stage_best}\")\n    print(f\"\\nComparaison finale:\")\n    print(f\"  Fusion freezé  : mean=0.6422, best=0.7794\")\n    print(f\"  Fusion non-frz : mean={mean_kappa:.4f}, best={best_overall:.4f}\")\n    if mean_kappa > 0.6422:\n        print(f\"  → Non-freezé MEILLEUR ! +{mean_kappa-0.6422:.4f}\")\n    else:\n        print(f\"  → Freezé meilleur. Diff={mean_kappa-0.6422:.4f}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\n=============================================================\nStage 6 — HybridCNNTransformer — Incremental Learning\nFine-tune depuis Stage 3 (best QWK=0.7845)\nDonnées cumulatives :\n  Shard 0 : patches_400              (399  WSIs,  4,319 patches)\n  Shard 1 : panda-patches-run1       (1000 WSIs, 14,842 patches)\n  Shard 2 : panda-patches-run2       (999  WSIs, 14,867 patches)\n  Shard 3 : panda-patches-run3       (999  WSIs, 14,835 patches)\n  Shard 4 : panda-patches-run4       (1000 WSIs, 14,808 patches)\n  Shard 5 : panda-patches-run5-parts (1000 WSIs, 14,833 patches)\n  Shard 6 : panda-patches-run6-parts (1000 WSIs, 14,726 patches)\nTotal     : ~5,994 WSIs, ~93,230 patches\nTarget    : WSI_QWK > 0.7845\n=============================================================\nAttache dans Kaggle :\n  khouloudbejaoui20/dataset\n  khouloudbejaoui20/panda-patches-run1\n  khouloudbejaoui20/panda-patches-run2\n  khouloudbejaoui20/panda-patches-run3\n  khouloudbejaoui20/panda-patches-run4\n  bejaouikhouloud/panda-patches-run5-parts\n  bejaouikhouloud/panda-patches-run6-parts\n  khouloudbejaoui20/best-models ← PREV_CKPT\nGPU : T4x2\n\"\"\"\n\nimport os, glob, json, random, shutil, subprocess, datetime, zipfile\nimport numpy as np\nfrom PIL import Image\nfrom collections import Counter\nfrom sklearn.model_selection import StratifiedGroupKFold\nfrom sklearn.metrics import cohen_kappa_score\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as T\nimport torchvision.models as models\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\n\n# ─────────────────────────────────────────────\n# 1. CONFIG\n# ─────────────────────────────────────────────\nSHARD_ID    = 6\nSTART_FOLD  = 0\nDONE_KAPPAS = []\n\n# PREV_CKPT = stage3_fold5_best.pth (meilleur checkpoint)\nPREV_CKPT = \"/kaggle/working/stage3_fold5_best.pth\"\n\nN_FOLDS     = 5\nSEED        = 42\nBATCH_SIZE  = 32\nNUM_WORKERS = 2\nEPOCHS      = 30\nPATIENCE    = 7\nLR_CNN      = 2e-5\nLR_TR       = 4e-5\n\nrandom.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)\ntorch.backends.cudnn.deterministic = True\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Device    : {DEVICE}\")\nprint(f\"Stage     : {SHARD_ID} (fine-tune depuis Stage 3 QWK=0.7845)\")\nprint(f\"EPOCHS={EPOCHS} | PATIENCE={PATIENCE} | LR_CNN={LR_CNN} | LR_TR={LR_TR}\")\n\n# ─────────────────────────────────────────────\n# 2. TELECHARGE PREV_CKPT\n# ─────────────────────────────────────────────\nif not os.path.exists(PREV_CKPT):\n    print(\"Téléchargement stage3_fold5_best.pth...\")\n    subprocess.run([\n        \"kaggle\", \"datasets\", \"download\",\n        \"khouloudbejaoui20/best-models\",\n        \"--file\", \"stage3_fold5_best.pth\",\n        \"-p\", \"/kaggle/working\", \"--unzip\"\n    ], capture_output=True)\n\nassert os.path.exists(PREV_CKPT), f\"PREV_CKPT non trouvé: {PREV_CKPT}\"\nprint(f\"OK Stage 3 checkpoint (QWK=0.7845)\")\n\n# ─────────────────────────────────────────────\n# 3. CHEMINS PATCHES\n# ─────────────────────────────────────────────\nBASE      = \"/kaggle/input/datasets/khouloudbejaoui20\"\nRUN5_BASE = \"/kaggle/input/datasets/bejaouikhouloud/panda-patches-run5-parts\"\nRUN6_BASE = \"/kaggle/input/datasets/bejaouikhouloud/panda-patches-run6-parts\"\nRUN6_OUT  = \"/kaggle/working/patches_panda_run6\"\n\nSHARD_PATHS = {\n    0: f\"{BASE}/dataset/patches_400-20260630T212405Z-3-001/patches_400\",\n    1: f\"{BASE}/panda-patches-run1/kaggle/working/patches_panda_new\",\n    2: f\"{BASE}/panda-patches-run2/kaggle/working/patches_panda_run2\",\n    3: f\"{BASE}/panda-patches-run3/kaggle/working/patches_panda_run3_correct\",\n    4: f\"{BASE}/panda-patches-run4\",\n}\n\nSAVE_DIR     = \"/kaggle/working\"\nKAGGLE_DS    = \"bejaouikhouloud/stage6-results\"\nRESULTS_JSON = f\"{SAVE_DIR}/stage6_cv_results.json\"\n\n# ─────────────────────────────────────────────\n# 4. EXTRAIT RUN6 depuis ZIPs\n# ─────────────────────────────────────────────\ndef extract_run6():\n    os.makedirs(RUN6_OUT, exist_ok=True)\n    existing = glob.glob(f\"{RUN6_OUT}/*.png\")\n    if len(existing) > 10000:\n        print(f\"  Run6 déjà extrait: {len(existing)} patches ✅\")\n        return\n\n    print(f\"  Extraction Run6 depuis ZIPs...\")\n    zip_files = sorted(glob.glob(f\"{RUN6_BASE}/*.zip\"))\n    print(f\"  ZIPs trouvés: {len(zip_files)}\")\n\n    for zf in zip_files:\n        print(f\"  Extracting {os.path.basename(zf)}...\")\n        try:\n            with zipfile.ZipFile(zf, 'r') as z:\n                z.extractall(RUN6_OUT)\n        except Exception as e:\n            print(f\"  ERREUR: {e}\")\n\n    files = glob.glob(f\"{RUN6_OUT}/*.png\")\n    # Cherche aussi dans sous-dossiers\n    if not files:\n        files = glob.glob(f\"{RUN6_OUT}/**/*.png\", recursive=True)\n        # Déplace vers RUN6_OUT\n        for f in files:\n            shutil.move(f, RUN6_OUT)\n\n    print(f\"  Run6 extrait: {len(glob.glob(f'{RUN6_OUT}/*.png'))} patches ✅\")\n\n# ─────────────────────────────────────────────\n# 5. LOAD PATCHES\n# ─────────────────────────────────────────────\nprint(f\"\\n── Verification ────────────────────────\")\nfor shard_id, path in SHARD_PATHS.items():\n    if os.path.exists(path):\n        files = glob.glob(f\"{path}/*.png\")\n        if not files:\n            files = glob.glob(f\"{path}/**/*.png\", recursive=True)\n        print(f\"  OK Shard {shard_id}: {len(files)} patches\")\n    else:\n        print(f\"  MANQUANT Shard {shard_id}: {path}\")\n\nprint(f\"\\n── Loading patches ──────────────────────\")\n\ndef load_patches(folder):\n    files = glob.glob(f\"{folder}/*.png\")\n    if not files:\n        files = glob.glob(f\"{folder}/**/*.png\", recursive=True)\n    records = []\n    for fp in files:\n        fname = os.path.basename(fp)\n        try:\n            parts  = fname.split(\"_slide\")\n            label  = int(parts[0].replace(\"grade\", \"\"))\n            wsi_id = parts[1].split(\"_\")[0]\n            records.append({\"path\": fp, \"wsi_id\": wsi_id, \"label\": label})\n        except:\n            continue\n    return records\n\nall_records = []\nfor shard_id, path in SHARD_PATHS.items():\n    if path and os.path.exists(path):\n        r = load_patches(path)\n        all_records.extend(r)\n        print(f\"  Shard {shard_id}: {len(r)} patches\")\n    else:\n        print(f\"  Shard {shard_id}: NON TROUVE\")\n\n# Shard 5 — 15 parties\nrun5_total = 0\nfor part in range(1, 16):\n    path = f\"{RUN5_BASE}/run5_part{part}/kaggle/working/patches_panda_run5\"\n    if os.path.exists(path):\n        r = load_patches(path)\n        all_records.extend(r)\n        run5_total += len(r)\nprint(f\"  Shard 5 (15 parts): {run5_total} patches\")\n\n# Shard 6 — extrait depuis ZIPs\nextract_run6()\nr6 = load_patches(RUN6_OUT)\nall_records.extend(r6)\nprint(f\"  Shard 6: {len(r6)} patches\")\n\n# Deduplicate\nseen = set()\nunique_records = []\nfor r in all_records:\n    fname = os.path.basename(r[\"path\"])\n    if fname not in seen:\n        seen.add(fname)\n        unique_records.append(r)\n\nall_wsis = list(set(r[\"wsi_id\"] for r in unique_records))\nprint(f\"\\n  Total unique : {len(unique_records)} patches\")\nprint(f\"  Total WSIs  : {len(all_wsis)}\")\ncnt = Counter(r[\"label\"] for r in unique_records)\nfor c in range(6):\n    print(f\"    ISUP {c}: {cnt.get(c, 0)}\")\n\nassert len(all_wsis) > 5500, f\"Trop peu de WSIs ({len(all_wsis)})\"\nprint(f\"\\n  OK {len(all_wsis)} WSIs confirmes\")\n\n# ─────────────────────────────────────────────\n# 6. TRANSFORMS\n# ─────────────────────────────────────────────\nMEAN = [0.485, 0.456, 0.406]\nSTD  = [0.229, 0.224, 0.225]\n\ntrain_tf = T.Compose([\n    T.Resize((224, 224)),\n    T.RandomHorizontalFlip(),\n    T.RandomVerticalFlip(),\n    T.RandomRotation(90),\n    T.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.2, hue=0.1),\n    T.RandomApply([T.GaussianBlur(5)], p=0.3),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n    T.RandomApply([T.RandomErasing(p=0.5, scale=(0.02, 0.2))], p=0.2),\n])\n\nval_tf = T.Compose([\n    T.Resize((224, 224)),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n])\n\nclass PatchDataset(Dataset):\n    def __init__(self, records, transform):\n        self.records   = records\n        self.transform = transform\n    def __len__(self): return len(self.records)\n    def __getitem__(self, idx):\n        r = self.records[idx]\n        try:\n            img = Image.open(r[\"path\"]).convert(\"RGB\")\n        except:\n            img = Image.new(\"RGB\", (224, 224), (255, 255, 255))\n        return self.transform(img), int(r[\"label\"])\n\n# ─────────────────────────────────────────────\n# 7. MODELE\n# ─────────────────────────────────────────────\nclass HybridCNNTransformer(nn.Module):\n    def __init__(self, num_classes=6, d_model=512, nhead=8,\n                 num_layers=2, dim_feedforward=2048, dropout=0.1):\n        super().__init__()\n        backbone = models.resnet50(weights=models.ResNet50_Weights.DEFAULT)\n        self.cnn_backbone = nn.Sequential(*list(backbone.children())[:-2])\n        self.pool         = nn.AdaptiveAvgPool2d((1, 1))\n        self.proj         = nn.Sequential(\n            nn.Identity(), nn.Identity(),\n            nn.Linear(2048, d_model), nn.LayerNorm(d_model)\n        )\n        self.cls_token = nn.Parameter(torch.zeros(1, 1, d_model))\n        self.pos_embed = nn.Parameter(torch.zeros(1, 2, d_model))\n        encoder_layer  = nn.TransformerEncoderLayer(\n            d_model=d_model, nhead=nhead,\n            dim_feedforward=dim_feedforward,\n            dropout=dropout, batch_first=True\n        )\n        self.transformer = nn.TransformerEncoder(\n            encoder_layer, num_layers=num_layers,\n            norm=nn.LayerNorm(d_model)\n        )\n        self.classifier = nn.Sequential(\n            nn.Linear(d_model, 256), nn.ReLU(),\n            nn.Dropout(dropout), nn.Linear(256, num_classes)\n        )\n        nn.init.trunc_normal_(self.cls_token, std=0.02)\n        nn.init.trunc_normal_(self.pos_embed, std=0.02)\n\n    def forward(self, x):\n        f   = self.pool(self.cnn_backbone(x)).flatten(1)\n        f   = self.proj(f).unsqueeze(1)\n        cls = self.cls_token.expand(x.size(0), -1, -1)\n        seq = torch.cat([cls, f], dim=1) + self.pos_embed\n        return self.classifier(self.transformer(seq)[:, 0])\n\ndef build_model():\n    model  = HybridCNNTransformer().to(DEVICE)\n    n_gpus = torch.cuda.device_count()\n    if n_gpus > 1:\n        model = nn.DataParallel(model)\n        print(f\"  OK DataParallel: {n_gpus} GPUs\")\n    n_param = sum(p.numel() for p in model.parameters())\n    assert n_param == 30_998_342, f\"Mismatch: {n_param}\"\n    print(f\"  OK Parameters: {n_param:,}\")\n    ckpt  = torch.load(PREV_CKPT, map_location=DEVICE, weights_only=False)\n    state = ckpt[\"model_state_dict\"]\n    if n_gpus > 1 and not list(state.keys())[0].startswith(\"module.\"):\n        state = {\"module.\" + k: v for k, v in state.items()}\n    elif n_gpus <= 1 and list(state.keys())[0].startswith(\"module.\"):\n        state = {k[7:]: v for k, v in state.items()}\n    model.load_state_dict(state)\n    print(f\"  OK Loaded Stage 3 (QWK={ckpt.get('kappa', '?')})\")\n    return model\n\n# ─────────────────────────────────────────────\n# 8. WSI-LEVEL QWK\n# ─────────────────────────────────────────────\ndef wsi_level_qwk(model, val_data, device):\n    model.eval()\n    wsi_preds  = {}\n    wsi_labels = {}\n    all_preds  = []\n    all_labels = []\n    loader = DataLoader(\n        PatchDataset(val_data, val_tf),\n        batch_size=BATCH_SIZE, shuffle=False,\n        num_workers=NUM_WORKERS\n    )\n    with torch.no_grad():\n        for imgs, labels in loader:\n            imgs  = imgs.to(device)\n            preds = model(imgs).argmax(1).cpu().numpy()\n            all_preds.extend(preds)\n            all_labels.extend(labels.numpy())\n    for i, r in enumerate(val_data):\n        wid = r[\"wsi_id\"]\n        if wid not in wsi_preds:\n            wsi_preds[wid]  = []\n            wsi_labels[wid] = r[\"label\"]\n        wsi_preds[wid].append(all_preds[i])\n    wf = []; wl = []\n    for wid in wsi_preds:\n        votes = wsi_preds[wid]\n        wf.append(max(set(votes), key=votes.count))\n        wl.append(wsi_labels[wid])\n    wsi_qwk = cohen_kappa_score(wl, wf, weights=\"quadratic\")\n    return wsi_qwk, len(wf)\n\n# ─────────────────────────────────────────────\n# 9. PUSH KAGGLE\n# ─────────────────────────────────────────────\ndef push_to_kaggle(ckpt_path, message):\n    try:\n        tmp = f\"{SAVE_DIR}/_push_stage6\"\n        os.makedirs(tmp, exist_ok=True)\n        for c in glob.glob(f\"{SAVE_DIR}/stage{SHARD_ID}_fold*.pth\"):\n            shutil.copy(c, tmp)\n        shutil.copy(ckpt_path, tmp)\n        if os.path.exists(RESULTS_JSON):\n            shutil.copy(RESULTS_JSON, tmp)\n        json.dump({\n            \"title\": \"Stage6 Results Bejaoui 2026\",\n            \"id\": KAGGLE_DS,\n            \"licenses\": [{\"name\": \"CC0-1.0\"}]\n        }, open(f\"{tmp}/dataset-metadata.json\", \"w\"))\n        r = subprocess.run(\n            [\"kaggle\", \"datasets\", \"create\",\n             \"-p\", tmp, \"--dir-mode\", \"skip\"],\n            capture_output=True, text=True, timeout=600\n        )\n        if r.returncode != 0:\n            r = subprocess.run(\n                [\"kaggle\", \"datasets\", \"version\",\n                 \"-p\", tmp, \"-m\", message, \"--dir-mode\", \"skip\"],\n                capture_output=True, text=True, timeout=600\n            )\n        print(f\"   {'OK!' if r.returncode==0 else 'ERREUR: '+r.stderr[:100]}\")\n    except Exception as e:\n        print(f\"   ERREUR: {e}\")\n\n# ─────────────────────────────────────────────\n# 10. TRAIN ONE FOLD\n# ─────────────────────────────────────────────\ndef train_one_fold(fold_id, train_data, val_data, n_folds=5):\n    print(f\"\\n  Fold {fold_id+1}/{n_folds} ──────────────────────────\")\n    print(f\"     Train: {len(train_data)} | Val: {len(val_data)}\")\n\n    train_loader = DataLoader(\n        PatchDataset(train_data, train_tf),\n        batch_size=BATCH_SIZE, shuffle=True,\n        num_workers=NUM_WORKERS, pin_memory=True\n    )\n\n    model = build_model()\n\n    cnt_tr = Counter(r[\"label\"] for r in train_data)\n    w = torch.tensor(\n        [1.0 / max(cnt_tr.get(c, 1), 1) for c in range(6)],\n        dtype=torch.float\n    )\n    w = (w / w.sum() * 6).to(DEVICE)\n    criterion = nn.CrossEntropyLoss(weight=w, label_smoothing=0.1)\n\n    base_model = model.module if hasattr(model, \"module\") else model\n    optimizer  = optim.AdamW([\n        {\"params\": base_model.cnn_backbone.parameters(), \"lr\": LR_CNN},\n        {\"params\": [p for n, p in base_model.named_parameters()\n                    if not n.startswith(\"cnn_backbone\")], \"lr\": LR_TR},\n    ], weight_decay=1e-4)\n\n    scheduler = CosineAnnealingLR(optimizer, T_max=EPOCHS, eta_min=1e-7)\n\n    best_kappa = 0.0\n    no_improve = 0\n    best_ckpt  = f\"{SAVE_DIR}/stage{SHARD_ID}_fold{fold_id+1}_best.pth\"\n\n    for epoch in range(1, EPOCHS + 1):\n        model.train()\n        t_loss, t_correct, t_total = 0.0, 0, 0\n        for imgs, labels in train_loader:\n            imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n            optimizer.zero_grad()\n            out  = model(imgs)\n            loss = criterion(out, labels)\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n            optimizer.step()\n            t_loss    += loss.item() * imgs.size(0)\n            t_correct += (out.argmax(1) == labels).sum().item()\n            t_total   += imgs.size(0)\n        scheduler.step()\n\n        wsi_kappa, n_wsis_val = wsi_level_qwk(model, val_data, DEVICE)\n        is_best = wsi_kappa > best_kappa\n\n        print(f\"     Ep[{epoch:02d}/{EPOCHS}] \"\n              f\"loss={t_loss/t_total:.4f} acc={t_correct/t_total:.3f} | \"\n              f\"WSI_QWK={wsi_kappa:.4f} ({n_wsis_val} WSIs)\"\n              + (\" *\" if is_best else \"\"))\n\n        if is_best:\n            best_kappa = wsi_kappa\n            no_improve = 0\n            torch.save({\n                \"epoch\":            epoch,\n                \"model_state_dict\": model.state_dict() if not hasattr(model, \"module\")\n                                    else model.module.state_dict(),\n                \"kappa\":            wsi_kappa,\n                \"shard_id\":         SHARD_ID,\n                \"fold_id\":          fold_id,\n                \"n_wsis\":           len(all_wsis),\n            }, best_ckpt)\n        else:\n            no_improve += 1\n            if no_improve >= PATIENCE:\n                print(f\"     Early stopping epoch {epoch}\")\n                break\n\n    print(f\"  OK Fold {fold_id+1} best WSI_QWK: {best_kappa:.4f}\")\n    push_to_kaggle(best_ckpt,\n        f\"Stage{SHARD_ID} Fold{fold_id+1} QWK={best_kappa:.4f}\")\n    return best_kappa, best_ckpt\n\n# ─────────────────────────────────────────────\n# 11. 5-FOLD CV\n# ─────────────────────────────────────────────\nprint(f\"\\n{'='*60}\")\nprint(f\"Stage {SHARD_ID} — {len(all_wsis)} WSIs — 5-Fold CV\")\nprint(f\"Fine-tune depuis Stage 3 (QWK=0.7845)\")\nprint(f\"Target : WSI_QWK > 0.7845\")\nprint(f\"{'='*60}\")\n\nwsi_ids    = [r[\"wsi_id\"] for r in unique_records]\nlabels_arr = [r[\"label\"]  for r in unique_records]\nindices    = list(range(len(unique_records)))\n\nsgkf = StratifiedGroupKFold(n_splits=N_FOLDS, shuffle=True, random_state=SEED)\n\nfold_kappas       = DONE_KAPPAS.copy()\nbest_overall      = max(DONE_KAPPAS) if DONE_KAPPAS else -1.0\nbest_ckpt_overall = None\n\nfor fold_id, (train_idx, val_idx) in enumerate(\n        sgkf.split(indices, labels_arr, groups=wsi_ids)):\n\n    if fold_id < START_FOLD:\n        continue\n\n    train_data = [unique_records[i] for i in train_idx]\n    val_data   = [unique_records[i] for i in val_idx]\n\n    tw = set(unique_records[i][\"wsi_id\"] for i in train_idx)\n    vw = set(unique_records[i][\"wsi_id\"] for i in val_idx)\n    assert len(tw & vw) == 0, f\"DATA LEAKAGE fold {fold_id+1}!\"\n\n    kappa, ckpt_path = train_one_fold(fold_id, train_data, val_data, N_FOLDS)\n    fold_kappas.append(kappa)\n\n    if kappa > best_overall:\n        best_overall      = kappa\n        best_ckpt_overall = ckpt_path\n\n# ─────────────────────────────────────────────\n# 12. RESUME FINAL\n# ─────────────────────────────────────────────\nmean_kappa = np.mean(fold_kappas)\nstd_kappa  = np.std(fold_kappas)\n\nprint(f\"\\n{'='*60}\")\nprint(f\"OK Stage {SHARD_ID} — 5-Fold CV Results\")\nprint(f\"{'='*60}\")\nfor i, k in enumerate(fold_kappas):\n    print(f\"   Fold {i+1}: QWK = {k:.4f}\")\nprint(f\"{'─'*40}\")\nprint(f\"   Mean QWK : {mean_kappa:.4f} +/- {std_kappa:.4f}\")\nprint(f\"   Best QWK : {best_overall:.4f}\")\nprint(f\"   Stage 3  : mean=0.7280, best=0.7845\")\nprint(f\"   Gain     : {mean_kappa - 0.7280:+.4f}\")\nprint(f\"   WSIs     : {len(all_wsis)}\")\nprint(f\"   SOTA [1] : 0.934\")\nprint(f\"{'='*60}\")\n\nresults = {\n    \"stage\":       SHARD_ID,\n    \"n_wsis\":      len(all_wsis),\n    \"n_patches\":   len(unique_records),\n    \"fold_kappas\": fold_kappas,\n    \"mean_kappa\":  mean_kappa,\n    \"std_kappa\":   std_kappa,\n    \"best_kappa\":  best_overall,\n    \"prev_ckpt\":   \"stage3_fold5_best.pth\",\n    \"timestamp\":   datetime.datetime.now().isoformat(),\n}\nwith open(RESULTS_JSON, \"w\") as f:\n    json.dump(results, f, indent=2)\n\nif best_ckpt_overall:\n    stage_best = f\"{SAVE_DIR}/stage{SHARD_ID}_best_qwk{best_overall:.4f}.pth\"\n    shutil.copy(best_ckpt_overall, stage_best)\n    print(f\"\\nOK Best checkpoint: {stage_best}\")\n    print(f\"\\n→ Pour Stage 7 :\")\n    print(f\"   SHARD_ID  = 7\")\n    print(f\"   PREV_CKPT = '{stage_best}'\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# EXTRACTION D'IMAGES WSI RÉELLES — PANDA Dataset\n# Génère une figure publication-quality depuis de vraies WSI\n# ============================================================\nimport os, glob, random\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport matplotlib.patches as mpatches\nimport matplotlib.gridspec as gridspec\nfrom matplotlib.patches import Rectangle\nimport tifffile\nfrom PIL import Image\n\n# ── Paramètres ────────────────────────────────────────────────\n# Sur Colab + Drive\n#TIFF_DIR    = \"/content/drive/MyDrive/Memoire_IHC/dataset/panda/tiff_samples\"\n#MASK_DIR    = \"/content/drive/MyDrive/Memoire_IHC/dataset/panda/masks\"\n#RESULTS_DIR = \"/content/drive/MyDrive/Memoire_IHC/results\"\n#os.makedirs(RESULTS_DIR, exist_ok=True)\n\n# Sur Kaggle (décommentez si Kaggle) :\nTIFF_DIR = \"/kaggle/input/prostate-cancer-grade-assessment/train_images\"\nMASK_DIR = \"/kaggle/input/prostate-cancer-grade-assessment/train_label_masks\"\n\nPATCH_SIZE  = 224\nISUP_COLORS = {\n    0: '#4CAF50',   # vert  — bénin\n    1: '#8BC34A',   # vert clair\n    2: '#FFC107',   # orange clair\n    3: '#FF9800',   # orange\n    4: '#F44336',   # rouge\n    5: '#9C27B0',   # violet\n}\nISUP_NAMES = {\n    0: 'ISUP 0 — Bénin',\n    1: 'ISUP 1 — Gleason 3+3',\n    2: 'ISUP 2 — Gleason 3+4',\n    3: 'ISUP 3 — Gleason 4+3',\n    4: 'ISUP 4 — Gleason 4+4',\n    5: 'ISUP 5 — Gleason 5+5',\n}\n\n# ── Fonction de lecture WSI ───────────────────────────────────\ndef read_wsi_thumbnail(tiff_path, max_size=1024):\n    \"\"\"\n    Lire une WSI TIFF multi-résolution et retourner\n    la vignette (niveau le plus bas = plus petite résolution)\n    \"\"\"\n    try:\n        with tifffile.TiffFile(tiff_path) as tif:\n            # Prendre le niveau le plus bas (thumbnail)\n            n_levels = len(tif.series[0].levels)\n            # Niveau intermédiaire pour bonne qualité\n            level = max(0, n_levels - 3)\n            img = tif.series[0].levels[level].asarray()\n            # Convertir en RGB si nécessaire\n            if img.ndim == 2:\n                img = np.stack([img]*3, axis=-1)\n            elif img.shape[0] == 3:\n                img = np.transpose(img, (1, 2, 0))\n            # Redimensionner si trop grand\n            h, w = img.shape[:2]\n            if max(h, w) > max_size:\n                scale = max_size / max(h, w)\n                new_h, new_w = int(h*scale), int(w*scale)\n                img = np.array(Image.fromarray(img).resize(\n                    (new_w, new_h), Image.LANCZOS))\n            return img\n    except Exception as e:\n        print(f\"Erreur lecture {tiff_path}: {e}\")\n        return None\n\ndef read_wsi_region(tiff_path, level=1, region=None):\n    \"\"\"\n    Lire une région spécifique de la WSI à un niveau donné\n    region = (x, y, width, height) en coordonnées du niveau 0\n    \"\"\"\n    try:\n        with tifffile.TiffFile(tif_path) as tif:\n            n_levels = len(tif.series[0].levels)\n            level = min(level, n_levels - 1)\n            img = tif.series[0].levels[level].asarray()\n            if img.ndim == 2:\n                img = np.stack([img]*3, axis=-1)\n            elif img.shape[0] in [3, 4]:\n                img = np.transpose(img, (1, 2, 0))\n            if img.shape[2] == 4:\n                img = img[:, :, :3]\n            return img\n    except Exception as e:\n        print(f\"Erreur: {e}\")\n        return None\n\ndef extract_patches_from_wsi(tiff_path, patch_size=224,\n                              n_patches=6, level=1):\n    \"\"\"\n    Extraire n patches aléatoires d'une WSI\n    (évite les zones vides/background)\n    \"\"\"\n    try:\n        with tifffile.TiffFile(tiff_path) as tif:\n            n_levels = len(tif.series[0].levels)\n            level = min(level, n_levels - 1)\n            img = tif.series[0].levels[level].asarray()\n            if img.ndim == 2:\n                img = np.stack([img]*3, axis=-1)\n            elif img.shape[0] in [3, 4]:\n                img = np.transpose(img, (1, 2, 0))\n            if img.shape[2] == 4:\n                img = img[:, :, :3]\n\n        h, w = img.shape[:2]\n        patches = []\n        positions = []\n        attempts = 0\n\n        while len(patches) < n_patches and attempts < 200:\n            attempts += 1\n            x = random.randint(0, max(0, w - patch_size))\n            y = random.randint(0, max(0, h - patch_size))\n            patch = img[y:y+patch_size, x:x+patch_size]\n\n            if patch.shape[0] < patch_size or patch.shape[1] < patch_size:\n                continue\n\n            # Filtrer patches trop blancs (background)\n            mean_val = patch.mean()\n            std_val  = patch.std()\n            if mean_val > 230 or std_val < 10:\n                continue\n\n            patches.append(patch)\n            positions.append((x, y))\n\n        return img, patches, positions\n\n    except Exception as e:\n        print(f\"Erreur extraction: {e}\")\n        return None, [], []\n\n# ── Trouver les WSI disponibles ───────────────────────────────\nprint(\"Recherche des WSI disponibles...\")\ntiff_files = sorted(glob.glob(os.path.join(TIFF_DIR, \"*.tiff\")))\nprint(f\"WSI trouvées : {len(tiff_files)}\")\n\nif len(tiff_files) == 0:\n    print(\"❌ Aucune WSI trouvée !\")\n    print(\"   Vérifiez le chemin TIFF_DIR\")\nelse:\n    # Afficher quelques exemples\n    for f in tiff_files[:5]:\n        size = os.path.getsize(f) / 1024 / 1024\n        print(f\"  {os.path.basename(f)} — {size:.1f} MB\")\n\n# ── Lire les labels ISUP depuis le CSV PANDA ─────────────────\nimport pandas as pd\n\nCSV_PATH = \"/content/drive/MyDrive/Memoire_IHC/dataset/panda/train.csv\"\n# Sur Kaggle :\n# CSV_PATH = \"/kaggle/input/prostate-cancer-grade-assessment/train.csv\"\n\nif os.path.exists(CSV_PATH):\n    df_labels = pd.read_csv(CSV_PATH)\n    df_labels = df_labels.set_index('image_id')\n    print(f\"\\nLabels chargés : {len(df_labels)} WSI\")\n    print(df_labels['isup_grade'].value_counts().sort_index())\nelse:\n    print(\"⚠️ CSV labels non trouvé — labels non disponibles\")\n    df_labels = None\n\n# ── FIGURE 1 : Vue d'ensemble WSI + patches ───────────────────\ndef generate_wsi_overview_figure(tiff_files, n_wsi=3,\n                                  n_patches_per_wsi=4):\n    \"\"\"\n    Figure principale : WSI thumbnail + patches extraits\n    \"\"\"\n    # Sélectionner des WSI de grades différents si possible\n    selected_wsi = []\n    if df_labels is not None:\n        for grade in [0, 2, 4]:\n            wsis_grade = [\n                f for f in tiff_files\n                if os.path.splitext(os.path.basename(f))[0]\n                in df_labels.index\n                and df_labels.loc[\n                    os.path.splitext(os.path.basename(f))[0],\n                    'isup_grade'] == grade\n            ]\n            if wsis_grade:\n                selected_wsi.append(random.choice(wsis_grade))\n    if len(selected_wsi) < n_wsi:\n        remaining = [f for f in tiff_files\n                     if f not in selected_wsi]\n        selected_wsi.extend(\n            random.sample(remaining,\n                          min(n_wsi - len(selected_wsi),\n                              len(remaining))))\n\n    fig = plt.figure(figsize=(18, 14), dpi=200,\n                     facecolor='white')\n    fig.patch.set_facecolor('white')\n\n    gs_main = gridspec.GridSpec(\n        len(selected_wsi), 1, figure=fig,\n        hspace=0.35)\n\n    for row_idx, tiff_path in enumerate(selected_wsi):\n        wsi_id = os.path.splitext(\n            os.path.basename(tiff_path))[0]\n\n        # Label ISUP\n        isup = df_labels.loc[wsi_id, 'isup_grade'] \\\n            if df_labels is not None and wsi_id in df_labels.index \\\n            else '?'\n        color = ISUP_COLORS.get(isup, '#888888')\n        label = ISUP_NAMES.get(isup, f'ISUP {isup}')\n\n        # Extraire thumbnail + patches\n        print(f\"Lecture {wsi_id} (ISUP {isup})...\")\n        img, patches, positions = extract_patches_from_wsi(\n            tiff_path, patch_size=PATCH_SIZE,\n            n_patches=n_patches_per_wsi, level=1)\n\n        if img is None:\n            continue\n\n        # Layout : 1 thumbnail + n patches\n        n_cols = 1 + n_patches_per_wsi\n        gs_row = gridspec.GridSpecFromSubplotSpec(\n            1, n_cols, subplot_spec=gs_main[row_idx],\n            wspace=0.08, width_ratios=[3] + [1]*n_patches_per_wsi)\n\n        # ── Thumbnail WSI ──────────────────────────────────────\n        ax_wsi = fig.add_subplot(gs_row[0])\n        h, w = img.shape[:2]\n        scale_w = min(800, w) / w\n        thumb = np.array(Image.fromarray(img).resize(\n            (int(w*scale_w), int(h*scale_w)), Image.LANCZOS))\n        ax_wsi.imshow(thumb)\n        ax_wsi.set_title(\n            f'{wsi_id[:16]}...\\n{label}',\n            fontsize=10, fontweight='bold',\n            color=color, pad=6)\n        ax_wsi.axis('off')\n\n        # Dessiner les positions des patches extraits\n        if positions:\n            for px, py in positions:\n                scale = scale_w\n                rect = Rectangle(\n                    (px*scale, py*scale),\n                    PATCH_SIZE*scale, PATCH_SIZE*scale,\n                    linewidth=1.5,\n                    edgecolor=color,\n                    facecolor='none', alpha=0.7)\n                ax_wsi.add_patch(rect)\n\n        # Barre colorée ISUP\n        ax_wsi.add_patch(Rectangle(\n            (0, 0), 1, 0.04,\n            transform=ax_wsi.transAxes,\n            facecolor=color, alpha=0.8, clip_on=False))\n\n        # ── Patches extraits ───────────────────────────────────\n        for p_idx, patch in enumerate(patches[:n_patches_per_wsi]):\n            ax_p = fig.add_subplot(gs_row[1 + p_idx])\n            ax_p.imshow(patch)\n            ax_p.set_title(\n                f'Patch {p_idx+1}\\n{PATCH_SIZE}×{PATCH_SIZE}',\n                fontsize=8, color='#444444', pad=4)\n            ax_p.axis('off')\n            # Bordure couleur ISUP\n            for spine in ax_p.spines.values():\n                spine.set_edgecolor(color)\n                spine.set_linewidth(2)\n\n    # ── Titre global ───────────────────────────────────────────\n    fig.suptitle(\n        'PANDA Dataset — Whole Slide Images (WSI) et patches extraits\\n'\n        'Biopsies prostatiques colorées H&E  ·  Grossissement 20×  '\n        '·  Format TIFF multi-résolution',\n        fontsize=13, fontweight='bold', y=0.98, color='#1E3A5F')\n\n    # ── Légende ISUP ───────────────────────────────────────────\n    legend_patches = [\n        mpatches.Patch(color=ISUP_COLORS[i],\n                       label=ISUP_NAMES[i])\n        for i in range(6)\n    ]\n    fig.legend(handles=legend_patches,\n               loc='lower center', ncol=3,\n               fontsize=9, framealpha=0.9,\n               bbox_to_anchor=(0.5, 0.0))\n\n    plt.tight_layout(rect=[0, 0.06, 1, 0.97])\n\n    out_path = os.path.join(\n        RESULTS_DIR, 'fig_wsi_real_examples.png')\n    fig.savefig(out_path, dpi=200, bbox_inches='tight',\n                facecolor='white')\n    plt.show()\n    print(f\"\\n✅ Sauvegardée : {out_path}\")\n    return fig\n\n# ── FIGURE 2 : Grille de patches par grade ISUP ───────────────\ndef generate_patches_grid_figure(patch_dirs=None):\n    \"\"\"\n    Grille 6×6 : 6 exemples par grade ISUP (0-5)\n    Depuis les patches déjà extraits (patches_phase2 ou patches_400)\n    \"\"\"\n    # Chemins des patches déjà extraits\n    PHASE2_DIR   = \"/content/drive/MyDrive/Memoire_IHC/dataset/panda/patches_phase2_clean\"\n    PHASE400_DIR = \"/content/drive/MyDrive/Memoire_IHC/dataset/panda/patches_400\"\n\n    all_patches = (\n        sorted(glob.glob(os.path.join(PHASE2_DIR, \"*.png\"))) +\n        sorted(glob.glob(os.path.join(PHASE400_DIR, \"*.png\")))\n    )\n    print(f\"Total patches disponibles : {len(all_patches)}\")\n\n    # Grouper par grade\n    import re\n    per_grade = {i: [] for i in range(6)}\n    for f in all_patches:\n        m = re.search(r\"grade(\\d+)_\", os.path.basename(f))\n        if m:\n            per_grade[int(m.group(1))].append(f)\n\n    for g, files in per_grade.items():\n        print(f\"  ISUP {g} : {len(files)} patches\")\n\n    # Figure\n    N_COLS = 6   # patches par grade\n    N_ROWS = 6   # grades (0-5)\n\n    fig, axes = plt.subplots(\n        N_ROWS, N_COLS,\n        figsize=(N_COLS * 2.2, N_ROWS * 2.4),\n        dpi=200, facecolor='white')\n\n    fig.suptitle(\n        'Exemples de patches extraits par grade ISUP\\n'\n        'Dataset PANDA  ·  224×224 pixels  ·  H&E staining',\n        fontsize=14, fontweight='bold',\n        y=0.98, color='#1E3A5F')\n\n    for grade in range(6):\n        files = per_grade[grade]\n        color = ISUP_COLORS[grade]\n        label = ISUP_NAMES[grade]\n\n        # Sélectionner N_COLS patches aléatoires\n        selected = random.sample(\n            files, min(N_COLS, len(files))) \\\n            if files else []\n\n        for col in range(N_COLS):\n            ax = axes[grade][col]\n\n            if col < len(selected):\n                img = np.array(\n                    Image.open(selected[col]).convert('RGB'))\n                ax.imshow(img)\n            else:\n                ax.set_facecolor('#F5F5F5')\n                ax.text(0.5, 0.5, 'N/A',\n                        ha='center', va='center',\n                        transform=ax.transAxes,\n                        color='#AAAAAA', fontsize=10)\n\n            ax.set_xticks([]); ax.set_yticks([])\n\n            # Bordure colorée par grade\n            for spine in ax.spines.values():\n                spine.set_edgecolor(color)\n                spine.set_linewidth(2.0)\n\n            # Label grade (première colonne seulement)\n            if col == 0:\n                ax.set_ylabel(\n                    label, fontsize=9,\n                    fontweight='bold', color=color,\n                    rotation=90, labelpad=6)\n\n            # Numéro patch (première ligne seulement)\n            if grade == 0:\n                ax.set_title(f'Ex. {col+1}',\n                             fontsize=9, color='#555555',\n                             pad=4)\n\n    plt.tight_layout(rect=[0, 0, 1, 0.96])\n\n    out_path = os.path.join(\n        RESULTS_DIR, 'fig_patches_grid_by_isup.png')\n    fig.savefig(out_path, dpi=200, bbox_inches='tight',\n                facecolor='white')\n    plt.show()\n    print(f\"\\n✅ Sauvegardée : {out_path}\")\n    return fig\n\n# ── FIGURE 3 : Zoom multi-échelle sur une WSI ─────────────────\ndef generate_multiscale_figure(tiff_path):\n    \"\"\"\n    Zoom progressif : WSI entière → région → patch 224×224\n    \"\"\"\n    wsi_id = os.path.splitext(\n        os.path.basename(tiff_path))[0]\n    isup = df_labels.loc[wsi_id, 'isup_grade'] \\\n        if df_labels is not None and wsi_id in df_labels.index \\\n        else '?'\n\n    print(f\"Génération figure multi-échelle : {wsi_id}\")\n\n    try:\n        with tifffile.TiffFile(tiff_path) as tif:\n            n_levels = len(tif.series[0].levels)\n            print(f\"Niveaux disponibles : {n_levels}\")\n\n            # Niveau 0 = full resolution (trop grand)\n            # Niveau -3 = thumbnail\n            levels_to_show = []\n            for lv in [n_levels-1, n_levels-2,\n                       max(0, n_levels-3)]:\n                img = tif.series[0].levels[lv].asarray()\n                if img.ndim == 2:\n                    img = np.stack([img]*3, axis=-1)\n                elif img.shape[0] in [3, 4]:\n                    img = np.transpose(img, (1, 2, 0))\n                if img.shape[2] == 4:\n                    img = img[:, :, :3]\n                levels_to_show.append((lv, img))\n\n    except Exception as e:\n        print(f\"Erreur: {e}\")\n        return\n\n    color = ISUP_COLORS.get(isup, '#888888')\n    label = ISUP_NAMES.get(isup, f'ISUP {isup}')\n\n    fig, axes = plt.subplots(\n        1, 4, figsize=(18, 5), dpi=200,\n        facecolor='white')\n\n    fig.suptitle(\n        f'WSI multi-échelle — {wsi_id[:20]}  |  {label}\\n'\n        f'PANDA Dataset  ·  Biopsie prostatique  ·  H&E',\n        fontsize=12, fontweight='bold',\n        y=1.02, color='#1E3A5F')\n\n    zoom_labels = [\n        ('Vue globale\\n(WSI entière)', '~1:50'),\n        ('Région tissulaire\\n(zoom ×4)', '~1:12'),\n        ('Zone d\\'intérêt\\n(zoom ×16)', '~1:3'),\n        (f'Patch 224×224 px\\n(entrée modèle)', '1:1'),\n    ]\n\n    for i, (ax, (zoom_label, scale_label)) in \\\n            enumerate(zip(axes, zoom_labels)):\n\n        if i < len(levels_to_show):\n            lv, img = levels_to_show[\n                min(i, len(levels_to_show)-1)]\n            h, w = img.shape[:2]\n\n            if i == 0:\n                # Vue globale — image entière\n                display = img\n            elif i == 1:\n                # Zoom sur région centrale\n                cx, cy = w//2, h//2\n                sz = min(w, h) // 2\n                display = img[\n                    max(0, cy-sz//2):cy+sz//2,\n                    max(0, cx-sz//2):cx+sz//2]\n            elif i == 2:\n                # Zoom encore plus fort\n                cx, cy = w//2, h//2\n                sz = min(w, h) // 4\n                display = img[\n                    max(0, cy-sz//2):cy+sz//2,\n                    max(0, cx-sz//2):cx+sz//2]\n            else:\n                # Patch 224×224\n                cx, cy = w//2, h//2\n                sz = 112\n                patch = img[\n                    max(0, cy-sz):cy+sz,\n                    max(0, cx-sz):cx+sz]\n                display = np.array(\n                    Image.fromarray(patch).resize(\n                        (224, 224), Image.LANCZOS))\n\n            ax.imshow(display)\n\n        ax.set_title(f'{zoom_label}\\n{scale_label}',\n                     fontsize=9, color='#333333', pad=6)\n        ax.axis('off')\n\n        # Bordure\n        for spine in ax.spines.values():\n            spine.set_edgecolor(color)\n            spine.set_linewidth(2)\n\n        # Flèche vers zoom suivant\n        if i < 3:\n            ax.annotate(\n                '', xy=(1.05, 0.5),\n                xytext=(0.95, 0.5),\n                xycoords='axes fraction',\n                textcoords='axes fraction',\n                arrowprops=dict(\n                    arrowstyle='->', color=color,\n                    lw=2.0))\n\n    plt.tight_layout()\n\n    out_path = os.path.join(\n        RESULTS_DIR,\n        f'fig_multiscale_{wsi_id[:12]}.png')\n    fig.savefig(out_path, dpi=200,\n                bbox_inches='tight', facecolor='white')\n    plt.show()\n    print(f\"\\n✅ Sauvegardée : {out_path}\")\n\n# ── LANCER LES FIGURES ────────────────────────────────────────\nprint(\"\\n\" + \"=\"*55)\nprint(\"GÉNÉRATION DES FIGURES WSI\")\nprint(\"=\"*55)\n\nif len(tiff_files) > 0:\n    # Figure 1 : Overview WSI + patches\n    print(\"\\n[1/3] Figure overview WSI + patches...\")\n    generate_wsi_overview_figure(\n        tiff_files, n_wsi=3, n_patches_per_wsi=4)\n\n    # Figure 2 : Grille patches par grade\n    print(\"\\n[2/3] Figure grille patches par ISUP...\")\n    generate_patches_grid_figure()\n\n    # Figure 3 : Multi-échelle sur la première WSI\n    print(\"\\n[3/3] Figure multi-échelle...\")\n    generate_multiscale_figure(tiff_files[0])\n\n    print(\"\\n✅ Toutes les figures générées dans results/\")\nelse:\n    print(\"❌ Pas de WSI disponibles — \"\n          \"vérifiez TIFF_DIR\")\n    print(\"   Génération figure patches uniquement...\")\n    generate_patches_grid_figure()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, glob\n\nDRIVE_BASE = \"/content/drive/MyDrive/Memoire_IHC/dataset/panda\"\n\nprint(\"=== Contenu dossier PANDA ===\\n\")\nif os.path.exists(DRIVE_BASE):\n    for item in os.listdir(DRIVE_BASE):\n        full = os.path.join(DRIVE_BASE, item)\n        if os.path.isdir(full):\n            files = os.listdir(full)\n            print(f\"📁 {item}/  → {len(files)} fichiers\")\n            if files:\n                print(f\"   ex: {files[0]}\")\n        else:\n            size = os.path.getsize(full)/1024/1024\n            print(f\"📄 {item}  ({size:.1f} MB)\")\nelse:\n    print(\"❌ Dossier PANDA non trouvé !\")\n    print(\"\\nRecherche dans tout Memoire_IHC...\")\n    for root, dirs, files in os.walk(\n            \"/content/drive/MyDrive/Memoire_IHC\"):\n        if len(files) > 0:\n            exts = set(os.path.splitext(f)[1]\n                       for f in files)\n            print(f\"  {root} → {len(files)} fichiers {exts}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, glob\n\nprint(\"=== Dataset PANDA disponible ===\\n\")\nfor root, dirs, files in os.walk(\"/kaggle/input\"):\n    if len(files) > 0:\n        exts = set(os.path.splitext(f)[1] for f in files)\n        size = sum(os.path.getsize(os.path.join(root, f))\n                   for f in files) / 1024 / 1024\n        print(f\"📁 {root}\")\n        print(f\"   → {len(files)} fichiers {exts} | {size:.0f} MB\")\n        print(f\"   → ex: {files[0]}\")\n        print()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, glob, re\nfrom collections import Counter\n\n# ── Compter tous les patches disponibles ──────────────────────\nprint(\"=== Inventaire complet des patches ===\\n\")\n\n# Chercher tous les PNG\nall_png = glob.glob(\n    \"/kaggle/input/**/*.png\", recursive=True)\n\nprint(f\"Total PNG trouvés : {len(all_png)}\")\n\n# Grouper par dossier source\nsources = Counter()\nfor f in all_png:\n    # Identifier la source\n    if 'run5' in f:\n        sources['panda-patches-run5'] += 1\n    elif 'patches_phase2' in f:\n        sources['patches_phase2_clean'] += 1\n    elif 'patches_400' in f:\n        sources['patches_400'] += 1\n    else:\n        sources['other'] += 1\n\nprint(\"\\nPar source :\")\nfor k, v in sources.most_common():\n    print(f\"  {k} : {v} patches\")\n\n# Distribution par grade\ndef get_grade(f):\n    m = re.search(r\"grade(\\d+)_\", os.path.basename(f))\n    return int(m.group(1)) if m else -1\n\ngrades = Counter(get_grade(f) for f in all_png)\nprint(\"\\nDistribution par grade ISUP :\")\nfor g in range(6):\n    print(f\"  ISUP {g} : {grades[g]:,} patches\")\nprint(f\"  Total  : {sum(grades.values()):,} patches\")\n\n# WSI uniques\nwsi_ids = set()\nfor f in all_png:\n    m = re.search(r\"slide([0-9a-fA-F]+)_\", os.path.basename(f))\n    if m:\n        wsi_ids.add(m.group(1))\nprint(f\"\\nWSI uniques : {len(wsi_ids)}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, glob\nfrom collections import Counter\n\nall_png = glob.glob(\"/kaggle/input/**/*.png\", recursive=True)\n\n# Trouver les dossiers \"other\"\nother_dirs = Counter()\nfor f in all_png:\n    if 'run5' not in f and \\\n       'patches_phase2' not in f and \\\n       'patches_400' not in f:\n        # Prendre le dossier parent\n        parent = os.path.basename(\n            os.path.dirname(f))\n        other_dirs[parent] += 1\n\nprint(\"Dossiers 'other' :\")\nfor k, v in other_dirs.most_common(10):\n    print(f\"  {k} : {v} patches\")\n\n# Exemple de nom de fichier other\nother_files = [f for f in all_png\n               if 'run5' not in f and\n               'patches_phase2' not in f and\n               'patches_400' not in f]\nprint(f\"\\nExemples :\")\nfor f in other_files[:5]:\n    print(f\"  {os.path.basename(f)}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# EXTRACTION IMAGE RÉELLE WSI depuis PANDA TIFF\n# Génère une figure publication-quality\n# ============================================================\nimport os, glob, random, re\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport matplotlib.gridspec as gridspec\nimport matplotlib.patches as mpatches\nfrom matplotlib.patches import Rectangle\nfrom PIL import Image\nimport tifffile\nimport pandas as pd\n\n# ── Chemins ───────────────────────────────────────────────────\n# Sur Kaggle avec dataset PANDA ajouté :\nTIFF_DIR = \"/kaggle/input/prostate-cancer-grade-assessment/train_images\"\nCSV_PATH = \"/kaggle/input/prostate-cancer-grade-assessment/train.csv\"\nRESULTS_DIR = \"/kaggle/working\"\n\n# ── Charger les labels ────────────────────────────────────────\ndf = pd.read_csv(CSV_PATH)\ndf = df.set_index('image_id')\nprint(f\"Labels chargés : {len(df)} WSI\")\nprint(df['isup_grade'].value_counts().sort_index())\n\nISUP_COLORS = {\n    0: '#4CAF50', 1: '#8BC34A', 2: '#FFC107',\n    3: '#FF9800', 4: '#F44336', 5: '#9C27B0'\n}\nISUP_NAMES = {\n    0: 'ISUP 0 — Benign',\n    1: 'ISUP 1 — Gleason 3+3',\n    2: 'ISUP 2 — Gleason 3+4',\n    3: 'ISUP 3 — Gleason 4+3',\n    4: 'ISUP 4 — Gleason 4+4',\n    5: 'ISUP 5 — Gleason 5+5',\n}\n\n# ── Fonction extraction WSI ───────────────────────────────────\ndef read_wsi_thumbnail(tiff_path, max_size=800):\n    \"\"\"Lire la vignette d'une WSI TIFF multi-résolution\"\"\"\n    with tifffile.TiffFile(tiff_path) as tif:\n        n_levels = len(tif.series[0].levels)\n        print(f\"  Niveaux disponibles : {n_levels}\")\n\n        # Prendre un niveau intermédiaire (pas trop grand)\n        level = max(0, n_levels - 2)\n        img = tif.series[0].levels[level].asarray()\n\n        # Convertir en RGB\n        if img.ndim == 2:\n            img = np.stack([img]*3, axis=-1)\n        elif img.shape[0] in [3, 4]:\n            img = np.transpose(img, (1, 2, 0))\n        if img.shape[-1] == 4:\n            img = img[:, :, :3]\n\n        print(f\"  Taille niveau {level}: {img.shape}\")\n\n        # Redimensionner si nécessaire\n        h, w = img.shape[:2]\n        if max(h, w) > max_size:\n            scale = max_size / max(h, w)\n            new_h = int(h * scale)\n            new_w = int(w * scale)\n            img = np.array(Image.fromarray(img).resize(\n                (new_w, new_h), Image.LANCZOS))\n        return img\n\ndef extract_tissue_patches(tiff_path, n=6,\n                            patch_size=224, level=1):\n    \"\"\"Extraire des patches de tissu (évite le fond blanc)\"\"\"\n    with tifffile.TiffFile(tiff_path) as tif:\n        n_levels = len(tif.series[0].levels)\n        level = min(level, n_levels - 1)\n        img = tif.series[0].levels[level].asarray()\n\n        if img.ndim == 2:\n            img = np.stack([img]*3, axis=-1)\n        elif img.shape[0] in [3, 4]:\n            img = np.transpose(img, (1, 2, 0))\n        if img.shape[-1] == 4:\n            img = img[:, :, :3]\n\n    h, w = img.shape[:2]\n    patches, positions = [], []\n    attempts = 0\n\n    while len(patches) < n and attempts < 500:\n        attempts += 1\n        x = random.randint(0, max(0, w - patch_size))\n        y = random.randint(0, max(0, h - patch_size))\n        patch = img[y:y+patch_size, x:x+patch_size]\n\n        if patch.shape[0] < patch_size or \\\n           patch.shape[1] < patch_size:\n            continue\n\n        # Éviter fond blanc et zones vides\n        mean = patch.mean()\n        std  = patch.std()\n        if mean > 235 or std < 8:\n            continue\n\n        patches.append(patch)\n        positions.append((x, y,\n                          img.shape[1], img.shape[0]))\n\n    return patches, positions, img\n\n# ── Sélectionner une WSI par grade ────────────────────────────\ndef get_wsi_by_grade(tiff_dir, df_labels, grade):\n    \"\"\"Trouver une WSI d'un grade ISUP spécifique\"\"\"\n    tiffs = glob.glob(os.path.join(tiff_dir, \"*.tiff\"))\n    candidates = [\n        f for f in tiffs\n        if os.path.splitext(os.path.basename(f))[0]\n        in df_labels.index\n        and df_labels.loc[\n            os.path.splitext(os.path.basename(f))[0],\n            'isup_grade'] == grade\n    ]\n    if not candidates:\n        return None\n    return random.choice(candidates)\n\n# ══════════════════════════════════════════════════════════════\n# FIGURE PRINCIPALE — 3 WSI avec leurs patches\n# (une bénigne, une intermédiaire, une agressive)\n# ══════════════════════════════════════════════════════════════\ngrades_to_show = [0, 2, 5]   # bénin, G3+4, G5+5\nN_PATCHES = 5\n\nfig = plt.figure(figsize=(20, 14), dpi=200,\n                 facecolor='white')\nfig.suptitle(\n    'PANDA Dataset — Whole Slide Images (WSI) réelles\\n'\n    'Biopsies prostatiques · H&E staining · '\n    'Format TIFF multi-résolution',\n    fontsize=14, fontweight='bold',\n    y=0.99, color='#1E3A5F')\n\ngs = gridspec.GridSpec(\n    len(grades_to_show), 1,\n    figure=fig, hspace=0.4)\n\nfor row, grade in enumerate(grades_to_show):\n    wsi_path = get_wsi_by_grade(TIFF_DIR, df, grade)\n    if wsi_path is None:\n        print(f\"❌ Aucune WSI grade {grade} trouvée\")\n        continue\n\n    wsi_id = os.path.splitext(\n        os.path.basename(wsi_path))[0]\n    color  = ISUP_COLORS[grade]\n    label  = ISUP_NAMES[grade]\n    size_mb = os.path.getsize(wsi_path) / 1024 / 1024\n\n    print(f\"\\nTraitement {wsi_id} ({label})...\")\n    print(f\"  Taille fichier : {size_mb:.1f} MB\")\n\n    # Lire thumbnail + patches\n    thumb = read_wsi_thumbnail(wsi_path, max_size=700)\n    patches, positions, full_img = \\\n        extract_tissue_patches(\n            wsi_path, n=N_PATCHES,\n            patch_size=224, level=1)\n\n    print(f\"  Thumbnail : {thumb.shape}\")\n    print(f\"  Patches extraits : {len(patches)}\")\n\n    # Layout : thumbnail (large) + patches\n    gs_row = gridspec.GridSpecFromSubplotSpec(\n        1, 1 + N_PATCHES,\n        subplot_spec=gs[row],\n        wspace=0.06,\n        width_ratios=[3.5] + [1]*N_PATCHES)\n\n    # ── Thumbnail WSI ──────────────────────────────────────────\n    ax_wsi = fig.add_subplot(gs_row[0])\n    ax_wsi.imshow(thumb)\n    ax_wsi.axis('off')\n\n    # Titre avec infos WSI\n    h_orig = full_img.shape[0]\n    w_orig = full_img.shape[1]\n    ax_wsi.set_title(\n        f'{label}\\n'\n        f'ID: {wsi_id[:20]}  ·  '\n        f'Résolution: {w_orig}×{h_orig} px  ·  '\n        f'{size_mb:.0f} MB',\n        fontsize=9, color=color,\n        fontweight='bold', pad=8)\n\n    # Dessiner les rectangles des patches\n    if positions and thumb is not None:\n        th = thumb.shape[0]\n        tw = thumb.shape[1]\n        scale_x = tw / w_orig\n        scale_y = th / h_orig\n        for px, py, orig_w, orig_h in positions:\n            rect = Rectangle(\n                (px * scale_x, py * scale_y),\n                224 * scale_x, 224 * scale_y,\n                linewidth=1.8,\n                edgecolor=color,\n                facecolor=color,\n                alpha=0.15)\n            ax_wsi.add_patch(rect)\n            rect2 = Rectangle(\n                (px * scale_x, py * scale_y),\n                224 * scale_x, 224 * scale_y,\n                linewidth=1.8,\n                edgecolor=color,\n                facecolor='none')\n            ax_wsi.add_patch(rect2)\n\n    # Barre colorée ISUP en bas\n    ax_wsi.add_patch(Rectangle(\n        (0, 0), 1, 0.035,\n        transform=ax_wsi.transAxes,\n        facecolor=color, alpha=0.85,\n        clip_on=False))\n\n    # ── Patches extraits ──────────────────────────────────────\n    for p_idx in range(N_PATCHES):\n        ax_p = fig.add_subplot(gs_row[1 + p_idx])\n\n        if p_idx < len(patches):\n            ax_p.imshow(patches[p_idx])\n        else:\n            ax_p.set_facecolor('#EEEEEE')\n            ax_p.text(0.5, 0.5, 'N/A',\n                      ha='center', va='center',\n                      transform=ax_p.transAxes,\n                      color='#AAAAAA', fontsize=9)\n\n        if p_idx == 0:\n            ax_p.set_ylabel(\n                '224×224\\npx', fontsize=7,\n                color='#666666', rotation=90)\n\n        ax_p.set_xticks([]); ax_p.set_yticks([])\n        ax_p.set_title(\n            f'Patch {p_idx+1}',\n            fontsize=8, color='#555555', pad=3)\n\n        for spine in ax_p.spines.values():\n            spine.set_edgecolor(color)\n            spine.set_linewidth(2)\n\n# ── Légende ───────────────────────────────────────────────────\nlegend_patches = [\n    mpatches.Patch(\n        color=ISUP_COLORS[i],\n        label=ISUP_NAMES[i])\n    for i in range(6)\n]\nfig.legend(\n    handles=legend_patches,\n    loc='lower center', ncol=3,\n    fontsize=9, framealpha=0.9,\n    bbox_to_anchor=(0.5, 0.0))\n\nplt.tight_layout(rect=[0, 0.05, 1, 0.98])\n\nout = f\"{RESULTS_DIR}/fig_wsi_real_panda.png\"\nfig.savefig(out, dpi=200,\n            bbox_inches='tight', facecolor='white')\nplt.show()\nprint(f\"\\n✅ Sauvegardée : {out}\")\n\n# ══════════════════════════════════════════════════════════════\n# FIGURE 2 — Zoom multi-échelle sur une WSI\n# ══════════════════════════════════════════════════════════════\nprint(\"\\nGénération figure multi-échelle...\")\n\nwsi_path = get_wsi_by_grade(TIFF_DIR, df, grade=3)\nif wsi_path:\n    wsi_id = os.path.splitext(\n        os.path.basename(wsi_path))[0]\n    isup = df.loc[wsi_id, 'isup_grade']\n    color = ISUP_COLORS[isup]\n    label = ISUP_NAMES[isup]\n\n    fig2, axes = plt.subplots(\n        1, 4, figsize=(20, 5), dpi=200,\n        facecolor='white')\n    fig2.suptitle(\n        f'Multi-scale view — {wsi_id[:24]}\\n{label}',\n        fontsize=12, fontweight='bold',\n        y=1.02, color='#1E3A5F')\n\n    zoom_titles = [\n        'WSI entière\\n(thumbnail)',\n        'Région × 4\\n(tissu)',\n        'Zone × 16\\n(glandes)',\n        'Patch 224×224\\n(entrée modèle)'\n    ]\n\n    with tifffile.TiffFile(wsi_path) as tif:\n        n_levels = len(tif.series[0].levels)\n        imgs = []\n        for lv in range(n_levels):\n            img = tif.series[0].levels[lv].asarray()\n            if img.ndim == 2:\n                img = np.stack([img]*3, axis=-1)\n            elif img.shape[0] in [3, 4]:\n                img = np.transpose(img, (1, 2, 0))\n            if img.shape[-1] == 4:\n                img = img[:, :, :3]\n            imgs.append(img)\n            print(f\"  Level {lv}: {img.shape}\")\n\n    # 4 niveaux de zoom\n    level_indices = [\n        len(imgs)-1,         # thumbnail\n        max(0, len(imgs)-3), # région\n        max(0, len(imgs)-5), # zone\n        0                    # plein résolution (224px crop)\n    ]\n\n    for i, (ax, title) in enumerate(\n            zip(axes, zoom_titles)):\n        lv = min(level_indices[i], len(imgs)-1)\n        img = imgs[lv]\n        h, w = img.shape[:2]\n\n        if i == 0:\n            # Thumbnail entier\n            display = np.array(\n                Image.fromarray(img).resize(\n                    (600, int(600*h/w)),\n                    Image.LANCZOS))\n        elif i < 3:\n            # Crop centré\n            cx, cy = w//2, h//2\n            half = min(cx, cy) // (i+1)\n            y1 = max(0, cy-half)\n            y2 = min(h, cy+half)\n            x1 = max(0, cx-half)\n            x2 = min(w, cx+half)\n            display = img[y1:y2, x1:x2]\n        else:\n            # Patch 224×224\n            cx, cy = w//2, h//2\n            y1 = max(0, cy-112)\n            y2 = min(h, cy+112)\n            x1 = max(0, cx-112)\n            x2 = min(w, cx+112)\n            crop = img[y1:y2, x1:x2]\n            display = np.array(\n                Image.fromarray(crop).resize(\n                    (224, 224), Image.LANCZOS))\n\n        ax.imshow(display)\n        ax.set_title(title, fontsize=10,\n                     color='#333333', pad=6)\n        ax.axis('off')\n        for spine in ax.spines.values():\n            spine.set_edgecolor(color)\n            spine.set_linewidth(2)\n\n        # Flèche de zoom\n        if i < 3:\n            ax.annotate(\n                '', xy=(1.06, 0.5),\n                xytext=(0.97, 0.5),\n                xycoords='axes fraction',\n                arrowprops=dict(\n                    arrowstyle='->', color=color,\n                    lw=2.5))\n\n    plt.tight_layout()\n    out2 = f\"{RESULTS_DIR}/fig_wsi_multiscale.png\"\n    fig2.savefig(out2, dpi=200,\n                 bbox_inches='tight',\n                 facecolor='white')\n    plt.show()\n    print(f\"✅ Sauvegardée : {out2}\")\n\nprint(\"\\n🎉 Figures WSI réelles générées !\")\nprint(f\"  {RESULTS_DIR}/fig_wsi_real_panda.png\")\nprint(f\"  {RESULTS_DIR}/fig_wsi_multiscale.png\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, glob\n\n# Chercher les TIFF dans tous les inputs Kaggle\ntiffs = glob.glob(\"/kaggle/input/**/*.tiff\", recursive=True)\ntiffs += glob.glob(\"/kaggle/input/**/*.tif\", recursive=True)\n\nprint(f\"TIFF trouvés : {len(tiffs)}\")\nfor f in tiffs[:5]:\n    size = os.path.getsize(f) / 1024 / 1024\n    print(f\"  {os.path.basename(f)} — {size:.1f} MB\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import tifffile\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport matplotlib.patches as mpatches\nimport pandas as pd\nimport os, glob, random\n\n# ── Chemins ───────────────────────────────────────────────────\nTIFF_DIR = os.path.dirname(tiffs[0])\nCSV_PATH = \"/kaggle/input/prostate-cancer-grade-assessment/train.csv\"\nRESULTS  = \"/kaggle/working\"\n\n# ── Labels ISUP ───────────────────────────────────────────────\ndf = pd.read_csv(CSV_PATH).set_index('image_id')\nprint(f\"Labels : {len(df)} WSI\")\n\nISUP_COLORS = {\n    0:'#4CAF50', 1:'#8BC34A', 2:'#FFC107',\n    3:'#FF9800', 4:'#F44336', 5:'#9C27B0'\n}\nISUP_NAMES = {\n    0:'ISUP 0 — Benign',     1:'ISUP 1 — Gleason 3+3',\n    2:'ISUP 2 — Gleason 3+4', 3:'ISUP 3 — Gleason 4+3',\n    4:'ISUP 4 — Gleason 4+4', 5:'ISUP 5 — Gleason 5+5',\n}\n\n# ── Lire WSI ──────────────────────────────────────────────────\ndef read_wsi(tiff_path, max_size=800):\n    with tifffile.TiffFile(tiff_path) as tif:\n        n = len(tif.series[0].levels)\n        img = tif.series[0].levels[max(0,n-2)].asarray()\n    if img.ndim == 2:\n        img = np.stack([img]*3, axis=-1)\n    elif img.shape[0] in [3,4]:\n        img = np.transpose(img,(1,2,0))\n    if img.shape[-1] == 4:\n        img = img[:,:,:3]\n    h, w = img.shape[:2]\n    if max(h,w) > max_size:\n        s = max_size/max(h,w)\n        img = np.array(__import__('PIL').Image.fromarray(img)\n                       .resize((int(w*s),int(h*s)),3))\n    return img\n\n# ── Sélectionner 1 WSI par grade ISUP ─────────────────────────\nselected = {}\nfor grade in range(6):\n    candidates = [\n        f for f in tiffs\n        if os.path.splitext(os.path.basename(f))[0]\n        in df.index\n        and df.loc[os.path.splitext(\n            os.path.basename(f))[0],'isup_grade']==grade\n    ]\n    if candidates:\n        selected[grade] = random.choice(candidates)\n\nprint(f\"WSI sélectionnées : {len(selected)}\")\n\n# ── FIGURE : 6 WSI réelles (une par grade ISUP) ───────────────\nfig, axes = plt.subplots(\n    2, 3, figsize=(18, 12), dpi=200,\n    facecolor='white')\naxes = axes.flatten()\n\nfig.suptitle(\n    'PANDA Dataset — Whole Slide Images réelles\\n'\n    'Une WSI par grade ISUP · H&E staining · '\n    'Biopsies prostatiques',\n    fontsize=14, fontweight='bold',\n    y=1.01, color='#1E3A5F')\n\nfor grade in range(6):\n    ax = axes[grade]\n\n    if grade not in selected:\n        ax.set_facecolor('#F5F5F5')\n        ax.text(0.5, 0.5, f'ISUP {grade}\\nN/A',\n                ha='center', va='center',\n                transform=ax.transAxes,\n                fontsize=12, color='#AAAAAA')\n        ax.axis('off')\n        continue\n\n    tiff_path = selected[grade]\n    wsi_id = os.path.splitext(\n        os.path.basename(tiff_path))[0]\n    color = ISUP_COLORS[grade]\n    label = ISUP_NAMES[grade]\n    size_mb = os.path.getsize(tiff_path)/1024/1024\n\n    print(f\"Lecture ISUP {grade}: {wsi_id[:16]}...\")\n    img = read_wsi(tiff_path, max_size=700)\n\n    ax.imshow(img)\n    ax.axis('off')\n    ax.set_title(\n        f'{label}\\n'\n        f'{img.shape[1]}×{img.shape[0]} px  ·  '\n        f'{size_mb:.0f} MB',\n        fontsize=10, fontweight='bold',\n        color=color, pad=8)\n\n    # Barre colorée en bas\n    from matplotlib.patches import Rectangle\n    ax.add_patch(Rectangle(\n        (0,0), 1, 0.04,\n        transform=ax.transAxes,\n        facecolor=color, alpha=0.85,\n        clip_on=False))\n\n    # Cadre coloré\n    for spine in ax.spines.values():\n        spine.set_edgecolor(color)\n        spine.set_linewidth(3)\n        spine.set_visible(True)\n\nplt.tight_layout(rect=[0,0,1,0.98])\n\nout = f\"{RESULTS}/fig_wsi_6grades_real.png\"\nplt.savefig(out, dpi=200,\n            bbox_inches='tight', facecolor='white')\nplt.show()\nprint(f\"\\n✅ Sauvegardée : {out}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, glob\n\n# Chercher le CSV train.csv partout\ncsvs = glob.glob(\"/kaggle/input/**/*.csv\", recursive=True)\nprint(\"CSV trouvés :\")\nfor c in csvs:\n    size = os.path.getsize(c)/1024\n    print(f\"  {c}  ({size:.0f} KB)\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, glob, random\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom matplotlib.patches import Rectangle\nimport tifffile\nfrom PIL import Image\n\n# ── Étape 1 : Trouver les chemins ─────────────────────────────\nCSV_PATH = \"/kaggle/input/competitions/prostate-cancer-grade-assessment/train.csv\"\n\n# Trouver les TIFF\ntiffs = glob.glob(\"/kaggle/input/**/*.tiff\", recursive=True)\nprint(f\"✅ TIFF trouvés : {len(tiffs)}\")\nprint(f\"   Exemple     : {tiffs[0]}\")\n\n# Charger labels\ndf = pd.read_csv(CSV_PATH).set_index('image_id')\nprint(f\"✅ Labels chargés : {len(df)} WSI\")\nprint(f\"\\nDistribution ISUP :\")\nprint(df['isup_grade'].value_counts().sort_index())","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── Étape 2 : Sélectionner 1 WSI par grade ISUP ───────────────\nISUP_COLORS = {\n    0:'#4CAF50', 1:'#8BC34A', 2:'#FFC107',\n    3:'#FF9800', 4:'#F44336', 5:'#9C27B0'\n}\nISUP_NAMES = {\n    0:'ISUP 0 — Benign',\n    1:'ISUP 1 — Gleason 3+3',\n    2:'ISUP 2 — Gleason 3+4',\n    3:'ISUP 3 — Gleason 4+3',\n    4:'ISUP 4 — Gleason 4+4',\n    5:'ISUP 5 — Gleason 5+5',\n}\n\nrandom.seed(42)\nselected = {}\n\nfor grade in range(6):\n    candidates = [\n        f for f in tiffs\n        if os.path.splitext(os.path.basename(f))[0]\n        in df.index\n        and df.loc[\n            os.path.splitext(os.path.basename(f))[0],\n            'isup_grade'] == grade\n    ]\n    if candidates:\n        selected[grade] = random.choice(candidates)\n        wsi_id = os.path.splitext(\n            os.path.basename(selected[grade]))[0]\n        size = os.path.getsize(selected[grade])/1024/1024\n        print(f\"ISUP {grade} : {wsi_id[:20]}... \"\n              f\"({size:.0f} MB)\")\n    else:\n        print(f\"ISUP {grade} : ❌ non trouvé\")\n\nprint(f\"\\n✅ {len(selected)} WSI sélectionnées\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── Étape 3 : Tester la lecture d'une WSI ─────────────────────\ndef read_wsi(tiff_path, max_size=700):\n    with tifffile.TiffFile(tiff_path) as tif:\n        n_levels = len(tif.series[0].levels)\n        level = max(0, n_levels - 2)\n        img = tif.series[0].levels[level].asarray()\n        print(f\"  Niveaux : {n_levels} | \"\n              f\"Niveau : {level} | \"\n              f\"Shape : {img.shape}\")\n\n    if img.ndim == 2:\n        img = np.stack([img]*3, axis=-1)\n    elif img.shape[0] in [3, 4]:\n        img = np.transpose(img, (1, 2, 0))\n    if img.shape[-1] == 4:\n        img = img[:, :, :3]\n\n    h, w = img.shape[:2]\n    if max(h, w) > max_size:\n        scale = max_size / max(h, w)\n        img = np.array(Image.fromarray(img).resize(\n            (int(w*scale), int(h*scale)),\n            Image.LANCZOS))\n    return img\n\n# Test ISUP 0\nprint(\"Test lecture ISUP 0...\")\nimg_test = read_wsi(selected[0])\nprint(f\"✅ Image lue : {img_test.shape}\")\n\nplt.figure(figsize=(6, 6), dpi=100)\nplt.imshow(img_test)\nplt.title(f\"Test WSI — ISUP 0\\n\"\n          f\"{img_test.shape[1]}×{img_test.shape[0]} px\",\n          fontsize=11, fontweight='bold',\n          color='#4CAF50')\nplt.axis('off')\nplt.tight_layout()\nplt.show()\nprint(\"✅ OK — lancez Cellule 4 !\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── Méthode alternative — lecture TIFF robuste ────────────────\ndef read_wsi_robust(tiff_path, max_size=700):\n    \"\"\"\n    Lecture WSI avec plusieurs méthodes de fallback\n    \"\"\"\n    with tifffile.TiffFile(tiff_path) as tif:\n        n_levels = len(tif.series[0].levels)\n        print(f\"  Niveaux disponibles : {n_levels}\")\n\n        # Essayer chaque niveau du plus petit au plus grand\n        img = None\n        for level in range(n_levels-1, -1, -1):\n            try:\n                img = tif.series[0].levels[level].asarray()\n                print(f\"  ✅ Niveau {level} lu : {img.shape}\")\n                break\n            except Exception as e:\n                print(f\"  ⚠️  Niveau {level} échoué : {e}\")\n                continue\n\n        # Fallback → lire page par page\n        if img is None:\n            print(\"  Tentative lecture page par page...\")\n            try:\n                pages = tif.pages\n                img = pages[0].asarray()\n                print(f\"  ✅ Page 0 lue : {img.shape}\")\n            except Exception as e:\n                print(f\"  ❌ Page 0 échouée : {e}\")\n\n    if img is None:\n        return None\n\n    # Convertir en RGB\n    if img.ndim == 2:\n        img = np.stack([img]*3, axis=-1)\n    elif img.ndim == 3 and img.shape[0] in [3, 4]:\n        img = np.transpose(img, (1, 2, 0))\n    if img.ndim == 3 and img.shape[-1] == 4:\n        img = img[:, :, :3]\n\n    # Redimensionner\n    h, w = img.shape[:2]\n    if max(h, w) > max_size:\n        scale = max_size / max(h, w)\n        img = np.array(Image.fromarray(\n            img.astype(np.uint8)).resize(\n            (int(w*scale), int(h*scale)),\n            Image.LANCZOS))\n\n    return img\n\n# ── Tester sur les 6 WSI sélectionnées ───────────────────────\nprint(\"=== Test lecture des 6 WSI ===\\n\")\nresults = {}\n\nfor grade, tiff_path in selected.items():\n    wsi_id = os.path.splitext(\n        os.path.basename(tiff_path))[0]\n    print(f\"ISUP {grade} — {wsi_id[:20]}...\")\n    try:\n        img = read_wsi_robust(tiff_path, max_size=600)\n        if img is not None:\n            results[grade] = img\n            print(f\"  ✅ Shape finale : {img.shape}\\n\")\n        else:\n            print(f\"  ❌ Échec total\\n\")\n    except Exception as e:\n        print(f\"  ❌ Erreur : {e}\\n\")\n\nprint(f\"✅ WSI lues avec succès : \"\n      f\"{len(results)}/6\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── Étape 1 : Installer imagecodecs ───────────────────────────\n!pip install imagecodecs -q\nprint(\"✅ imagecodecs installé !\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import importlib, tifffile\nimportlib.reload(tifffile)\n\nimport os, glob, random\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nimport pandas as pd\n\nCSV_PATH = \"/kaggle/input/competitions/prostate-cancer-grade-assessment/train.csv\"\ndf = pd.read_csv(CSV_PATH).set_index('image_id')\ntiffs = glob.glob(\"/kaggle/input/**/*.tiff\", recursive=True)\n\nISUP_COLORS = {\n    0:'#4CAF50', 1:'#8BC34A', 2:'#FFC107',\n    3:'#FF9800', 4:'#F44336', 5:'#9C27B0'\n}\nISUP_NAMES = {\n    0:'ISUP 0 — Benign',      1:'ISUP 1 — Gleason 3+3',\n    2:'ISUP 2 — Gleason 3+4', 3:'ISUP 3 — Gleason 4+3',\n    4:'ISUP 4 — Gleason 4+4', 5:'ISUP 5 — Gleason 5+5',\n}\n\nrandom.seed(42)\nselected = {}\nfor grade in range(6):\n    candidates = [\n        f for f in tiffs\n        if os.path.splitext(os.path.basename(f))[0]\n        in df.index\n        and df.loc[\n            os.path.splitext(os.path.basename(f))[0],\n            'isup_grade'] == grade\n    ]\n    if candidates:\n        selected[grade] = random.choice(candidates)\n\ndef read_wsi(tiff_path, max_size=700):\n    with tifffile.TiffFile(tiff_path) as tif:\n        n_levels = len(tif.series[0].levels)\n        level = n_levels - 1\n        img = tif.series[0].levels[level].asarray()\n    if img.ndim == 2:\n        img = np.stack([img]*3, axis=-1)\n    elif img.shape[0] in [3, 4]:\n        img = np.transpose(img, (1, 2, 0))\n    if img.shape[-1] == 4:\n        img = img[:, :, :3]\n    h, w = img.shape[:2]\n    if max(h, w) > max_size:\n        scale = max_size / max(h, w)\n        img = np.array(Image.fromarray(\n            img.astype(np.uint8)).resize(\n            (int(w*scale), int(h*scale)),\n            Image.LANCZOS))\n    return img\n\n# Test ISUP 0\nprint(\"Test lecture ISUP 0...\")\nimg_test = read_wsi(selected[0])\nprint(f\"✅ Shape : {img_test.shape}\")\n\nplt.figure(figsize=(5,5), dpi=100)\nplt.imshow(img_test)\nplt.title(\"Test WSI ISUP 0\", fontsize=11,\n          fontweight='bold', color='#4CAF50')\nplt.axis('off')\nplt.tight_layout()\nplt.show()\nprint(\"✅ OK !\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# LECTURE WSI PANDA — version complète\n# ============================================================\nimport subprocess\nsubprocess.run([\"pip\", \"install\", \"imagecodecs\",\n                \"--quiet\"], check=True)\n\nimport os, glob, random\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom matplotlib.patches import Rectangle\nimport matplotlib.patches as mpatches\nimport pandas as pd\nfrom PIL import Image\nimport tifffile\n\n# ── Chemins ───────────────────────────────────────────────────\nCSV_PATH = \"/kaggle/input/competitions/prostate-cancer-grade-assessment/train.csv\"\nRESULTS  = \"/kaggle/working\"\n\ndf = pd.read_csv(CSV_PATH).set_index('image_id')\ntiffs = glob.glob(\"/kaggle/input/**/*.tiff\",\n                  recursive=True)\nprint(f\"✅ TIFF : {len(tiffs)} | Labels : {len(df)}\")\n\nISUP_COLORS = {\n    0:'#4CAF50', 1:'#8BC34A', 2:'#FFC107',\n    3:'#FF9800', 4:'#F44336', 5:'#9C27B0'\n}\nISUP_NAMES = {\n    0:'ISUP 0 — Benign',\n    1:'ISUP 1 — Gleason 3+3',\n    2:'ISUP 2 — Gleason 3+4',\n    3:'ISUP 3 — Gleason 4+3',\n    4:'ISUP 4 — Gleason 4+4',\n    5:'ISUP 5 — Gleason 5+5',\n}\n\n# ── Sélectionner 1 WSI par grade ──────────────────────────────\nrandom.seed(42)\nselected = {}\nfor grade in range(6):\n    candidates = [\n        f for f in tiffs\n        if os.path.splitext(os.path.basename(f))[0]\n        in df.index\n        and df.loc[\n            os.path.splitext(os.path.basename(f))[0],\n            'isup_grade'] == grade\n    ]\n    if candidates:\n        selected[grade] = random.choice(candidates)\n        print(f\"ISUP {grade} : \"\n              f\"{os.path.basename(selected[grade])[:20]}...\")\n\n# ── Lecture WSI ───────────────────────────────────────────────\ndef read_wsi(tiff_path, max_size=700):\n    with tifffile.TiffFile(tiff_path) as tif:\n        n = len(tif.series[0].levels)\n        # Prendre le niveau le plus petit\n        img = tif.series[0].levels[n-1].asarray()\n    if img.ndim == 2:\n        img = np.stack([img]*3, axis=-1)\n    elif img.shape[0] in [3, 4]:\n        img = np.transpose(img, (1, 2, 0))\n    if img.shape[-1] == 4:\n        img = img[:, :, :3]\n    h, w = img.shape[:2]\n    if max(h, w) > max_size:\n        s = max_size / max(h, w)\n        img = np.array(Image.fromarray(\n            img.astype(np.uint8)).resize(\n            (int(w*s), int(h*s)), Image.LANCZOS))\n    return img\n\n# ── Figure 6 WSI réelles ──────────────────────────────────────\nfig, axes = plt.subplots(\n    2, 3, figsize=(18, 13),\n    dpi=200, facecolor='white')\naxes = axes.flatten()\n\nfig.suptitle(\n    'PANDA Dataset — Whole Slide Images (WSI) réelles\\n'\n    'Une WSI par grade ISUP  ·  H&E staining  ·  '\n    'Biopsies prostatiques',\n    fontsize=13, fontweight='bold',\n    y=1.01, color='#1E3A5F')\n\nfor grade in range(6):\n    ax = axes[grade]\n    if grade not in selected:\n        ax.text(0.5, 0.5, f'ISUP {grade}\\nN/A',\n                ha='center', va='center',\n                transform=ax.transAxes,\n                fontsize=12, color='#AAAAAA')\n        ax.axis('off')\n        continue\n\n    tiff_path = selected[grade]\n    wsi_id  = os.path.splitext(\n        os.path.basename(tiff_path))[0]\n    color   = ISUP_COLORS[grade]\n    label   = ISUP_NAMES[grade]\n    size_mb = os.path.getsize(tiff_path)/1024/1024\n\n    print(f\"Lecture ISUP {grade}...\", end=' ')\n    try:\n        img = read_wsi(tiff_path, max_size=700)\n        print(f\"✅ {img.shape[1]}×{img.shape[0]}\")\n    except Exception as e:\n        print(f\"❌ {e}\")\n        ax.text(0.5, 0.5,\n                f'ISUP {grade}\\nErreur\\n{str(e)[:30]}',\n                ha='center', va='center',\n                transform=ax.transAxes,\n                fontsize=8, color='#F44336')\n        ax.axis('off')\n        continue\n\n    ax.imshow(img)\n    ax.axis('off')\n    ax.set_title(\n        f'{label}\\n'\n        f'{wsi_id[:18]}...\\n'\n        f'{img.shape[1]}×{img.shape[0]} px  ·  '\n        f'{size_mb:.0f} MB',\n        fontsize=9, fontweight='bold',\n        color=color, pad=8)\n    ax.add_patch(Rectangle(\n        (0, 0), 1, 0.04,\n        transform=ax.transAxes,\n        facecolor=color, alpha=0.9,\n        clip_on=False))\n    for spine in ax.spines.values():\n        spine.set_edgecolor(color)\n        spine.set_linewidth(3)\n        spine.set_visible(True)\n\nplt.tight_layout(rect=[0, 0, 1, 0.98])\nout = f\"{RESULTS}/fig_wsi_6grades_real.png\"\nplt.savefig(out, dpi=200,\n            bbox_inches='tight', facecolor='white')\nplt.show()\nprint(f\"\\n✅ Sauvegardée : {out}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\n=============================================================\nStage 3 — HybridUNITransformer — FROM SCRATCH\nUNI (MahmoodLab) foundation model + Transformer Encoder\nDonnées cumulatives :\n  Shard 0 : patches_400        (399  WSIs,  4,319 patches)\n  Shard 1 : panda-patches-run1 (1000 WSIs, 14,842 patches)\n  Shard 2 : panda-patches-run2 (999  WSIs, 14,867 patches)\n  Shard 3 : panda-patches-run3 (999  WSIs, 14,835 patches)\nTotal     : ~3,398 WSIs, ~48,863 patches\n\nComparaison :\n  ResNet-50 Stage 3 : mean=0.7280, best=0.7845\n  UNI Stage 3       : ???\n=============================================================\nAttache dans Kaggle :\n  khouloudbejaoui20/dataset\n  khouloudbejaoui20/panda-patches-run1\n  khouloudbejaoui20/panda-patches-run2\n  khouloudbejaoui20/panda-patches-run3\nGPU : T4×2\n\"\"\"\n\nimport os, glob, json, random, shutil, subprocess, datetime\nimport numpy as np\nfrom PIL import Image\nfrom collections import Counter\nfrom sklearn.model_selection import StratifiedGroupKFold\nfrom sklearn.metrics import cohen_kappa_score\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as T\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\n\n# ─────────────────────────────────────────────\n# 1. LOGIN HUGGINGFACE\n# ─────────────────────────────────────────────\nfrom huggingface_hub import login, whoami\nimport timm\n\nHF_TOKEN = \"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\"\nlogin(token=HF_TOKEN)\ninfo = whoami()\nprint(f\"OK HuggingFace: {info['name']}\")\n\n# ─────────────────────────────────────────────\n# 2. CONFIG\n# ─────────────────────────────────────────────\nSHARD_ID    = 3\nSTART_FOLD  = 1\nDONE_KAPPAS = [0.6339]\n\n# UNI Stage 3 = FROM SCRATCH (pas de PREV_CKPT)\nPREV_CKPT   = PREV_CKPT = \"/kaggle/input/datasets/bejaouikhouloud/stage3-uni-results/stage3_uni_fold1_best.pth\"\n\n# FREEZE UNI backbone pour éviter overfitting !\nFREEZE_UNI  = True   # ← clé pour réduire overfitting\n\nN_FOLDS     = 5\nSEED        = 42\nBATCH_SIZE  = 8      # réduit pour éviter OOM avec UNI\nNUM_WORKERS = 2\nEPOCHS      = 20\nPATIENCE    = 5\nLR_UNI      = 1e-5   # utilisé seulement si FREEZE_UNI=False\nLR_TR       = 4e-5   # Transformer + MLP head\n\nrandom.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)\ntorch.backends.cudnn.deterministic = True\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Device     : {DEVICE}\")\nprint(f\"Stage      : {SHARD_ID} UNI — FROM SCRATCH\")\nprint(f\"Freeze UNI : {FREEZE_UNI}\")\nprint(f\"BATCH={BATCH_SIZE} | EPOCHS={EPOCHS} | PATIENCE={PATIENCE}\")\nprint(f\"LR_TR={LR_TR}\")\n\n# ─────────────────────────────────────────────\n# 3. CHEMINS PATCHES\n# ─────────────────────────────────────────────\nBASE = \"/kaggle/input/datasets/khouloudbejaoui20\"\n\nSHARD_PATHS = {\n    0: f\"{BASE}/dataset/patches_400-20260630T212405Z-3-001/patches_400\",\n    1: f\"{BASE}/panda-patches-run1/kaggle/working/patches_panda_new\",\n    2: f\"{BASE}/panda-patches-run2/kaggle/working/patches_panda_run2\",\n    3: f\"{BASE}/panda-patches-run3/kaggle/working/patches_panda_run3_correct\",\n}\n\nSAVE_DIR     = \"/kaggle/working\"\nKAGGLE_DS    = \"bejaouikhouloud/stage3-uni-results\"\nRESULTS_JSON = f\"{SAVE_DIR}/stage3_uni_cv_results.json\"\n\n# ─────────────────────────────────────────────\n# 4. VERIFICATION SHARDS\n# ─────────────────────────────────────────────\nprint(f\"\\n── Verification ────────────────────────\")\nfor shard_id, path in SHARD_PATHS.items():\n    if os.path.exists(path):\n        files = glob.glob(f\"{path}/*.png\")\n        if not files:\n            files = glob.glob(f\"{path}/**/*.png\", recursive=True)\n        print(f\"  OK Shard {shard_id}: {len(files)} patches\")\n    else:\n        print(f\"  MANQUANT Shard {shard_id}: {path}\")\n\n# ─────────────────────────────────────────────\n# 5. CHARGEMENT PATCHES\n# ─────────────────────────────────────────────\nprint(f\"\\n── Loading patches ──────────────────────\")\n\ndef load_patches(folder):\n    files = glob.glob(f\"{folder}/*.png\")\n    if not files:\n        files = glob.glob(f\"{folder}/**/*.png\", recursive=True)\n    records = []\n    for fp in files:\n        fname = os.path.basename(fp)\n        try:\n            parts  = fname.split(\"_slide\")\n            label  = int(parts[0].replace(\"grade\", \"\"))\n            wsi_id = parts[1].split(\"_\")[0]\n            records.append({\"path\": fp, \"wsi_id\": wsi_id, \"label\": label})\n        except:\n            continue\n    return records\n\nall_records = []\nfor shard_id, path in SHARD_PATHS.items():\n    if path and os.path.exists(path):\n        r = load_patches(path)\n        all_records.extend(r)\n        print(f\"  Shard {shard_id}: {len(r)} patches\")\n    else:\n        print(f\"  Shard {shard_id}: NON TROUVE\")\n\n# Deduplicate\nseen = set()\nunique_records = []\nfor r in all_records:\n    fname = os.path.basename(r[\"path\"])\n    if fname not in seen:\n        seen.add(fname)\n        unique_records.append(r)\n\nall_wsis = list(set(r[\"wsi_id\"] for r in unique_records))\nprint(f\"\\n  Total unique : {len(unique_records)} patches\")\nprint(f\"  Total WSIs  : {len(all_wsis)}\")\ncnt = Counter(r[\"label\"] for r in unique_records)\nfor c in range(6):\n    print(f\"    ISUP {c}: {cnt.get(c, 0)}\")\n\nassert len(all_wsis) > 3000, f\"Trop peu de WSIs ({len(all_wsis)})\"\nprint(f\"\\n  OK {len(all_wsis)} WSIs confirmes\")\n\n# ─────────────────────────────────────────────\n# 6. TRANSFORMS\n# ─────────────────────────────────────────────\nMEAN = [0.485, 0.456, 0.406]\nSTD  = [0.229, 0.224, 0.225]\n\ntrain_tf = T.Compose([\n    T.Resize((224, 224)),\n    T.RandomHorizontalFlip(),\n    T.RandomVerticalFlip(),\n    T.RandomRotation(90),\n    T.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.2, hue=0.1),\n    T.RandomApply([T.GaussianBlur(5)], p=0.3),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n    T.RandomApply([T.RandomErasing(p=0.5, scale=(0.02, 0.2))], p=0.2),\n])\n\nval_tf = T.Compose([\n    T.Resize((224, 224)),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n])\n\nclass PatchDataset(Dataset):\n    def __init__(self, records, transform):\n        self.records   = records\n        self.transform = transform\n    def __len__(self): return len(self.records)\n    def __getitem__(self, idx):\n        r = self.records[idx]\n        try:\n            img = Image.open(r[\"path\"]).convert(\"RGB\")\n        except:\n            img = Image.new(\"RGB\", (224, 224), (255, 255, 255))\n        return self.transform(img), int(r[\"label\"])\n\n# ─────────────────────────────────────────────\n# 7. MODELE — HybridUNITransformer\n# ─────────────────────────────────────────────\nclass HybridUNITransformer(nn.Module):\n    def __init__(self, num_classes=6, d_model=512, nhead=8,\n                 num_layers=2, dim_feedforward=2048, dropout=0.1,\n                 freeze_uni=True):\n        super().__init__()\n\n        # UNI backbone\n        self.uni_backbone = timm.create_model(\n            \"hf-hub:MahmoodLab/UNI\",\n            pretrained=True,\n            num_classes=0,\n            init_values=1e-5,\n            dynamic_img_size=True\n        )\n        uni_dim = self.uni_backbone.num_features  # 1024\n\n        # Freeze UNI si demandé\n        if freeze_uni:\n            for param in self.uni_backbone.parameters():\n                param.requires_grad = False\n            print(f\"  UNI backbone FREEZÉ\")\n        else:\n            print(f\"  UNI backbone trainable\")\n\n        # Projection\n        self.proj = nn.Sequential(\n            nn.Linear(uni_dim, d_model),\n            nn.LayerNorm(d_model)\n        )\n\n        # CLS token + positional embedding\n        self.cls_token = nn.Parameter(torch.zeros(1, 1, d_model))\n        self.pos_embed = nn.Parameter(torch.zeros(1, 2, d_model))\n\n        # Transformer Encoder\n        encoder_layer = nn.TransformerEncoderLayer(\n            d_model=d_model, nhead=nhead,\n            dim_feedforward=dim_feedforward,\n            dropout=dropout, batch_first=True\n        )\n        self.transformer = nn.TransformerEncoder(\n            encoder_layer, num_layers=num_layers,\n            norm=nn.LayerNorm(d_model)\n        )\n\n        # MLP Head\n        self.classifier = nn.Sequential(\n            nn.Linear(d_model, 256),\n            nn.ReLU(),\n            nn.Dropout(dropout),\n            nn.Linear(256, num_classes)\n        )\n\n        nn.init.trunc_normal_(self.cls_token, std=0.02)\n        nn.init.trunc_normal_(self.pos_embed, std=0.02)\n\n    def forward(self, x):\n        with torch.no_grad() if not self.uni_backbone.training else torch.enable_grad():\n            f = self.uni_backbone(x)\n        f   = self.proj(f).unsqueeze(1)\n        cls = self.cls_token.expand(x.size(0), -1, -1)\n        seq = torch.cat([cls, f], dim=1) + self.pos_embed\n        return self.classifier(self.transformer(seq)[:, 0])\n\ndef build_model():\n    # Libère mémoire GPU\n    torch.cuda.empty_cache()\n    import gc; gc.collect()\n\n    model  = HybridUNITransformer(freeze_uni=FREEZE_UNI).to(DEVICE)\n    n_gpus = torch.cuda.device_count()\n    if n_gpus > 1:\n        model = nn.DataParallel(model)\n        print(f\"  OK DataParallel: {n_gpus} GPUs\")\n\n    # Paramètres trainables\n    n_total    = sum(p.numel() for p in model.parameters())\n    n_trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    print(f\"  Total params    : {n_total:,}\")\n    print(f\"  Trainable params: {n_trainable:,}\")\n    return model\n\n# ─────────────────────────────────────────────\n# 8. WSI-LEVEL QWK\n# ─────────────────────────────────────────────\ndef wsi_level_qwk(model, val_data, device):\n    model.eval()\n    wsi_preds  = {}\n    wsi_labels = {}\n    all_preds  = []\n    all_labels = []\n    loader = DataLoader(\n        PatchDataset(val_data, val_tf),\n        batch_size=BATCH_SIZE, shuffle=False,\n        num_workers=NUM_WORKERS\n    )\n    with torch.no_grad():\n        for imgs, labels in loader:\n            imgs  = imgs.to(device)\n            preds = model(imgs).argmax(1).cpu().numpy()\n            all_preds.extend(preds)\n            all_labels.extend(labels.numpy())\n    for i, r in enumerate(val_data):\n        wid = r[\"wsi_id\"]\n        if wid not in wsi_preds:\n            wsi_preds[wid]  = []\n            wsi_labels[wid] = r[\"label\"]\n        wsi_preds[wid].append(all_preds[i])\n    wf = []; wl = []\n    for wid in wsi_preds:\n        votes = wsi_preds[wid]\n        wf.append(max(set(votes), key=votes.count))\n        wl.append(wsi_labels[wid])\n    wsi_qwk = cohen_kappa_score(wl, wf, weights=\"quadratic\")\n    return wsi_qwk, len(wf)\n\n# ─────────────────────────────────────────────\n# 9. PUSH KAGGLE\n# ─────────────────────────────────────────────\ndef push_to_kaggle(ckpt_path, message):\n    try:\n        tmp = f\"{SAVE_DIR}/_push_uni3\"\n        os.makedirs(tmp, exist_ok=True)\n        for c in glob.glob(f\"{SAVE_DIR}/stage{SHARD_ID}_uni*.pth\"):\n            shutil.copy(c, tmp)\n        shutil.copy(ckpt_path, tmp)\n        if os.path.exists(RESULTS_JSON):\n            shutil.copy(RESULTS_JSON, tmp)\n        json.dump({\n            \"title\": \"Stage3 UNI Results Bejaoui 2026\",\n            \"id\": KAGGLE_DS,\n            \"licenses\": [{\"name\": \"CC0-1.0\"}]\n        }, open(f\"{tmp}/dataset-metadata.json\", \"w\"))\n        r = subprocess.run(\n            [\"kaggle\", \"datasets\", \"create\",\n             \"-p\", tmp, \"--dir-mode\", \"skip\"],\n            capture_output=True, text=True, timeout=600\n        )\n        if r.returncode != 0:\n            r = subprocess.run(\n                [\"kaggle\", \"datasets\", \"version\",\n                 \"-p\", tmp, \"-m\", message, \"--dir-mode\", \"skip\"],\n                capture_output=True, text=True, timeout=600\n            )\n        print(f\"   {'OK!' if r.returncode==0 else 'ERREUR: '+r.stderr[:100]}\")\n    except Exception as e:\n        print(f\"   ERREUR: {e}\")\n\n# ─────────────────────────────────────────────\n# 10. TRAIN ONE FOLD\n# ─────────────────────────────────────────────\ndef train_one_fold(fold_id, train_data, val_data, n_folds=5):\n    print(f\"\\n  Fold {fold_id+1}/{n_folds} ──────────────────────────\")\n    print(f\"     Train: {len(train_data)} | Val: {len(val_data)}\")\n    print(f\"     Freeze UNI: {FREEZE_UNI}\")\n\n    train_loader = DataLoader(\n        PatchDataset(train_data, train_tf),\n        batch_size=BATCH_SIZE, shuffle=True,\n        num_workers=NUM_WORKERS, pin_memory=True\n    )\n\n    model = build_model()\n\n    # Class weights\n    cnt_tr = Counter(r[\"label\"] for r in train_data)\n    w = torch.tensor(\n        [1.0 / max(cnt_tr.get(c, 1), 1) for c in range(6)],\n        dtype=torch.float\n    )\n    w = (w / w.sum() * 6).to(DEVICE)\n    criterion = nn.CrossEntropyLoss(weight=w, label_smoothing=0.1)\n\n    # Optimizer — seulement les paramètres trainables\n    base_model = model.module if hasattr(model, \"module\") else model\n\n    if FREEZE_UNI:\n        # Entraîne seulement proj + transformer + classifier\n        trainable_params = [p for p in model.parameters() if p.requires_grad]\n        optimizer = optim.AdamW(trainable_params, lr=LR_TR, weight_decay=1e-4)\n        print(f\"     Optimizer: AdamW lr={LR_TR} (trainable only)\")\n    else:\n        optimizer = optim.AdamW([\n            {\"params\": base_model.uni_backbone.parameters(), \"lr\": LR_UNI},\n            {\"params\": [p for n, p in base_model.named_parameters()\n                        if not n.startswith(\"uni_backbone\")], \"lr\": LR_TR},\n        ], weight_decay=1e-4)\n        print(f\"     Optimizer: AdamW LR_UNI={LR_UNI} LR_TR={LR_TR}\")\n\n    scheduler = CosineAnnealingLR(optimizer, T_max=EPOCHS, eta_min=1e-7)\n\n    best_kappa = 0.0\n    no_improve = 0\n    best_ckpt  = f\"{SAVE_DIR}/stage{SHARD_ID}_uni_fold{fold_id+1}_best.pth\"\n\n    for epoch in range(1, EPOCHS + 1):\n        model.train()\n        if FREEZE_UNI:\n            # S'assure que UNI reste en eval mode\n            base_model = model.module if hasattr(model, \"module\") else model\n            base_model.uni_backbone.eval()\n\n        t_loss, t_correct, t_total = 0.0, 0, 0\n        for imgs, labels in train_loader:\n            imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n            optimizer.zero_grad()\n            out  = model(imgs)\n            loss = criterion(out, labels)\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n            optimizer.step()\n            t_loss    += loss.item() * imgs.size(0)\n            t_correct += (out.argmax(1) == labels).sum().item()\n            t_total   += imgs.size(0)\n        scheduler.step()\n\n        wsi_kappa, n_wsis_val = wsi_level_qwk(model, val_data, DEVICE)\n        is_best = wsi_kappa > best_kappa\n\n        print(f\"     Ep[{epoch:02d}/{EPOCHS}] \"\n              f\"loss={t_loss/t_total:.4f} acc={t_correct/t_total:.3f} | \"\n              f\"WSI_QWK={wsi_kappa:.4f} ({n_wsis_val} WSIs)\"\n              + (\" *\" if is_best else \"\"))\n\n        if is_best:\n            best_kappa = wsi_kappa\n            no_improve = 0\n            torch.save({\n                \"epoch\":            epoch,\n                \"model_state_dict\": model.state_dict() if not hasattr(model, \"module\")\n                                    else model.module.state_dict(),\n                \"kappa\":            wsi_kappa,\n                \"shard_id\":         SHARD_ID,\n                \"fold_id\":          fold_id,\n                \"backbone\":         \"UNI\",\n                \"freeze_uni\":       FREEZE_UNI,\n                \"n_wsis\":           len(all_wsis),\n            }, best_ckpt)\n        else:\n            no_improve += 1\n            if no_improve >= PATIENCE:\n                print(f\"     Early stopping epoch {epoch}\")\n                break\n\n    print(f\"  OK Fold {fold_id+1} best WSI_QWK: {best_kappa:.4f}\")\n    push_to_kaggle(best_ckpt,\n        f\"Stage{SHARD_ID} UNI Fold{fold_id+1} QWK={best_kappa:.4f}\")\n    return best_kappa, best_ckpt\n\n# ─────────────────────────────────────────────\n# 11. 5-FOLD CV\n# ─────────────────────────────────────────────\nprint(f\"\\n{'='*60}\")\nprint(f\"Stage {SHARD_ID} UNI — {len(all_wsis)} WSIs — 5-Fold CV\")\nprint(f\"Backbone: UNI (MahmoodLab) | Freeze: {FREEZE_UNI}\")\nprint(f\"ResNet-50 Stage 3 baseline: mean=0.7280, best=0.7845\")\nprint(f\"{'='*60}\")\n\nwsi_ids    = [r[\"wsi_id\"] for r in unique_records]\nlabels_arr = [r[\"label\"]  for r in unique_records]\nindices    = list(range(len(unique_records)))\n\nsgkf = StratifiedGroupKFold(n_splits=N_FOLDS, shuffle=True, random_state=SEED)\n\nfold_kappas       = DONE_KAPPAS.copy()\nbest_overall      = max(DONE_KAPPAS) if DONE_KAPPAS else -1.0\nbest_ckpt_overall = None\n\nfor fold_id, (train_idx, val_idx) in enumerate(\n        sgkf.split(indices, labels_arr, groups=wsi_ids)):\n\n    if fold_id < START_FOLD:\n        continue\n\n    train_data = [unique_records[i] for i in train_idx]\n    val_data   = [unique_records[i] for i in val_idx]\n\n    tw = set(unique_records[i][\"wsi_id\"] for i in train_idx)\n    vw = set(unique_records[i][\"wsi_id\"] for i in val_idx)\n    assert len(tw & vw) == 0, f\"DATA LEAKAGE fold {fold_id+1}!\"\n\n    kappa, ckpt_path = train_one_fold(fold_id, train_data, val_data, N_FOLDS)\n    fold_kappas.append(kappa)\n\n    if kappa > best_overall:\n        best_overall      = kappa\n        best_ckpt_overall = ckpt_path\n\n# ─────────────────────────────────────────────\n# 12. RESUME FINAL\n# ─────────────────────────────────────────────\nmean_kappa = np.mean(fold_kappas)\nstd_kappa  = np.std(fold_kappas)\n\nprint(f\"\\n{'='*60}\")\nprint(f\"OK Stage {SHARD_ID} UNI — 5-Fold CV Results\")\nprint(f\"{'='*60}\")\nfor i, k in enumerate(fold_kappas):\n    print(f\"   Fold {i+1}: QWK = {k:.4f}\")\nprint(f\"{'─'*40}\")\nprint(f\"   Mean QWK UNI        : {mean_kappa:.4f} +/- {std_kappa:.4f}\")\nprint(f\"   Best QWK UNI        : {best_overall:.4f}\")\nprint(f\"   ResNet-50 Stage 3   : mean=0.7280, best=0.7845\")\nprint(f\"   Gain vs ResNet-50   : {mean_kappa - 0.7280:+.4f}\")\nprint(f\"   WSIs                : {len(all_wsis)}\")\nprint(f\"   SOTA [1]            : 0.934\")\nprint(f\"{'='*60}\")\n\nresults = {\n    \"stage\":           SHARD_ID,\n    \"backbone\":        \"UNI\",\n    \"freeze_uni\":      FREEZE_UNI,\n    \"n_wsis\":          len(all_wsis),\n    \"n_patches\":       len(unique_records),\n    \"fold_kappas\":     fold_kappas,\n    \"mean_kappa\":      mean_kappa,\n    \"std_kappa\":       std_kappa,\n    \"best_kappa\":      best_overall,\n    \"resnet50_mean\":   0.7280,\n    \"resnet50_best\":   0.7845,\n    \"gain_vs_resnet\":  mean_kappa - 0.7280,\n    \"timestamp\":       datetime.datetime.now().isoformat(),\n}\nwith open(RESULTS_JSON, \"w\") as f:\n    json.dump(results, f, indent=2)\n\nif best_ckpt_overall:\n    stage_best = f\"{SAVE_DIR}/stage{SHARD_ID}_uni_best_qwk{best_overall:.4f}.pth\"\n    shutil.copy(best_ckpt_overall, stage_best)\n    print(f\"\\nOK Best checkpoint: {stage_best}\")\n    print(f\"\\nComparaison finale:\")\n    print(f\"  ResNet-50 Stage 3 : mean=0.7280, best=0.7845\")\n    print(f\"  UNI Stage 3       : mean={mean_kappa:.4f}, best={best_overall:.4f}\")\n    print(f\"  Gain mean         : {mean_kappa - 0.7280:+.4f}\")\n    print(f\"  Gain best         : {best_overall - 0.7845:+.4f}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, glob, torch\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# Cherche tous les checkpoints UNI Stage 3\nckpts = sorted(glob.glob(\"/kaggle/working/stage3_uni_fold*.pth\"))\n\nprint(f\"Checkpoints trouvés: {len(ckpts)}\")\nprint(\"-\" * 50)\n\nfor ckpt_path in ckpts:\n    try:\n        ckpt = torch.load(ckpt_path, map_location=\"cpu\", weights_only=False)\n        fold = ckpt.get(\"fold_id\", \"?\")\n        qwk  = ckpt.get(\"kappa\", \"?\")\n        ep   = ckpt.get(\"epoch\", \"?\")\n        mb   = os.path.getsize(ckpt_path)/(1024*1024*1024)\n        print(f\"Fold {fold+1 if isinstance(fold,int) else fold} : \"\n              f\"QWK={qwk:.4f} | Ep={ep} | {mb:.2f} GB\")\n        print(f\"  Path: {os.path.basename(ckpt_path)}\")\n    except Exception as e:\n        print(f\"ERREUR: {ckpt_path} → {e}\")\n\nprint(\"-\" * 50)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\n=============================================================\nStage 3 — HybridUNITransformer — FROM SCRATCH\nUNI (MahmoodLab) foundation model + Transformer Encoder\nDonnées cumulatives :\n  Shard 0 : patches_400        (399  WSIs,  4,319 patches)\n  Shard 1 : panda-patches-run1 (1000 WSIs, 14,842 patches)\n  Shard 2 : panda-patches-run2 (999  WSIs, 14,867 patches)\n  Shard 3 : panda-patches-run3 (999  WSIs, 14,835 patches)\nTotal     : ~3,398 WSIs, ~48,863 patches\n\nComparaison :\n  ResNet-50 Stage 3 : mean=0.7280, best=0.7845\n  UNI Stage 3       : ???\n=============================================================\nAttache dans Kaggle :\n  khouloudbejaoui20/dataset\n  khouloudbejaoui20/panda-patches-run1\n  khouloudbejaoui20/panda-patches-run2\n  khouloudbejaoui20/panda-patches-run3\nGPU : T4×2\n\"\"\"\n\nimport os, glob, json, random, shutil, subprocess, datetime\nimport numpy as np\nfrom PIL import Image\nfrom collections import Counter\nfrom sklearn.model_selection import StratifiedGroupKFold\nfrom sklearn.metrics import cohen_kappa_score\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as T\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\n\n# ─────────────────────────────────────────────\n# 1. LOGIN HUGGINGFACE\n# ─────────────────────────────────────────────\nfrom huggingface_hub import login, whoami\nimport timm\n\nHF_TOKEN = \"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\"\nlogin(token=HF_TOKEN)\ninfo = whoami()\nprint(f\"OK HuggingFace: {info['name']}\")\n\n# ─────────────────────────────────────────────\n# 2. CONFIG\n# ─────────────────────────────────────────────\nSHARD_ID    = 3\nSTART_FOLD  = 3\n\nDONE_KAPPAS = [0.6339, 0.6622, 0.6658]  #fold4\n# UNI Stage 3 = FROM SCRATCH (pas de PREV_CKPT)\nPREV_CKPT   =None\n\n# FREEZE UNI backbone pour éviter overfitting !\nFREEZE_UNI  = True   # ← clé pour réduire overfitting\n\nN_FOLDS     = 5\nSEED        = 42\nBATCH_SIZE  = 8      # réduit pour éviter OOM avec UNI\nNUM_WORKERS = 2\nEPOCHS      = 20\nPATIENCE    = 5\nLR_UNI      = 1e-5   # utilisé seulement si FREEZE_UNI=False\nLR_TR       = 4e-5   # Transformer + MLP head\n\nrandom.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)\ntorch.backends.cudnn.deterministic = True\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Device     : {DEVICE}\")\nprint(f\"Stage      : {SHARD_ID} UNI — FROM SCRATCH\")\nprint(f\"Freeze UNI : {FREEZE_UNI}\")\nprint(f\"BATCH={BATCH_SIZE} | EPOCHS={EPOCHS} | PATIENCE={PATIENCE}\")\nprint(f\"LR_TR={LR_TR}\")\n\n# ─────────────────────────────────────────────\n# 3. CHEMINS PATCHES\n# ─────────────────────────────────────────────\nBASE = \"/kaggle/input/datasets/khouloudbejaoui20\"\n\nSHARD_PATHS = {\n    0: f\"{BASE}/dataset/patches_400-20260630T212405Z-3-001/patches_400\",\n    1: f\"{BASE}/panda-patches-run1/kaggle/working/patches_panda_new\",\n    2: f\"{BASE}/panda-patches-run2/kaggle/working/patches_panda_run2\",\n    3: f\"{BASE}/panda-patches-run3/kaggle/working/patches_panda_run3_correct\",\n}\n\nSAVE_DIR     = \"/kaggle/working\"\nKAGGLE_DS    = \"bejaouikhouloud/stage3-uni-results\"\nRESULTS_JSON = f\"{SAVE_DIR}/stage3_uni_cv_results.json\"\n\n# ─────────────────────────────────────────────\n# 4. VERIFICATION SHARDS\n# ─────────────────────────────────────────────\nprint(f\"\\n── Verification ────────────────────────\")\nfor shard_id, path in SHARD_PATHS.items():\n    if os.path.exists(path):\n        files = glob.glob(f\"{path}/*.png\")\n        if not files:\n            files = glob.glob(f\"{path}/**/*.png\", recursive=True)\n        print(f\"  OK Shard {shard_id}: {len(files)} patches\")\n    else:\n        print(f\"  MANQUANT Shard {shard_id}: {path}\")\n\n# ─────────────────────────────────────────────\n# 5. CHARGEMENT PATCHES\n# ─────────────────────────────────────────────\nprint(f\"\\n── Loading patches ──────────────────────\")\n\ndef load_patches(folder):\n    files = glob.glob(f\"{folder}/*.png\")\n    if not files:\n        files = glob.glob(f\"{folder}/**/*.png\", recursive=True)\n    records = []\n    for fp in files:\n        fname = os.path.basename(fp)\n        try:\n            parts  = fname.split(\"_slide\")\n            label  = int(parts[0].replace(\"grade\", \"\"))\n            wsi_id = parts[1].split(\"_\")[0]\n            records.append({\"path\": fp, \"wsi_id\": wsi_id, \"label\": label})\n        except:\n            continue\n    return records\n\nall_records = []\nfor shard_id, path in SHARD_PATHS.items():\n    if path and os.path.exists(path):\n        r = load_patches(path)\n        all_records.extend(r)\n        print(f\"  Shard {shard_id}: {len(r)} patches\")\n    else:\n        print(f\"  Shard {shard_id}: NON TROUVE\")\n\n# Deduplicate\nseen = set()\nunique_records = []\nfor r in all_records:\n    fname = os.path.basename(r[\"path\"])\n    if fname not in seen:\n        seen.add(fname)\n        unique_records.append(r)\n\nall_wsis = list(set(r[\"wsi_id\"] for r in unique_records))\nprint(f\"\\n  Total unique : {len(unique_records)} patches\")\nprint(f\"  Total WSIs  : {len(all_wsis)}\")\ncnt = Counter(r[\"label\"] for r in unique_records)\nfor c in range(6):\n    print(f\"    ISUP {c}: {cnt.get(c, 0)}\")\n\nassert len(all_wsis) > 3000, f\"Trop peu de WSIs ({len(all_wsis)})\"\nprint(f\"\\n  OK {len(all_wsis)} WSIs confirmes\")\n\n# ─────────────────────────────────────────────\n# 6. TRANSFORMS\n# ─────────────────────────────────────────────\nMEAN = [0.485, 0.456, 0.406]\nSTD  = [0.229, 0.224, 0.225]\n\ntrain_tf = T.Compose([\n    T.Resize((224, 224)),\n    T.RandomHorizontalFlip(),\n    T.RandomVerticalFlip(),\n    T.RandomRotation(90),\n    T.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.2, hue=0.1),\n    T.RandomApply([T.GaussianBlur(5)], p=0.3),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n    T.RandomApply([T.RandomErasing(p=0.5, scale=(0.02, 0.2))], p=0.2),\n])\n\nval_tf = T.Compose([\n    T.Resize((224, 224)),\n    T.ToTensor(),\n    T.Normalize(MEAN, STD),\n])\n\nclass PatchDataset(Dataset):\n    def __init__(self, records, transform):\n        self.records   = records\n        self.transform = transform\n    def __len__(self): return len(self.records)\n    def __getitem__(self, idx):\n        r = self.records[idx]\n        try:\n            img = Image.open(r[\"path\"]).convert(\"RGB\")\n        except:\n            img = Image.new(\"RGB\", (224, 224), (255, 255, 255))\n        return self.transform(img), int(r[\"label\"])\n\n# ─────────────────────────────────────────────\n# 7. MODELE — HybridUNITransformer\n# ─────────────────────────────────────────────\nclass HybridUNITransformer(nn.Module):\n    def __init__(self, num_classes=6, d_model=512, nhead=8,\n                 num_layers=2, dim_feedforward=2048, dropout=0.1,\n                 freeze_uni=True):\n        super().__init__()\n\n        # UNI backbone\n        self.uni_backbone = timm.create_model(\n            \"hf-hub:MahmoodLab/UNI\",\n            pretrained=True,\n            num_classes=0,\n            init_values=1e-5,\n            dynamic_img_size=True\n        )\n        uni_dim = self.uni_backbone.num_features  # 1024\n\n        # Freeze UNI si demandé\n        if freeze_uni:\n            for param in self.uni_backbone.parameters():\n                param.requires_grad = False\n            print(f\"  UNI backbone FREEZÉ\")\n        else:\n            print(f\"  UNI backbone trainable\")\n\n        # Projection\n        self.proj = nn.Sequential(\n            nn.Linear(uni_dim, d_model),\n            nn.LayerNorm(d_model)\n        )\n\n        # CLS token + positional embedding\n        self.cls_token = nn.Parameter(torch.zeros(1, 1, d_model))\n        self.pos_embed = nn.Parameter(torch.zeros(1, 2, d_model))\n\n        # Transformer Encoder\n        encoder_layer = nn.TransformerEncoderLayer(\n            d_model=d_model, nhead=nhead,\n            dim_feedforward=dim_feedforward,\n            dropout=dropout, batch_first=True\n        )\n        self.transformer = nn.TransformerEncoder(\n            encoder_layer, num_layers=num_layers,\n            norm=nn.LayerNorm(d_model)\n        )\n\n        # MLP Head\n        self.classifier = nn.Sequential(\n            nn.Linear(d_model, 256),\n            nn.ReLU(),\n            nn.Dropout(dropout),\n            nn.Linear(256, num_classes)\n        )\n\n        nn.init.trunc_normal_(self.cls_token, std=0.02)\n        nn.init.trunc_normal_(self.pos_embed, std=0.02)\n\n    def forward(self, x):\n        with torch.no_grad() if not self.uni_backbone.training else torch.enable_grad():\n            f = self.uni_backbone(x)\n        f   = self.proj(f).unsqueeze(1)\n        cls = self.cls_token.expand(x.size(0), -1, -1)\n        seq = torch.cat([cls, f], dim=1) + self.pos_embed\n        return self.classifier(self.transformer(seq)[:, 0])\n\ndef build_model():\n    # Libère mémoire GPU\n    torch.cuda.empty_cache()\n    import gc; gc.collect()\n\n    model  = HybridUNITransformer(freeze_uni=FREEZE_UNI).to(DEVICE)\n    n_gpus = torch.cuda.device_count()\n    if n_gpus > 1:\n        model = nn.DataParallel(model)\n        print(f\"  OK DataParallel: {n_gpus} GPUs\")\n\n    # Paramètres trainables\n    n_total    = sum(p.numel() for p in model.parameters())\n    n_trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    print(f\"  Total params    : {n_total:,}\")\n    print(f\"  Trainable params: {n_trainable:,}\")\n    return model\n\n# ─────────────────────────────────────────────\n# 8. WSI-LEVEL QWK\n# ─────────────────────────────────────────────\ndef wsi_level_qwk(model, val_data, device):\n    model.eval()\n    wsi_preds  = {}\n    wsi_labels = {}\n    all_preds  = []\n    all_labels = []\n    loader = DataLoader(\n        PatchDataset(val_data, val_tf),\n        batch_size=BATCH_SIZE, shuffle=False,\n        num_workers=NUM_WORKERS\n    )\n    with torch.no_grad():\n        for imgs, labels in loader:\n            imgs  = imgs.to(device)\n            preds = model(imgs).argmax(1).cpu().numpy()\n            all_preds.extend(preds)\n            all_labels.extend(labels.numpy())\n    for i, r in enumerate(val_data):\n        wid = r[\"wsi_id\"]\n        if wid not in wsi_preds:\n            wsi_preds[wid]  = []\n            wsi_labels[wid] = r[\"label\"]\n        wsi_preds[wid].append(all_preds[i])\n    wf = []; wl = []\n    for wid in wsi_preds:\n        votes = wsi_preds[wid]\n        wf.append(max(set(votes), key=votes.count))\n        wl.append(wsi_labels[wid])\n    wsi_qwk = cohen_kappa_score(wl, wf, weights=\"quadratic\")\n    return wsi_qwk, len(wf)\n\n# ─────────────────────────────────────────────\n# 9. PUSH KAGGLE\n# ─────────────────────────────────────────────\ndef push_to_kaggle(ckpt_path, message):\n    try:\n        tmp = f\"{SAVE_DIR}/_push_uni3\"\n        os.makedirs(tmp, exist_ok=True)\n        for c in glob.glob(f\"{SAVE_DIR}/stage{SHARD_ID}_uni*.pth\"):\n            shutil.copy(c, tmp)\n        shutil.copy(ckpt_path, tmp)\n        if os.path.exists(RESULTS_JSON):\n            shutil.copy(RESULTS_JSON, tmp)\n        json.dump({\n            \"title\": \"Stage3 UNI Results Bejaoui 2026\",\n            \"id\": KAGGLE_DS,\n            \"licenses\": [{\"name\": \"CC0-1.0\"}]\n        }, open(f\"{tmp}/dataset-metadata.json\", \"w\"))\n        r = subprocess.run(\n            [\"kaggle\", \"datasets\", \"create\",\n             \"-p\", tmp, \"--dir-mode\", \"skip\"],\n            capture_output=True, text=True, timeout=600\n        )\n        if r.returncode != 0:\n            r = subprocess.run(\n                [\"kaggle\", \"datasets\", \"version\",\n                 \"-p\", tmp, \"-m\", message, \"--dir-mode\", \"skip\"],\n                capture_output=True, text=True, timeout=600\n            )\n        print(f\"   {'OK!' if r.returncode==0 else 'ERREUR: '+r.stderr[:100]}\")\n    except Exception as e:\n        print(f\"   ERREUR: {e}\")\n\n# ─────────────────────────────────────────────\n# 10. TRAIN ONE FOLD\n# ─────────────────────────────────────────────\ndef train_one_fold(fold_id, train_data, val_data, n_folds=5):\n    print(f\"\\n  Fold {fold_id+1}/{n_folds} ──────────────────────────\")\n    print(f\"     Train: {len(train_data)} | Val: {len(val_data)}\")\n    print(f\"     Freeze UNI: {FREEZE_UNI}\")\n\n    train_loader = DataLoader(\n        PatchDataset(train_data, train_tf),\n        batch_size=BATCH_SIZE, shuffle=True,\n        num_workers=NUM_WORKERS, pin_memory=True\n    )\n\n    model = build_model()\n\n    # Class weights\n    cnt_tr = Counter(r[\"label\"] for r in train_data)\n    w = torch.tensor(\n        [1.0 / max(cnt_tr.get(c, 1), 1) for c in range(6)],\n        dtype=torch.float\n    )\n    w = (w / w.sum() * 6).to(DEVICE)\n    criterion = nn.CrossEntropyLoss(weight=w, label_smoothing=0.1)\n\n    # Optimizer — seulement les paramètres trainables\n    base_model = model.module if hasattr(model, \"module\") else model\n\n    if FREEZE_UNI:\n        # Entraîne seulement proj + transformer + classifier\n        trainable_params = [p for p in model.parameters() if p.requires_grad]\n        optimizer = optim.AdamW(trainable_params, lr=LR_TR, weight_decay=1e-4)\n        print(f\"     Optimizer: AdamW lr={LR_TR} (trainable only)\")\n    else:\n        optimizer = optim.AdamW([\n            {\"params\": base_model.uni_backbone.parameters(), \"lr\": LR_UNI},\n            {\"params\": [p for n, p in base_model.named_parameters()\n                        if not n.startswith(\"uni_backbone\")], \"lr\": LR_TR},\n        ], weight_decay=1e-4)\n        print(f\"     Optimizer: AdamW LR_UNI={LR_UNI} LR_TR={LR_TR}\")\n\n    scheduler = CosineAnnealingLR(optimizer, T_max=EPOCHS, eta_min=1e-7)\n\n    best_kappa = 0.0\n    no_improve = 0\n    best_ckpt  = f\"{SAVE_DIR}/stage{SHARD_ID}_uni_fold{fold_id+1}_best.pth\"\n\n    for epoch in range(1, EPOCHS + 1):\n        model.train()\n        if FREEZE_UNI:\n            # S'assure que UNI reste en eval mode\n            base_model = model.module if hasattr(model, \"module\") else model\n            base_model.uni_backbone.eval()\n\n        t_loss, t_correct, t_total = 0.0, 0, 0\n        for imgs, labels in train_loader:\n            imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n            optimizer.zero_grad()\n            out  = model(imgs)\n            loss = criterion(out, labels)\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n            optimizer.step()\n            t_loss    += loss.item() * imgs.size(0)\n            t_correct += (out.argmax(1) == labels).sum().item()\n            t_total   += imgs.size(0)\n        scheduler.step()\n\n        wsi_kappa, n_wsis_val = wsi_level_qwk(model, val_data, DEVICE)\n        is_best = wsi_kappa > best_kappa\n\n        print(f\"     Ep[{epoch:02d}/{EPOCHS}] \"\n              f\"loss={t_loss/t_total:.4f} acc={t_correct/t_total:.3f} | \"\n              f\"WSI_QWK={wsi_kappa:.4f} ({n_wsis_val} WSIs)\"\n              + (\" *\" if is_best else \"\"))\n\n        if is_best:\n            best_kappa = wsi_kappa\n            no_improve = 0\n            torch.save({\n                \"epoch\":            epoch,\n                \"model_state_dict\": model.state_dict() if not hasattr(model, \"module\")\n                                    else model.module.state_dict(),\n                \"kappa\":            wsi_kappa,\n                \"shard_id\":         SHARD_ID,\n                \"fold_id\":          fold_id,\n                \"backbone\":         \"UNI\",\n                \"freeze_uni\":       FREEZE_UNI,\n                \"n_wsis\":           len(all_wsis),\n            }, best_ckpt)\n        else:\n            no_improve += 1\n            if no_improve >= PATIENCE:\n                print(f\"     Early stopping epoch {epoch}\")\n                break\n\n    print(f\"  OK Fold {fold_id+1} best WSI_QWK: {best_kappa:.4f}\")\n    push_to_kaggle(best_ckpt,\n        f\"Stage{SHARD_ID} UNI Fold{fold_id+1} QWK={best_kappa:.4f}\")\n    return best_kappa, best_ckpt\n\n# ─────────────────────────────────────────────\n# 11. 5-FOLD CV\n# ─────────────────────────────────────────────\nprint(f\"\\n{'='*60}\")\nprint(f\"Stage {SHARD_ID} UNI — {len(all_wsis)} WSIs — 5-Fold CV\")\nprint(f\"Backbone: UNI (MahmoodLab) | Freeze: {FREEZE_UNI}\")\nprint(f\"ResNet-50 Stage 3 baseline: mean=0.7280, best=0.7845\")\nprint(f\"{'='*60}\")\n\nwsi_ids    = [r[\"wsi_id\"] for r in unique_records]\nlabels_arr = [r[\"label\"]  for r in unique_records]\nindices    = list(range(len(unique_records)))\n\nsgkf = StratifiedGroupKFold(n_splits=N_FOLDS, shuffle=True, random_state=SEED)\n\nfold_kappas       = DONE_KAPPAS.copy()\nbest_overall      = max(DONE_KAPPAS) if DONE_KAPPAS else -1.0\nbest_ckpt_overall = None\n\nfor fold_id, (train_idx, val_idx) in enumerate(\n        sgkf.split(indices, labels_arr, groups=wsi_ids)):\n\n    if fold_id < START_FOLD:\n        continue\n\n    train_data = [unique_records[i] for i in train_idx]\n    val_data   = [unique_records[i] for i in val_idx]\n\n    tw = set(unique_records[i][\"wsi_id\"] for i in train_idx)\n    vw = set(unique_records[i][\"wsi_id\"] for i in val_idx)\n    assert len(tw & vw) == 0, f\"DATA LEAKAGE fold {fold_id+1}!\"\n\n    kappa, ckpt_path = train_one_fold(fold_id, train_data, val_data, N_FOLDS)\n    fold_kappas.append(kappa)\n\n    if kappa > best_overall:\n        best_overall      = kappa\n        best_ckpt_overall = ckpt_path\n\n# ─────────────────────────────────────────────\n# 12. RESUME FINAL\n# ─────────────────────────────────────────────\nmean_kappa = np.mean(fold_kappas)\nstd_kappa  = np.std(fold_kappas)\n\nprint(f\"\\n{'='*60}\")\nprint(f\"OK Stage {SHARD_ID} UNI — 5-Fold CV Results\")\nprint(f\"{'='*60}\")\nfor i, k in enumerate(fold_kappas):\n    print(f\"   Fold {i+1}: QWK = {k:.4f}\")\nprint(f\"{'─'*40}\")\nprint(f\"   Mean QWK UNI        : {mean_kappa:.4f} +/- {std_kappa:.4f}\")\nprint(f\"   Best QWK UNI        : {best_overall:.4f}\")\nprint(f\"   ResNet-50 Stage 3   : mean=0.7280, best=0.7845\")\nprint(f\"   Gain vs ResNet-50   : {mean_kappa - 0.7280:+.4f}\")\nprint(f\"   WSIs                : {len(all_wsis)}\")\nprint(f\"   SOTA [1]            : 0.934\")\nprint(f\"{'='*60}\")\n\nresults = {\n    \"stage\":           SHARD_ID,\n    \"backbone\":        \"UNI\",\n    \"freeze_uni\":      FREEZE_UNI,\n    \"n_wsis\":          len(all_wsis),\n    \"n_patches\":       len(unique_records),\n    \"fold_kappas\":     fold_kappas,\n    \"mean_kappa\":      mean_kappa,\n    \"std_kappa\":       std_kappa,\n    \"best_kappa\":      best_overall,\n    \"resnet50_mean\":   0.7280,\n    \"resnet50_best\":   0.7845,\n    \"gain_vs_resnet\":  mean_kappa - 0.7280,\n    \"timestamp\":       datetime.datetime.now().isoformat(),\n}\nwith open(RESULTS_JSON, \"w\") as f:\n    json.dump(results, f, indent=2)\n\nif best_ckpt_overall:\n    stage_best = f\"{SAVE_DIR}/stage{SHARD_ID}_uni_best_qwk{best_overall:.4f}.pth\"\n    shutil.copy(best_ckpt_overall, stage_best)\n    print(f\"\\nOK Best checkpoint: {stage_best}\")\n    print(f\"\\nComparaison finale:\")\n    print(f\"  ResNet-50 Stage 3 : mean=0.7280, best=0.7845\")\n    print(f\"  UNI Stage 3       : mean={mean_kappa:.4f}, best={best_overall:.4f}\")\n    print(f\"  Gain mean         : {mean_kappa - 0.7280:+.4f}\")\n    print(f\"  Gain best         : {best_overall - 0.7845:+.4f}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}