{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":4104,"databundleVersionId":46661,"sourceType":"competition"},{"sourceId":7866129,"sourceType":"datasetVersion","datasetId":4614938},{"sourceId":7869237,"sourceType":"datasetVersion","datasetId":4617269},{"sourceId":12489454,"sourceType":"datasetVersion","datasetId":7881389}],"dockerImageVersionId":31089,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true},"papermill":{"default_parameters":{},"duration":13.107283,"end_time":"2025-03-03T11:50:15.228776","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2025-03-03T11:50:02.121493","version":"2.5.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Retinopathy : entraînement du modèle\n\nOn entraîne un modèle de classification d’images pour détecter les stades de rétinopathie diabétique.\\\nObjectif : **maximiser la Quadratic Weighted Kappa (QWK)** en validation, en gardant un pipeline simple, reproductible et rapide à itérer.\n\n## Importation des packages","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport tensorflow as tf\nfrom glob import glob\nimport os\nimport cv2\n\nfrom sklearn.model_selection import train_test_split\n\nfrom tensorflow.keras.applications import ConvNeXtTiny\nfrom tensorflow.keras import layers, models, callbacks\nfrom tensorflow.keras import mixed_precision\n\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-31T12:23:32.881324Z","iopub.execute_input":"2025-08-31T12:23:32.881957Z","iopub.status.idle":"2025-08-31T12:23:45.404781Z","shell.execute_reply.started":"2025-08-31T12:23:32.881934Z","shell.execute_reply":"2025-08-31T12:23:45.404160Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"file_lbl=\"/kaggle/input/diabetic-retinopathy-detection/trainLabels.csv.zip\"\ndf_train=pd.read_csv(file_lbl,sep=',')\ndisplay(df_train)\n\npaths_train= glob('/kaggle/input/diabetic-retinopathy-train-unzipped/train/*.jpeg')\nimage=cv2.imread(paths_train[0])\nplt.imshow(image)\n\npaths_test = glob('/kaggle/input/diabetic-retinopathy-test-unzipped/test/*.jpeg')\nimage=cv2.imread(paths_test[0])\nplt.imshow(image)\n\nfile_sub=\"/kaggle/input/diabetic-retinopathy-detection/sampleSubmission.csv.zip\"\ndf_submission=pd.read_csv(file_sub,sep=',')\ndf_submission.loc[0, 'level']=1\ndisplay(df_submission)\n\ndf_submission.to_csv('submission.csv', index=False)","metadata":{"papermill":{"duration":1.649134,"end_time":"2025-03-03T11:50:07.748810","exception":false,"start_time":"2025-03-03T11:50:06.099676","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2025-08-31T12:23:45.405763Z","iopub.execute_input":"2025-08-31T12:23:45.406307Z","iopub.status.idle":"2025-08-31T12:23:49.687628Z","shell.execute_reply.started":"2025-08-31T12:23:45.406281Z","shell.execute_reply":"2025-08-31T12:23:49.686890Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Prétraitement des images de rétine\nDétection du disque optique, recadrage circulaire, redimensionnement à 512×512 px et normalisation z-score.","metadata":{}},{"cell_type":"code","source":"def preprocess(img_path, target_size=(512, 512)):\n    # Lecture de l'image et conversion en RGB\n    img = cv2.imread(img_path)\n    img_rgb = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    \n    # Détection du masque de la rétine (seuil simple)\n    gray = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)\n    _, mask = cv2.threshold(gray, 15, 255, cv2.THRESH_BINARY)\n    \n    # Extraction du plus grand contour (zone circulaire)\n    contours, _ = cv2.findContours(mask, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n    cnt = max(contours, key=cv2.contourArea)\n    x, y, w, h = cv2.boundingRect(cnt)\n    \n    # Recadrage à la zone d'intérêt\n    cropped = img_rgb[y:y+h, x:x+w]\n    \n    # Redimensionnement uniforme\n    resized = cv2.resize(cropped, target_size, interpolation=cv2.INTER_AREA)\n    \n    # Normalisation : passage en [0,1], puis centrage-réduction\n    norm = resized.astype('float32') / 255.0\n    mean = norm.mean(axis=(0,1), keepdims=True)\n    std = norm.std(axis=(0,1), keepdims=True)\n    norm = (norm - mean) / (std + 1e-8)\n    \n    return img_rgb, resized, norm\n\n# Pour démonstration, on ne traite que les 5 premières images (éviter de surcharger l'affichage)\nsample_paths = paths_train[:5]\n\n# Prétraitement et affichage pour chaque image de l'échantillon\nfor img_path in sample_paths:\n    orig, resized, norm = preprocess(img_path)\n    plt.figure(figsize=(15, 5))\n    plt.subplot(1, 3, 1)\n    plt.imshow(orig)\n    plt.title('Original')\n    plt.axis('off')\n    \n    plt.subplot(1, 3, 2)\n    plt.imshow(resized)\n    plt.title('Recadrée & Redimensionnée')\n    plt.axis('off')\n    \n    plt.subplot(1, 3, 3)\n    disp = (norm - norm.min()) / (norm.max() - norm.min())\n    plt.imshow(disp)\n    plt.title('Normalisée (z-score)')\n    plt.axis('off')\n    \n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-31T12:23:49.688364Z","iopub.execute_input":"2025-08-31T12:23:49.688579Z","iopub.status.idle":"2025-08-31T12:23:56.928742Z","shell.execute_reply.started":"2025-08-31T12:23:49.688563Z","shell.execute_reply":"2025-08-31T12:23:56.928037Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 1. Charger le CSV des labels depuis /mnt/data\nlabels_zip = '/kaggle/input/diabetic-retinopathy-detection/trainLabels.csv.zip' \nlabels_df = pd.read_csv(labels_zip, compression='zip')\n\nlabels_df.rename(columns={'image': 'image_id'}, inplace=True)\n\n# 2. Lister toutes les images d’entraînement dans /mnt/data\npaths_train = glob('/kaggle/input/diabetic-retinopathy-train-unzipped/train/*.jpeg')\n\n# 3. Construire le DataFrame chemins ↔ image_id\ndf_paths = pd.DataFrame({\n    'path': paths_train,\n    'image_id': [os.path.splitext(os.path.basename(p))[0] for p in paths_train]\n})\n\n# 4. Fusionner labels et chemins\ndf = df_paths.merge(labels_df, on='image_id')\nprint(\"Total images disponibles :\", len(df))\n\n# 5. Tentative de split stratifié, fallback sans stratification si nécessaire\ntry:\n    train_df, val_df = train_test_split(\n        df,\n        test_size=0.2,\n        stratify=df['level'],\n        random_state=42\n    )\n    print(\"Split stratifié réalisé.\")\nexcept ValueError:\n    train_df, val_df = train_test_split(\n        df,\n        test_size=0.2,\n        random_state=42\n    )\n    print(\"Stratification impossible (échantillons insuffisants) : split aléatoire réalisé.\")\n\nprint(f\"Jeu d'entraînement : {train_df.shape[0]} images\")\nprint(f\"Jeu de validation : {val_df.shape[0]} images\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-31T12:23:56.930344Z","iopub.execute_input":"2025-08-31T12:23:56.930843Z","iopub.status.idle":"2025-08-31T12:23:57.089344Z","shell.execute_reply.started":"2025-08-31T12:23:56.930821Z","shell.execute_reply":"2025-08-31T12:23:57.088660Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Imports & configuration \n\n- **Choix `PREPROCESS_MODE`** :  \n  `fast` = simple et rapide, `cached` = mêmes étapes mais mis en cache pour accélérer les epochs suivantes, `robust` = pipeline plus défensif (meilleurs crops/masques) pour des images bruitées, le tout pour équilibrer **vitesse** vs **qualité** selon le contexte.\n\n- **Backbone (`resnet18` / `resnet50`)** :  \n  ResNet18 = **itération rapide** (débogage, prototypage), ResNet50 = **capacité** et potentiel de score plus élevé. Il permet d'adapter le **compromis perf/temps** à l’étape du projet.\n\n- On commence avec un LR plus haut pour apprendre la tête, puis on affine le backbone avec un LR plus bas dans le but de stabiliser l’entraînement et extraire des **représentations fines** sans tout “casser”.\n\n- On fixe les graines aléatoires **seed** et on règle cuDNN pour assurer la **reproductibilité** des résultats tout en gardant de **bonnes performances** quand la taille des batches/images est fixe.\n","metadata":{}},{"cell_type":"code","source":"import os, time, cv2, numpy as np, torch\nfrom torch import nn\nfrom torch.optim import SGD\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import models, transforms\nfrom sklearn.metrics import cohen_kappa_score\nimport torch.backends.cudnn as cudnn\n\n# ---- Config ----\nPREPROCESS_MODE = \"fast\"   # \"fast\" | \"cached\" | \"robust\"\nIMG_SIZE = 224\nBACKBONE = \"resnet18\"      # \"resnet18\" | \"resnet50\" \nBATCH_SIZE = 64            # 32 | 64 \nNUM_EPOCHS = 6             # 4 | 5 | 6\nFINE_TUNE_EPOCHS = 3       # epochs supplémentaires après dégel\nPATIENCE = 2               # early stopping patience (sur κ)\nLR_BASE = 3e-3\nLR_FT = 1e-3\nWEIGHT_DECAY = 1e-4\nMOMENTUM = 0.9\nBEST_PATH = \"/kaggle/working/best_model.pth\"\nCACHE_DIR = \"/kaggle/working/cache224\"\nos.makedirs(CACHE_DIR, exist_ok=True)\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\ntorch.set_float32_matmul_precision(\"high\")\ncudnn.benchmark = True\n\n# Fonctions utiles au processus\ndef _center_square_crop(img):\n    h, w = img.shape[:2]\n    side = min(h, w)\n    y0 = (h - side) // 2\n    x0 = (w - side) // 2\n    return img[y0:y0+side, x0:x0+side]\n\ndef preprocess_fast(img_path, target_size=(IMG_SIZE, IMG_SIZE)):\n    img = cv2.imread(img_path, cv2.IMREAD_COLOR)\n    if img is None: raise FileNotFoundError(img_path)\n    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    img = _center_square_crop(img)\n    img = cv2.resize(img, target_size, interpolation=cv2.INTER_AREA)\n    norm = img.astype('float32')/255.0\n    return img, img, norm\n\ndef _cache_path(p):\n    base = os.path.splitext(os.path.basename(p))[0] + f\"_{IMG_SIZE}.jpg\"\n    return os.path.join(CACHE_DIR, base)\n\ndef preprocess_cached(img_path, target_size=(IMG_SIZE, IMG_SIZE)):\n    cp = _cache_path(img_path)\n    if os.path.exists(cp):\n        img = cv2.imread(cp, cv2.IMREAD_COLOR)\n        if img is None:  # si cache cassé\n            os.remove(cp)\n            return preprocess_cached(img_path, target_size)\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    else:\n        img = cv2.imread(img_path, cv2.IMREAD_COLOR)\n        if img is None: raise FileNotFoundError(img_path)\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        img = _center_square_crop(img)\n        img = cv2.resize(img, target_size, interpolation=cv2.INTER_AREA)\n        cv2.imwrite(cp, cv2.cvtColor(img, cv2.COLOR_RGB2BGR), [int(cv2.IMWRITE_JPEG_QUALITY), 95])\n    norm = img.astype('float32')/255.0\n    return img, img, norm\n\ndef preprocess_robuste(img_path, target_size=(IMG_SIZE, IMG_SIZE), min_area_ratio=0.05):\n    img = cv2.imread(img_path, cv2.IMREAD_COLOR)\n    if img is None:\n        raise FileNotFoundError(f\"Impossible de lire {img_path}\")\n    img_rgb = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n\n    gray = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)\n    gray_blur = cv2.GaussianBlur(gray, (0,0), 3)\n    _, th_otsu = cv2.threshold(gray_blur, 0, 255, cv2.THRESH_BINARY+cv2.THRESH_OTSU)\n    mask_gray = (gray > 10).astype(np.uint8) * 255\n    hsv = cv2.cvtColor(img, cv2.COLOR_BGR2HSV)\n    v = hsv[:,:,2]\n    mask_v = (v > 15).astype(np.uint8) * 255\n    mask = cv2.bitwise_or(th_otsu, cv2.bitwise_and(mask_gray, mask_v))\n    mask = cv2.morphologyEx(mask, cv2.MORPH_CLOSE, np.ones((9,9), np.uint8))\n    mask = cv2.morphologyEx(mask, cv2.MORPH_OPEN,  np.ones((5,5), np.uint8))\n\n    contours, _ = cv2.findContours(mask, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n    cropped = None\n    if contours:\n        cnt = max(contours, key=cv2.contourArea)\n        x,y,w,h = cv2.boundingRect(cnt)\n        H,W = img_rgb.shape[:2]\n        if (w*h) >= (min_area_ratio * H * W):\n            side = int(max(w, h))\n            cx, cy = x + w//2, y + h//2\n            x0 = max(0, cx - side//2)\n            y0 = max(0, cy - side//2)\n            x1 = min(W, x0 + side)\n            y1 = min(H, y0 + side)\n            cropped = img_rgb[y0:y1, x0:x1]\n    if cropped is None or cropped.size == 0:\n        cropped = _center_square_crop(img_rgb)\n\n    resized = cv2.resize(cropped, target_size, interpolation=cv2.INTER_AREA)\n    norm = resized.astype('float32')/255.0\n    return img_rgb, resized, norm\n\n# Map des fonctions\nPREPROC_FN = {\n    \"fast\": preprocess_fast,\n    \"cached\": preprocess_cached,\n    \"robust\": preprocess_robuste\n}[PREPROCESS_MODE]\n\n#  Dataset\nclass RetinopathyDataset(Dataset):\n    def __init__(self, df, transform=None, preprocess_func=preprocess_fast):\n        self.df = df.reset_index(drop=True)\n        self.transform = transform\n        self.preproc = preprocess_func\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img_path, label = row['path'], int(row['level'])\n        try:\n            _, resized, norm = self.preproc(img_path)\n            img_uint8 = resized.astype('uint8')\n            if self.transform:\n                img_t = self.transform(img_uint8)\n            else:\n                img_t = torch.tensor(np.transpose(norm, (2,0,1)), dtype=torch.float32)\n            y_t = torch.tensor(label, dtype=torch.long)\n            return img_t, y_t\n        except Exception:\n            # Secours minimal\n            img = cv2.imread(img_path, cv2.IMREAD_COLOR)\n            if img is None: raise\n            img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n            img = _center_square_crop(img)\n            img = cv2.resize(img, (IMG_SIZE, IMG_SIZE), interpolation=cv2.INTER_AREA)\n            if self.transform:\n                img_t = self.transform(img.astype('uint8'))\n            else:\n                img_t = torch.from_numpy(np.transpose(img.astype('float32')/255.0, (2,0,1)))\n            y_t = torch.tensor(label, dtype=torch.long)\n            return img_t, y_t\n\n#  Transforms & Loaders\ntrain_transform = transforms.Compose([\n    transforms.ToPILImage(),\n    transforms.RandomResizedCrop((IMG_SIZE, IMG_SIZE), scale=(0.9, 1.0)),\n    transforms.RandomHorizontalFlip(),\n    transforms.RandomRotation(15),\n    transforms.ToTensor(),\n    transforms.Normalize([0.485,0.456,0.406],[0.229,0.224,0.225]),\n])\n\nval_transform = transforms.Compose([\n    transforms.ToPILImage(),\n    transforms.Resize((IMG_SIZE, IMG_SIZE)),\n    transforms.ToTensor(),\n    transforms.Normalize([0.485,0.456,0.406],[0.229,0.224,0.225]),\n])\n\ntrain_ds = RetinopathyDataset(train_df, transform=train_transform, preprocess_func=PREPROC_FN)\nval_ds   = RetinopathyDataset(val_df,   transform=val_transform,   preprocess_func=PREPROC_FN)\n\nnum_workers = min(4, os.cpu_count() or 2)\ntrain_loader = DataLoader(\n    train_ds, batch_size=BATCH_SIZE, shuffle=True,\n    num_workers=num_workers, pin_memory=True,\n    persistent_workers=True, prefetch_factor=2\n)\nval_loader = DataLoader(\n    val_ds, batch_size=BATCH_SIZE, shuffle=False,\n    num_workers=num_workers, pin_memory=True,\n    persistent_workers=True, prefetch_factor=2\n)\n\nprint(f\"Batches train: {len(train_loader)} | Batches val: {len(val_loader)}\")\n\n#  Modèle\nif BACKBONE == \"resnet50\":\n    model = models.resnet50(weights=models.ResNet50_Weights.IMAGENET1K_V2)\nelse:\n    model = models.resnet18(weights=models.ResNet18_Weights.IMAGENET1K_V1)\n\nin_feats = model.fc.in_features\nmodel.fc = nn.Linear(in_feats, 5)\nmodel = model.to(device)\nmodel = model.to(memory_format=torch.channels_last)   # accélération mémoire\n\n# Freeze \nfor name, p in model.named_parameters():\n    if not name.startswith(\"fc.\"):\n        p.requires_grad = False\n        \nclass_counts = train_df['level'].value_counts().sort_index().values\nclass_weights = (class_counts.sum() / (class_counts + 1e-9))\nclass_weights = class_weights / class_weights.mean()\nclass_weights = torch.tensor(class_weights, dtype=torch.float32, device=device)\n\ncriterion = nn.CrossEntropyLoss(weight=class_weights)\n\noptimizer = SGD(model.parameters(), lr=LR_BASE, momentum=MOMENTUM, nesterov=True, weight_decay=WEIGHT_DECAY)\nscheduler = ReduceLROnPlateau(optimizer, mode='max', factor=0.5, patience=2)\n\nscaler = torch.amp.GradScaler(\"cuda\", enabled=torch.cuda.is_available())\n\n#  pour train/eval\ndef run_one_epoch(dataloader, model, train=True):\n    model.train(train)\n    losses, preds_all, y_all = [], [], []\n    with torch.set_grad_enabled(train):\n        for xb, yb in dataloader:\n            xb = xb.to(device, non_blocking=True).to(memory_format=torch.channels_last)\n            yb = yb.to(device, non_blocking=True)\n\n            if train:\n                optimizer.zero_grad(set_to_none=True)\n                with torch.amp.autocast(\"cuda\", enabled=torch.cuda.is_available()):\n                    logits = model(xb)\n                    loss = criterion(logits, yb)\n                scaler.scale(loss).backward()\n                scaler.step(optimizer)\n                scaler.update()\n            else:\n                with torch.amp.autocast(\"cuda\", enabled=torch.cuda.is_available()):\n                    logits = model(xb)\n                    loss = criterion(logits, yb)\n\n            losses.append(loss.detach().item())\n            preds = torch.argmax(logits, dim=1)\n            preds_all.append(preds.detach().cpu().numpy())\n            y_all.append(yb.detach().cpu().numpy())\n\n    preds_all = np.concatenate(preds_all) if preds_all else np.array([])\n    y_all = np.concatenate(y_all) if y_all else np.array([])\n    kappa = cohen_kappa_score(y_all, preds_all, weights='quadratic') if len(y_all) else 0.0\n    return float(np.mean(losses)), float(kappa)\n\n#  Train loop + Early stop\nbest_kappa = -1.0\nepochs_no_improve = 0\n\nprint(f\"=== Entraînement (freeze backbone) {NUM_EPOCHS} epochs ===\")\nfor epoch in range(1, NUM_EPOCHS + 1):\n    t0 = time.time()\n    train_loss, train_kappa = run_one_epoch(train_loader, model, train=True)\n    val_loss,   val_kappa   = run_one_epoch(val_loader,   model, train=False)\n\n    scheduler.step(val_kappa)\n    lr_now = scheduler.get_last_lr()[0]\n    dt = time.time() - t0\n\n    print(f\"[Epoch {epoch:02d}] \"\n          f\"Train loss={train_loss:.4f} | Train κ={train_kappa:.4f}  ||  \"\n          f\"Val loss={val_loss:.4f} | Val κ={val_kappa:.4f} | LR={lr_now:.6f} | {dt:.1f}s\")\n\n    if val_kappa > best_kappa + 1e-5:\n        best_kappa = val_kappa\n        epochs_no_improve = 0\n        torch.save(model.state_dict(), BEST_PATH)\n        print(f\"✓ Nouveau meilleur modèle (κ={best_kappa:.4f}) → {BEST_PATH}\")\n    else:\n        epochs_no_improve += 1\n        if epochs_no_improve >= PATIENCE:\n            print(\"⏸ Early stopping (pas d'amélioration).\")\n            break\n\n#  Fine-Tuning (dégèl complet)\nprint(\"\\n=== Fine-tuning (dégel complet) ===\")\nfor p in model.parameters():\n    p.requires_grad = True\n\noptimizer = SGD(model.parameters(), lr=LR_FT, momentum=MOMENTUM, nesterov=True, weight_decay=WEIGHT_DECAY)\nscheduler = ReduceLROnPlateau(optimizer, mode='max', factor=0.5, patience=2)\n\nepochs_no_improve = 0\nfor epoch in range(1, FINE_TUNE_EPOCHS + 1):\n    t0 = time.time()\n    train_loss, train_kappa = run_one_epoch(train_loader, model, train=True)\n    val_loss,   val_kappa   = run_one_epoch(val_loader,   model, train=False)\n\n    scheduler.step(val_kappa)\n    lr_now = scheduler.get_last_lr()[0]\n    dt = time.time() - t0\n\n    print(f\"[FT {epoch:02d}] \"\n          f\"Train loss={train_loss:.4f} | Train κ={train_kappa:.4f}  ||  \"\n          f\"Val loss={val_loss:.4f} | Val κ={val_kappa:.4f} | LR={lr_now:.6f} | {dt:.1f}s\")\n\n    if val_kappa > best_kappa + 1e-5:\n        best_kappa = val_kappa\n        epochs_no_improve = 0\n        torch.save(model.state_dict(), BEST_PATH)\n        print(f\"✓ Nouveau meilleur modèle (κ={best_kappa:.4f}) → {BEST_PATH}\")\n    else:\n        epochs_no_improve += 1\n        if epochs_no_improve >= PATIENCE:\n            print(\"⏸ Early stopping FT.\")\n            break\n\nprint(f\"\\nMeilleur κ valid. atteint: {best_kappa:.4f}. Poids: {BEST_PATH}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-31T12:23:57.090044Z","iopub.execute_input":"2025-08-31T12:23:57.090273Z","iopub.status.idle":"2025-08-31T14:41:44.253468Z","shell.execute_reply.started":"2025-08-31T12:23:57.090255Z","shell.execute_reply":"2025-08-31T14:41:44.252277Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Inférence sur les données de test","metadata":{}},{"cell_type":"code","source":"# =========================\n#  Inference — Retinopathy\n# =========================\nimport os, glob, time, cv2, numpy as np, pandas as pd\nimport torch, torch.backends.cudnn as cudnn\nfrom torch import nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import models, transforms\n\n# ---- Adresses adaptées à ton entraînement ----\nWEIGHTS_PATH = \"/kaggle/working/best_model.pth\"   # <- poids sauvegardés par ton training\nOUTPUT_CSV   = \"/kaggle/working/predictions.csv\"  # <- sortie\nBACKBONE     = \"resnet18\"                         # <- comme dans ton training\nIMG_SIZE     = 224                                # <- comme dans ton training\nNUM_CLASSES  = 5\nBATCH_SIZE   = 64\nNUM_WORKERS  = min(4, os.cpu_count() or 2)\nPREPROCESS_MODE = \"fast\"      # \"fast\" (rapide) ou \"robust\" (morpho, plus lent)\nTTA_N = 0                     # 0 pour aller vite; 1..4 pour TTA\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\ntorch.set_float32_matmul_precision(\"high\")\ncudnn.benchmark = True\nprint(\"Device:\", device)\n\n# -------------------------\n#  Auto-détection TEST_DIR\n# -------------------------\ndef find_test_dir(base=\"/kaggle/input\"):\n    candidates = []\n    for root, dirs, files in os.walk(base):\n        # compte d'images dans ce dossier\n        imgs = sum(1 for f in files if f.lower().endswith((\".jpg\",\".jpeg\",\".png\",\".bmp\",\".tif\",\".tiff\")))\n        if imgs > 0:\n            candidates.append((root, imgs))\n    if not candidates:\n        raise FileNotFoundError(\"Aucune image trouvée sous /kaggle/input. Renseigne TEST_DIR manuellement.\")\n    # priorité aux dossiers nommés *test*\n    test_like = [(p,c) for (p,c) in candidates if \"test\" in os.path.basename(p).lower() or \"test\" in p.lower()]\n    if test_like:\n        test_like.sort(key=lambda x: x[1], reverse=True)\n        return test_like[0][0]\n    # sinon, dossier avec le plus d'images\n    candidates.sort(key=lambda x: x[1], reverse=True)\n    return candidates[0][0]\n\nTEST_DIR = find_test_dir()\nprint(\"TEST_DIR choisi:\", TEST_DIR)\n\n# -------------------------\n#  Prétraitements\n# -------------------------\ndef _center_square_crop(img):\n    h, w = img.shape[:2]\n    side = min(h, w)\n    y0 = (h - side) // 2\n    x0 = (w - side) // 2\n    return img[y0:y0+side, x0:x0+side]\n\ndef preprocess_fast(img_path, target_size=(IMG_SIZE, IMG_SIZE)):\n    img = cv2.imread(img_path, cv2.IMREAD_COLOR)\n    if img is None: \n        raise FileNotFoundError(img_path)\n    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    img = _center_square_crop(img)\n    img = cv2.resize(img, target_size, interpolation=cv2.INTER_AREA)\n    return img  # uint8 RGB\n\ndef preprocess_robust(img_path, target_size=(IMG_SIZE, IMG_SIZE), min_area_ratio=0.05):\n    img = cv2.imread(img_path, cv2.IMREAD_COLOR)\n    if img is None: \n        raise FileNotFoundError(img_path)\n    img_rgb = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    gray = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)\n    gray_blur = cv2.GaussianBlur(gray, (0,0), 3)\n    _, th_otsu = cv2.threshold(gray_blur, 0, 255, cv2.THRESH_BINARY+cv2.THRESH_OTSU)\n    mask_gray = (gray > 10).astype(np.uint8) * 255\n    hsv = cv2.cvtColor(img, cv2.COLOR_BGR2HSV); v = hsv[:,:,2]\n    mask_v = (v > 15).astype(np.uint8) * 255\n    mask = cv2.bitwise_or(th_otsu, cv2.bitwise_and(mask_gray, mask_v))\n    mask = cv2.morphologyEx(mask, cv2.MORPH_CLOSE, np.ones((9,9), np.uint8))\n    mask = cv2.morphologyEx(mask, cv2.MORPH_OPEN,  np.ones((5,5), np.uint8))\n    contours, _ = cv2.findContours(mask, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n    cropped = None\n    if contours:\n        cnt = max(contours, key=cv2.contourArea)\n        x,y,w,h = cv2.boundingRect(cnt)\n        H,W = img_rgb.shape[:2]\n        if (w*h) >= (min_area_ratio * H * W):\n            side = int(max(w, h)); cx, cy = x+w//2, y+h//2\n            x0 = max(0, cx - side//2); y0 = max(0, cy - side//2)\n            x1 = min(W, x0 + side);    y1 = min(H, y0 + side)\n            cropped = img_rgb[y0:y1, x0:x1]\n    if cropped is None or cropped.size == 0:\n        cropped = _center_square_crop(img_rgb)\n    return cv2.resize(cropped, target_size, interpolation=cv2.INTER_AREA)\n\ndef preprocess_infer(img_path, mode=\"fast\"):\n    return preprocess_fast(img_path) if mode==\"fast\" else preprocess_robust(img_path)\n\n# -------------------------\n#  Dataset & transforms\n# -------------------------\nclass TestDataset(Dataset):\n    def __init__(self, img_paths, preprocess_mode=\"fast\"):\n        self.paths = img_paths\n        self.mode = preprocess_mode\n    def __len__(self): return len(self.paths)\n    def __getitem__(self, idx):\n        p = self.paths[idx]\n        img = preprocess_infer(p, mode=self.mode)  # uint8 RGB\n        return img, os.path.basename(p)\n\nbase_transform = transforms.Compose([\n    transforms.ToPILImage(),\n    transforms.Resize((IMG_SIZE, IMG_SIZE)),\n    transforms.ToTensor(),\n    transforms.Normalize([0.485,0.456,0.406],[0.229,0.224,0.225]),\n])\n\n\ntta_list = [\n    transforms.Compose([\n        transforms.ToPILImage(),\n        transforms.RandomHorizontalFlip(p=1.0),\n        transforms.Resize((IMG_SIZE, IMG_SIZE)),\n        transforms.ToTensor(),\n        transforms.Normalize([0.485,0.456,0.406],[0.229,0.224,0.225]),\n    ]),\n    transforms.Compose([\n        transforms.ToPILImage(),\n        transforms.RandomRotation(15),  # couvre [-15°, +15°]\n        transforms.Resize((IMG_SIZE, IMG_SIZE)),\n        transforms.ToTensor(),\n        transforms.Normalize([0.485,0.456,0.406],[0.229,0.224,0.225]),\n    ]),\n    transforms.Compose([\n        transforms.ToPILImage(),\n        transforms.RandomRotation((-15, -15)),  # exactement -15°\n        transforms.Resize((IMG_SIZE, IMG_SIZE)),\n        transforms.ToTensor(),\n        transforms.Normalize([0.485,0.456,0.406],[0.229,0.224,0.225]),\n    ]),\n    transforms.Compose([\n        transforms.ToPILImage(),\n        transforms.RandomVerticalFlip(p=1.0),\n        transforms.Resize((IMG_SIZE, IMG_SIZE)),\n        transforms.ToTensor(),\n        transforms.Normalize([0.485,0.456,0.406],[0.229,0.224,0.225]),\n    ]),\n]\n\n\n# -------------------------\n#  Modèle\n# -------------------------\ndef build_model(backbone=\"resnet18\", num_classes=5):\n    if backbone == \"resnet50\":\n        m = models.resnet50(weights=None)\n    else:\n        m = models.resnet18(weights=None)\n    in_feats = m.fc.in_features\n    m.fc = nn.Linear(in_feats, num_classes)\n    return m\n\nassert os.path.exists(WEIGHTS_PATH), f\"Poids introuvables: {WEIGHTS_PATH}\"\nmodel = build_model(BACKBONE, NUM_CLASSES).to(device)\nstate = torch.load(WEIGHTS_PATH, map_location=\"cpu\")\nmodel.load_state_dict(state, strict=True)\nmodel.eval()\nmodel = model.to(memory_format=torch.channels_last)\nprint(\"Modèle chargé:\", BACKBONE)\n\n# -------------------------\n#  Fichiers test & loader\n# -------------------------\nimg_paths = []\nfor ext in (\"*.jpg\",\"*.jpeg\",\"*.png\",\"*.bmp\",\"*.tif\",\"*.tiff\"):\n    img_paths.extend(glob.glob(os.path.join(TEST_DIR, \"**\", ext), recursive=True))\nimg_paths = sorted(img_paths)\nprint(\"Images trouvées:\", len(img_paths))\nassert len(img_paths) > 0, \"Aucune image dans TEST_DIR.\"\n\nds = TestDataset(img_paths, preprocess_mode=PREPROCESS_MODE)\nloader = DataLoader(ds, batch_size=BATCH_SIZE, shuffle=False,\n                    num_workers=NUM_WORKERS, pin_memory=True,\n                    persistent_workers=(NUM_WORKERS>0))\n\n# -------------------------\n#  Inférence + TTA\n# -------------------------\nall_names, all_preds = [], []\nt0 = time.time()\nwith torch.no_grad():\n    for imgs_uint8, names in loader:\n        # base pass\n        batch_base = []\n        for arr in imgs_uint8:\n            if torch.is_tensor(arr): arr = arr.cpu().numpy()\n            x = base_transform(arr).unsqueeze(0)\n            batch_base.append(x)\n        xb = torch.cat(batch_base, dim=0).to(device).to(memory_format=torch.channels_last)\n        with torch.amp.autocast(\"cuda\", enabled=torch.cuda.is_available()):\n            logits = model(xb)\n        logits_sum = logits\n\n        # TTA passes\n        TTA_USE = min(TTA_N, len(tta_list))\n        for i in range(TTA_USE):\n            batch_tta = []\n            trans = tta_list[i]\n            for arr in imgs_uint8:\n                if torch.is_tensor(arr): arr = arr.cpu().numpy()\n                xt = trans(arr).unsqueeze(0)\n                batch_tta.append(xt)\n            x_tta = torch.cat(batch_tta, dim=0).to(device).to(memory_format=torch.channels_last)\n            with torch.amp.autocast(\"cuda\", enabled=torch.cuda.is_available()):\n                logits_t = model(x_tta)\n            logits_sum = logits_sum + logits_t\n\n        probs = torch.softmax(logits_sum / (1 + TTA_USE), dim=1)\n        preds = torch.argmax(probs, dim=1).detach().cpu().numpy()\n        all_names.extend(list(names))\n        all_preds.extend(list(preds))\n\ndt = time.time() - t0\nprint(f\"Inférence terminée en {dt/60:.1f} min pour {len(all_names)} images.\")\n\n# -------------------------\n#  Export CSV\n# -------------------------\nsub = pd.DataFrame({\"filename\": all_names, \"level\": all_preds})\nsub.to_csv(OUTPUT_CSV, index=False)\nprint(\"CSV écrit ->\", OUTPUT_CSV)\nprint(sub.head())\nprint(\"Distribution classes préd.:\")\nprint(sub[\"level\"].value_counts().sort_index().to_string())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-31T14:41:44.260094Z","iopub.execute_input":"2025-08-31T14:41:44.260412Z","iopub.status.idle":"2025-08-31T15:09:23.154404Z","shell.execute_reply.started":"2025-08-31T14:41:44.260370Z","shell.execute_reply":"2025-08-31T15:09:23.153451Z"}},"outputs":[],"execution_count":null}]}