{"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":"import torch\n\nprint(\"PyTorch :\", torch.__version__)\nprint(\"CUDA    :\", torch.version.cuda)\nprint(\"GPU 0   :\", torch.cuda.get_device_name(0))\nprint(\"Nombre de GPU :\", torch.cuda.device_count())\n\nfor i in range(torch.cuda.device_count()):\n    print(\n        f\"GPU {i} :\",\n        torch.cuda.get_device_name(i),\n        \"- capacité\",\n        torch.cuda.get_device_capability(i)\n    )\n\n# Test réel\nfor i in range(torch.cuda.device_count()):\n    x = torch.randn(2, 3, 64, 64, device=f\"cuda:{i}\")\n    conv = torch.nn.Conv2d(3, 8, 3, padding=1).to(f\"cuda:{i}\")\n    y = conv(x)\n    print(f\"✅ GPU {i} opérationnel :\", y.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T13:58:08.969809Z","iopub.execute_input":"2026-08-11T13:58:08.970199Z","iopub.status.idle":"2026-08-11T13:58:14.275734Z","shell.execute_reply.started":"2026-08-11T13:58:08.970168Z","shell.execute_reply":"2026-08-11T13:58:14.275027Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# FINAL PIPELINE — BLOC 1\n# SPLIT PATIENT + CACHES 512\n# Compatible T4 x2\n# ============================================================\n\nimport os\nimport cv2\nimport time\nimport shutil\nimport random\nimport warnings\n\nfrom pathlib import Path\nfrom concurrent.futures import ThreadPoolExecutor\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\nfrom tqdm.auto import tqdm\nfrom sklearn.model_selection import train_test_split\n\n\n# ============================================================\n# 1. CONFIG\n# ============================================================\n\nwarnings.filterwarnings(\"ignore\")\n\nSEED = 42\nIMAGE_SIZE = 512\nJPEG_QUALITY = 92\n\nrandom.seed(SEED)\nnp.random.seed(SEED)\n\ncv2.setNumThreads(1)\n\n\n# ============================================================\n# 2. PATHS\n# ============================================================\n\nCOMPETITION_DIR = Path(\n    \"/kaggle/input/competitions/\"\n    \"diabetic-retinopathy-detection\"\n)\n\nTRAIN_IMAGE_DIR = Path(\n    \"/kaggle/input/datasets/\"\n    \"josephrynkiewicz/\"\n    \"diabetic-retinopathy-train-unzipped/\"\n    \"train\"\n)\n\nTRAIN_LABELS_PATH = (\n    COMPETITION_DIR\n    / \"trainLabels.csv.zip\"\n)\n\nWORK_DIR = Path(\n    \"/kaggle/working/dr_final\"\n)\n\nNATURAL_DIR = (\n    WORK_DIR / \"natural_512\"\n)\n\nGRAHAM_DIR = (\n    WORK_DIR / \"graham_512\"\n)\n\nCHECKPOINT_DIR = (\n    WORK_DIR / \"checkpoints\"\n)\n\nfor directory in [\n    WORK_DIR,\n    NATURAL_DIR,\n    GRAHAM_DIR,\n    CHECKPOINT_DIR\n]:\n    directory.mkdir(\n        parents=True,\n        exist_ok=True\n    )\n\n\n# ============================================================\n# 3. VERIFY INPUT\n# ============================================================\n\nassert TRAIN_LABELS_PATH.exists(), (\n    TRAIN_LABELS_PATH\n)\n\nassert TRAIN_IMAGE_DIR.exists(), (\n    TRAIN_IMAGE_DIR\n)\n\ndf = pd.read_csv(\n    TRAIN_LABELS_PATH\n)\n\nassert len(df) == 35126\n\n\n# ============================================================\n# 4. PATIENT / EYE\n# ============================================================\n\ndf[\"patient_id\"] = (\n    df[\"image\"]\n    .str.rsplit(\"_\", n=1)\n    .str[0]\n)\n\ndf[\"eye\"] = (\n    df[\"image\"]\n    .str.rsplit(\"_\", n=1)\n    .str[1]\n)\n\nprint(\"=\" * 65)\nprint(\"DATASET\")\nprint(\"=\" * 65)\n\nprint(\n    \"Images   :\",\n    len(df)\n)\n\nprint(\n    \"Patients :\",\n    df[\"patient_id\"].nunique()\n)\n\nprint()\n\nprint(\n    df[\"level\"]\n    .value_counts()\n    .sort_index()\n)\n\n\n# ============================================================\n# 5. HOLD-OUT PATIENT 80/20\n#\n# On revient au split 20 % :\n# - plus stable pour mesurer les modèles\n# - calibration seuils + évaluation indépendantes\n# - aucun patient partagé\n# ============================================================\n\npatient_df = (\n    df.groupby(\n        \"patient_id\",\n        as_index=False\n    )[\"level\"]\n    .max()\n    .rename(\n        columns={\n            \"level\":\n            \"patient_level\"\n        }\n    )\n)\n\n\ntrain_patients, val_patients = (\n    train_test_split(\n\n        patient_df,\n\n        test_size=0.20,\n\n        random_state=SEED,\n\n        stratify=patient_df[\n            \"patient_level\"\n        ]\n    )\n)\n\n\nTRAIN_PATIENTS = set(\n    train_patients[\n        \"patient_id\"\n    ]\n)\n\nVAL_PATIENTS = set(\n    val_patients[\n        \"patient_id\"\n    ]\n)\n\n\nassert TRAIN_PATIENTS.isdisjoint(\n    VAL_PATIENTS\n)\n\n\ntrain_df = (\n    df[\n        df[\"patient_id\"]\n        .isin(TRAIN_PATIENTS)\n    ]\n    .copy()\n    .reset_index(drop=True)\n)\n\n\nval_df = (\n    df[\n        df[\"patient_id\"]\n        .isin(VAL_PATIENTS)\n    ]\n    .copy()\n    .reset_index(drop=True)\n)\n\n\nprint(\"\\n\" + \"=\" * 65)\nprint(\"SPLIT PATIENT\")\nprint(\"=\" * 65)\n\nprint(\n    \"Patients train :\",\n    len(TRAIN_PATIENTS)\n)\n\nprint(\n    \"Patients val   :\",\n    len(VAL_PATIENTS)\n)\n\nprint(\n    \"Images train   :\",\n    len(train_df)\n)\n\nprint(\n    \"Images val     :\",\n    len(val_df)\n)\n\nprint(\n    \"Patients communs :\",\n    len(\n        TRAIN_PATIENTS\n        &\n        VAL_PATIENTS\n    )\n)\n\n\n# ============================================================\n# 6. DISTRIBUTION\n# ============================================================\n\ndistribution = pd.DataFrame({\n\n    \"Total\":\n        df[\"level\"]\n        .value_counts(\n            normalize=True\n        )\n        .sort_index(),\n\n    \"Train\":\n        train_df[\"level\"]\n        .value_counts(\n            normalize=True\n        )\n        .sort_index(),\n\n    \"Validation\":\n        val_df[\"level\"]\n        .value_counts(\n            normalize=True\n        )\n        .sort_index()\n\n}) * 100\n\n\nprint()\ndisplay(\n    distribution.round(2)\n)\n\n\n# ============================================================\n# 7. SAVE SPLITS\n# ============================================================\n\ntrain_df.to_csv(\n    WORK_DIR\n    / \"train_split.csv\",\n    index=False\n)\n\nval_df.to_csv(\n    WORK_DIR\n    / \"val_split.csv\",\n    index=False\n)\n\n\n# ============================================================\n# 8. CROP BLACK BORDER\n# ============================================================\n\ndef crop_retina(image):\n\n    intensity = (\n        image.max(axis=2)\n    )\n\n    mask = (\n        intensity > 7\n    ).astype(\n        np.uint8\n    )\n\n\n    if mask.sum() == 0:\n\n        return image\n\n\n    kernel = np.ones(\n        (5, 5),\n        np.uint8\n    )\n\n\n    mask = cv2.morphologyEx(\n        mask,\n        cv2.MORPH_CLOSE,\n        kernel\n    )\n\n\n    number_labels, labels, stats, _ = (\n        cv2.connectedComponentsWithStats(\n            mask,\n            connectivity=8\n        )\n    )\n\n\n    if number_labels <= 1:\n\n        return image\n\n\n    component = (\n        1\n        +\n        np.argmax(\n            stats[\n                1:,\n                cv2.CC_STAT_AREA\n            ]\n        )\n    )\n\n\n    x = stats[\n        component,\n        cv2.CC_STAT_LEFT\n    ]\n\n    y = stats[\n        component,\n        cv2.CC_STAT_TOP\n    ]\n\n    width = stats[\n        component,\n        cv2.CC_STAT_WIDTH\n    ]\n\n    height = stats[\n        component,\n        cv2.CC_STAT_HEIGHT\n    ]\n\n\n    margin_x = int(\n        width * 0.015\n    )\n\n    margin_y = int(\n        height * 0.015\n    )\n\n\n    x1 = max(\n        0,\n        x - margin_x\n    )\n\n    y1 = max(\n        0,\n        y - margin_y\n    )\n\n    x2 = min(\n        image.shape[1],\n        x + width + margin_x\n    )\n\n    y2 = min(\n        image.shape[0],\n        y + height + margin_y\n    )\n\n\n    crop = image[\n        y1:y2,\n        x1:x2\n    ]\n\n\n    if crop.size == 0:\n\n        return image\n\n\n    return crop\n\n\n# ============================================================\n# 9. PAD SQUARE\n# ============================================================\n\ndef pad_square(image):\n\n    height, width = (\n        image.shape[:2]\n    )\n\n    side = max(\n        height,\n        width\n    )\n\n\n    result = np.zeros(\n        (\n            side,\n            side,\n            3\n        ),\n        dtype=np.uint8\n    )\n\n\n    top = (\n        side - height\n    ) // 2\n\n    left = (\n        side - width\n    ) // 2\n\n\n    result[\n        top:top + height,\n        left:left + width\n    ] = image\n\n\n    return result\n\n\n# ============================================================\n# 10. NATURAL 512\n# ============================================================\n\ndef create_natural(image):\n\n    image = crop_retina(\n        image\n    )\n\n    image = pad_square(\n        image\n    )\n\n    image = cv2.resize(\n\n        image,\n\n        (\n            IMAGE_SIZE,\n            IMAGE_SIZE\n        ),\n\n        interpolation=cv2.INTER_AREA\n    )\n\n    return image\n\n\n# ============================================================\n# 11. BEN-GRAHAM STYLE\n# ============================================================\n\ndef create_graham(\n    image\n):\n\n    gray = cv2.cvtColor(\n        image,\n        cv2.COLOR_BGR2GRAY\n    )\n\n\n    retina_mask = (\n        gray > 5\n    )\n\n\n    blurred = cv2.GaussianBlur(\n\n        image,\n\n        (0, 0),\n\n        sigmaX=IMAGE_SIZE / 30\n    )\n\n\n    processed = cv2.addWeighted(\n\n        image,\n        4.0,\n\n        blurred,\n        -4.0,\n\n        128\n    )\n\n\n    processed = np.clip(\n        processed,\n        0,\n        255\n    ).astype(\n        np.uint8\n    )\n\n\n    processed[\n        ~retina_mask\n    ] = 0\n\n\n    return processed\n\n\n# ============================================================\n# 12. PROCESS ONE IMAGE\n# ============================================================\n\ndef process_image(\n    image_id\n):\n\n    source = (\n        TRAIN_IMAGE_DIR\n        / f\"{image_id}.jpeg\"\n    )\n\n\n    natural_path = (\n        NATURAL_DIR\n        / f\"{image_id}.jpg\"\n    )\n\n\n    graham_path = (\n        GRAHAM_DIR\n        / f\"{image_id}.jpg\"\n    )\n\n\n    # Resume capable\n    if (\n        natural_path.exists()\n        and\n        graham_path.exists()\n    ):\n\n        return True\n\n\n    image = cv2.imread(\n        str(source),\n        cv2.IMREAD_COLOR\n    )\n\n\n    if image is None:\n\n        return False\n\n\n    natural = create_natural(\n        image\n    )\n\n\n    graham = create_graham(\n        natural\n    )\n\n\n    ok1 = cv2.imwrite(\n\n        str(\n            natural_path\n        ),\n\n        natural,\n\n        [\n            cv2.IMWRITE_JPEG_QUALITY,\n            JPEG_QUALITY\n        ]\n    )\n\n\n    ok2 = cv2.imwrite(\n\n        str(\n            graham_path\n        ),\n\n        graham,\n\n        [\n            cv2.IMWRITE_JPEG_QUALITY,\n            JPEG_QUALITY\n        ]\n    )\n\n\n    return (\n        bool(ok1)\n        and\n        bool(ok2)\n    )\n\n\n# ============================================================\n# 13. BUILD CACHE\n# ============================================================\n\nimage_ids = (\n    df[\"image\"]\n    .astype(str)\n    .tolist()\n)\n\n\nworkers = min(\n    6,\n    os.cpu_count() or 4\n)\n\n\nprint(\"\\n\" + \"=\" * 65)\nprint(\"CREATION CACHE 512\")\nprint(\"=\" * 65)\n\nprint(\n    \"Workers :\",\n    workers\n)\n\n\nstart_time = time.time()\n\n\nwith ThreadPoolExecutor(\n    max_workers=workers\n) as executor:\n\n    results = list(\n\n        tqdm(\n\n            executor.map(\n                process_image,\n                image_ids\n            ),\n\n            total=len(\n                image_ids\n            ),\n\n            desc=(\n                \"Natural + Graham\"\n            )\n        )\n    )\n\n\nduration = (\n    time.time()\n    -\n    start_time\n) / 60\n\n\n# ============================================================\n# 14. VERIFY\n# ============================================================\n\nnatural_files = list(\n    NATURAL_DIR.glob(\n        \"*.jpg\"\n    )\n)\n\ngraham_files = list(\n    GRAHAM_DIR.glob(\n        \"*.jpg\"\n    )\n)\n\n\nfailed = (\n    len(results)\n    -\n    sum(results)\n)\n\n\nprint(\"\\n\" + \"=\" * 65)\nprint(\"CACHE RESULT\")\nprint(\"=\" * 65)\n\nprint(\n    \"Natural :\",\n    len(natural_files)\n)\n\nprint(\n    \"Graham  :\",\n    len(graham_files)\n)\n\nprint(\n    \"Failed  :\",\n    failed\n)\n\nprint(\n    \"Duration:\",\n    round(\n        duration,\n        1\n    ),\n    \"minutes\"\n)\n\n\nassert (\n    len(natural_files)\n    ==\n    35126\n)\n\nassert (\n    len(graham_files)\n    ==\n    35126\n)\n\nassert failed == 0\n\n\n# ============================================================\n# 15. DISK SIZE\n# ============================================================\n\nnatural_gb = (\n    sum(\n        file.stat().st_size\n        for file in natural_files\n    )\n    /\n    1024**3\n)\n\n\ngraham_gb = (\n    sum(\n        file.stat().st_size\n        for file in graham_files\n    )\n    /\n    1024**3\n)\n\n\nprint(\n    \"Natural size :\",\n    round(\n        natural_gb,\n        2\n    ),\n    \"GB\"\n)\n\nprint(\n    \"Graham size  :\",\n    round(\n        graham_gb,\n        2\n    ),\n    \"GB\"\n)\n\n\n# ============================================================\n# 16. VISUAL CHECK\n# ============================================================\n\nfig, axes = plt.subplots(\n    5,\n    2,\n    figsize=(10, 22)\n)\n\n\nfor level in range(5):\n\n    sample = (\n        df[\n            df[\"level\"]\n            ==\n            level\n        ]\n        .sample(\n            1,\n            random_state=SEED + level\n        )\n        .iloc[0]\n    )\n\n\n    image_id = sample[\n        \"image\"\n    ]\n\n\n    natural = cv2.imread(\n        str(\n            NATURAL_DIR\n            /\n            f\"{image_id}.jpg\"\n        )\n    )\n\n\n    graham = cv2.imread(\n        str(\n            GRAHAM_DIR\n            /\n            f\"{image_id}.jpg\"\n        )\n    )\n\n\n    natural = cv2.cvtColor(\n        natural,\n        cv2.COLOR_BGR2RGB\n    )\n\n\n    graham = cv2.cvtColor(\n        graham,\n        cv2.COLOR_BGR2RGB\n    )\n\n\n    axes[\n        level,\n        0\n    ].imshow(\n        natural\n    )\n\n\n    axes[\n        level,\n        0\n    ].set_title(\n        f\"Classe {level} - Natural\"\n    )\n\n\n    axes[\n        level,\n        1\n    ].imshow(\n        graham\n    )\n\n\n    axes[\n        level,\n        1\n    ].set_title(\n        f\"Classe {level} - Graham\"\n    )\n\n\n    axes[\n        level,\n        0\n    ].axis(\"off\")\n\n    axes[\n        level,\n        1\n    ].axis(\"off\")\n\n\nplt.tight_layout()\nplt.show()\n\n\n# ============================================================\n# 17. FINAL\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 65)\nprint(\"✅ BLOC 1 TERMINÉ\")\nprint(\"=\" * 65)\n\nprint(\n    \"Train images :\",\n    len(train_df)\n)\n\nprint(\n    \"Val images   :\",\n    len(val_df)\n)\n\nprint(\n    \"Natural      :\",\n    len(natural_files)\n)\n\nprint(\n    \"Graham       :\",\n    len(graham_files)\n)\n\nprint(\n    \"Working dir  :\",\n    WORK_DIR\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T14:00:47.873321Z","iopub.execute_input":"2026-08-11T14:00:47.874005Z","iopub.status.idle":"2026-08-11T16:17:14.071352Z","shell.execute_reply.started":"2026-08-11T14:00:47.873972Z","shell.execute_reply":"2026-08-11T16:17:14.070437Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# FINAL MODEL A — EFFICIENTNET-B4\n# DDP 2 x Tesla T4\n# 384 -> 512\n# Regression MSE\n# Dynamic weighted sampling\n# SyncBatchNorm\n# AMP sécurisé\n# QWK calibration/evaluation\n# TTA + eye fusion final\n# ============================================================\n\nfrom pathlib import Path\nimport textwrap\n\nSCRIPT_PATH = Path(\n    \"/kaggle/working/train_effb4_ddp.py\"\n)\n\nscript = r'''\nimport os\nimport gc\nimport cv2\nimport math\nimport time\nimport json\nimport random\nimport warnings\nimport zipfile\n\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\n\nimport torch\nimport torch.nn as nn\nimport torch.distributed as dist\n\nfrom torch.nn.parallel import DistributedDataParallel as DDP\n\nfrom torch.utils.data import (\n    Dataset,\n    DataLoader,\n    Sampler\n)\n\nfrom torchvision import models\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import cohen_kappa_score\n\nfrom scipy.optimize import minimize\n\nfrom tqdm import tqdm\n\n\n# ============================================================\n# 1. DDP\n# ============================================================\n\nwarnings.filterwarnings(\"ignore\")\n\ndist.init_process_group(\n    backend=\"nccl\"\n)\n\nLOCAL_RANK = int(\n    os.environ[\"LOCAL_RANK\"]\n)\n\nRANK = dist.get_rank()\nWORLD_SIZE = dist.get_world_size()\n\ntorch.cuda.set_device(\n    LOCAL_RANK\n)\n\nDEVICE = torch.device(\n    f\"cuda:{LOCAL_RANK}\"\n)\n\n\n# ============================================================\n# 2. SEEDS\n# ============================================================\n\nSEED = 42\n\nrandom.seed(\n    SEED + RANK\n)\n\nnp.random.seed(\n    SEED + RANK\n)\n\ntorch.manual_seed(\n    SEED + RANK\n)\n\ntorch.cuda.manual_seed_all(\n    SEED + RANK\n)\n\ntorch.backends.cudnn.benchmark = True\n\ncv2.setNumThreads(1)\n\n\n# ============================================================\n# 3. CONFIG\n# ============================================================\n\nWORK_DIR = Path(\n    \"/kaggle/working/dr_final\"\n)\n\nGRAHAM_DIR = (\n    WORK_DIR / \"graham_512\"\n)\n\nCHECKPOINT_DIR = (\n    WORK_DIR / \"checkpoints\"\n)\n\nCHECKPOINT_DIR.mkdir(\n    parents=True,\n    exist_ok=True\n)\n\nTRAIN_CSV = (\n    WORK_DIR / \"train_split.csv\"\n)\n\nVAL_CSV = (\n    WORK_DIR / \"val_split.csv\"\n)\n\nBEST_MODEL_PATH = (\n    CHECKPOINT_DIR\n    / \"efficientnet_b4_best.pth\"\n)\n\nLATEST_MODEL_PATH = (\n    CHECKPOINT_DIR\n    / \"efficientnet_b4_latest.pth\"\n)\n\nHISTORY_PATH = (\n    WORK_DIR\n    / \"efficientnet_b4_history.csv\"\n)\n\nFINAL_TTA_PATH = (\n    WORK_DIR\n    / \"efficientnet_b4_tta_validation.csv\"\n)\n\n\n# ============================================================\n# 4. PHASES\n#\n# 2 epochs 384\n# 5 epochs 512\n# ============================================================\n\nPHASES = [\n\n    {\n        \"name\":\n            \"384\",\n\n        \"image_size\":\n            384,\n\n        \"epochs\":\n            2,\n\n        \"batch_size\":\n            12,\n\n        \"backbone_lr\":\n            8e-5,\n\n        \"head_lr\":\n            4e-4,\n\n        \"sampling_alphas\":\n            [\n                0.35,\n                0.28\n            ]\n    },\n\n    {\n        \"name\":\n            \"512\",\n\n        \"image_size\":\n            512,\n\n        \"epochs\":\n            5,\n\n        \"batch_size\":\n            8,\n\n        \"backbone_lr\":\n            3e-5,\n\n        \"head_lr\":\n            1.2e-4,\n\n        \"sampling_alphas\":\n            [\n                0.26,\n                0.20,\n                0.14,\n                0.08,\n                0.04\n            ]\n    }\n]\n\n\n# ============================================================\n# 5. LOAD DATA\n# ============================================================\n\ntrain_df = pd.read_csv(\n    TRAIN_CSV\n)\n\nval_df = pd.read_csv(\n    VAL_CSV\n)\n\ntrain_df[\n    \"patient_id\"\n] = (\n    train_df[\n        \"patient_id\"\n    ].astype(str)\n)\n\nval_df[\n    \"patient_id\"\n] = (\n    val_df[\n        \"patient_id\"\n    ].astype(str)\n)\n\n\nassert len(train_df) == 28100\nassert len(val_df) == 7026\n\n\n# ============================================================\n# 6. CALIBRATION / EVALUATION\n#\n# moitié des patients validation pour thresholds\n# moitié totalement indépendante pour QWK\n# ============================================================\n\nval_patient_df = (\n\n    val_df\n    .groupby(\n        \"patient_id\",\n        as_index=False\n    )[\"level\"]\n    .max()\n\n    .rename(\n        columns={\n            \"level\":\n            \"patient_level\"\n        }\n    )\n)\n\n\ncalibration_patients, evaluation_patients = (\n    train_test_split(\n\n        val_patient_df,\n\n        test_size=0.50,\n\n        random_state=SEED,\n\n        stratify=val_patient_df[\n            \"patient_level\"\n        ]\n    )\n)\n\n\nCALIBRATION_IDS = set(\n    calibration_patients[\n        \"patient_id\"\n    ].astype(str)\n)\n\n\nEVALUATION_IDS = set(\n    evaluation_patients[\n        \"patient_id\"\n    ].astype(str)\n)\n\n\nassert CALIBRATION_IDS.isdisjoint(\n    EVALUATION_IDS\n)\n\n\n# ============================================================\n# 7. PRINT ENV\n# ============================================================\n\nif RANK == 0:\n\n    print(\"=\" * 72)\n    print(\"FINAL MODEL A — EFFICIENTNET-B4\")\n    print(\"=\" * 72)\n\n    print(\n        \"PyTorch      :\",\n        torch.__version__\n    )\n\n    print(\n        \"CUDA         :\",\n        torch.version.cuda\n    )\n\n    print(\n        \"Nombre GPU   :\",\n        WORLD_SIZE\n    )\n\n    for i in range(\n        torch.cuda.device_count()\n    ):\n\n        print(\n            f\"GPU {i}        :\",\n            torch.cuda.get_device_name(i)\n        )\n\n    print(\n        \"Train images :\",\n        len(train_df)\n    )\n\n    print(\n        \"Val images   :\",\n        len(val_df)\n    )\n\n    print(\n        \"Patients calibration :\",\n        len(CALIBRATION_IDS)\n    )\n\n    print(\n        \"Patients evaluation  :\",\n        len(EVALUATION_IDS)\n    )\n\n\n# ============================================================\n# 8. AUGMENTATIONS\n# ============================================================\n\nIMAGENET_MEAN = (\n    0.485,\n    0.456,\n    0.406\n)\n\nIMAGENET_STD = (\n    0.229,\n    0.224,\n    0.225\n)\n\n\ndef train_transform(\n    image_size\n):\n\n    return A.Compose([\n\n        A.Resize(\n            image_size,\n            image_size\n        ),\n\n        A.HorizontalFlip(\n            p=0.50\n        ),\n\n        A.VerticalFlip(\n            p=0.50\n        ),\n\n        A.Affine(\n\n            scale=(\n                0.90,\n                1.10\n            ),\n\n            translate_percent=(\n                -0.05,\n                0.05\n            ),\n\n            rotate=(\n                -180,\n                180\n            ),\n\n            shear=(\n                -5,\n                5\n            ),\n\n            border_mode=(\n                cv2.BORDER_CONSTANT\n            ),\n\n            fill=0,\n\n            p=0.85\n        ),\n\n        A.OneOf([\n\n            A.RandomBrightnessContrast(\n\n                brightness_limit=0.10,\n                contrast_limit=0.10,\n\n                p=1.0\n            ),\n\n            A.RandomGamma(\n\n                gamma_limit=(\n                    88,\n                    112\n                ),\n\n                p=1.0\n            )\n\n        ], p=0.25),\n\n        A.OneOf([\n\n            A.GaussianBlur(\n\n                blur_limit=(\n                    3,\n                    5\n                ),\n\n                p=1.0\n            ),\n\n            A.GaussNoise(\n\n                std_range=(\n                    0.005,\n                    0.020\n                ),\n\n                p=1.0\n            )\n\n        ], p=0.05),\n\n        A.Normalize(\n\n            mean=IMAGENET_MEAN,\n            std=IMAGENET_STD\n        ),\n\n        ToTensorV2()\n    ])\n\n\ndef val_transform(\n    image_size\n):\n\n    return A.Compose([\n\n        A.Resize(\n            image_size,\n            image_size\n        ),\n\n        A.Normalize(\n\n            mean=IMAGENET_MEAN,\n            std=IMAGENET_STD\n        ),\n\n        ToTensorV2()\n    ])\n\n\n# ============================================================\n# 9. DATASET\n# ============================================================\n\nclass DRDataset(\n    Dataset\n):\n\n    def __init__(\n        self,\n        dataframe,\n        transform,\n        validation=False\n    ):\n\n        self.df = (\n            dataframe\n            .reset_index(\n                drop=True\n            )\n        )\n\n        self.transform = (\n            transform\n        )\n\n        self.validation = (\n            validation\n        )\n\n\n    def __len__(self):\n\n        return len(\n            self.df\n        )\n\n\n    def __getitem__(\n        self,\n        idx\n    ):\n\n        row = self.df.iloc[\n            idx\n        ]\n\n        image_id = str(\n            row[\"image\"]\n        )\n\n        image_path = (\n            GRAHAM_DIR\n            /\n            f\"{image_id}.jpg\"\n        )\n\n        image = cv2.imread(\n            str(image_path),\n            cv2.IMREAD_COLOR\n        )\n\n        if image is None:\n\n            raise RuntimeError(\n                f\"Image introuvable : \"\n                f\"{image_path}\"\n            )\n\n        image = cv2.cvtColor(\n            image,\n            cv2.COLOR_BGR2RGB\n        )\n\n        image = self.transform(\n            image=image\n        )[\"image\"]\n\n        target = torch.tensor(\n            float(\n                row[\"level\"]\n            ),\n            dtype=torch.float32\n        )\n\n        if not self.validation:\n\n            return (\n                image,\n                target\n            )\n\n        return {\n\n            \"image\":\n                image,\n\n            \"target\":\n                target,\n\n            \"image_id\":\n                image_id,\n\n            \"patient_id\":\n                str(\n                    row[\"patient_id\"]\n                ),\n\n            \"eye\":\n                str(\n                    row[\"eye\"]\n                )\n        }\n\n\n# ============================================================\n# 10. DISTRIBUTED WEIGHTED SAMPLER\n# ============================================================\n\nclass DistributedWeightedSampler(\n    Sampler\n):\n\n    def __init__(\n        self,\n        labels,\n        alpha,\n        num_replicas,\n        rank,\n        seed=42\n    ):\n\n        self.labels = np.asarray(\n            labels,\n            dtype=np.int64\n        )\n\n        self.alpha = float(\n            alpha\n        )\n\n        self.num_replicas = (\n            num_replicas\n        )\n\n        self.rank = rank\n\n        self.seed = seed\n\n        self.epoch = 0\n\n        self.num_samples = (\n            int(\n                math.ceil(\n                    len(self.labels)\n                    /\n                    self.num_replicas\n                )\n            )\n        )\n\n        self.total_size = (\n            self.num_samples\n            *\n            self.num_replicas\n        )\n\n\n    def set_epoch(\n        self,\n        epoch\n    ):\n\n        self.epoch = epoch\n\n\n    def __len__(self):\n\n        return self.num_samples\n\n\n    def __iter__(\n        self\n    ):\n\n        counts = np.bincount(\n            self.labels,\n            minlength=5\n        )\n\n        class_weights = (\n            len(self.labels)\n            /\n            counts\n        ) ** self.alpha\n\n        sample_weights = (\n            class_weights[\n                self.labels\n            ]\n        )\n\n        sample_weights = torch.tensor(\n            sample_weights,\n            dtype=torch.double\n        )\n\n        generator = torch.Generator()\n\n        generator.manual_seed(\n            self.seed\n            +\n            self.epoch\n        )\n\n        global_indices = (\n            torch.multinomial(\n\n                sample_weights,\n\n                self.total_size,\n\n                replacement=True,\n\n                generator=generator\n            )\n            .tolist()\n        )\n\n        indices = global_indices[\n            self.rank:\n            self.total_size:\n            self.num_replicas\n        ]\n\n        return iter(\n            indices\n        )\n\n\n# ============================================================\n# 11. PRETRAINED MODEL\n#\n# Rank 0 downloads once.\n# ============================================================\n\nif RANK == 0:\n\n    _tmp = models.efficientnet_b4(\n\n        weights=(\n            models\n            .EfficientNet_B4_Weights\n            .DEFAULT\n        )\n    )\n\n    del _tmp\n\ndist.barrier()\n\n\ndef create_model():\n\n    network = models.efficientnet_b4(\n\n        weights=(\n            models\n            .EfficientNet_B4_Weights\n            .DEFAULT\n        )\n    )\n\n    features = (\n        network\n        .classifier[1]\n        .in_features\n    )\n\n    network.classifier = nn.Sequential(\n\n        nn.Dropout(\n            p=0.40\n        ),\n\n        nn.Linear(\n            features,\n            1\n        )\n    )\n\n    return network\n\n\nmodel = create_model()\n\n\n# Sync BN :\n# batch 8/GPU => statistiques basées\n# sur les deux GPU ensemble\nmodel = (\n    nn.SyncBatchNorm\n    .convert_sync_batchnorm(\n        model\n    )\n)\n\n\nmodel = model.to(\n    DEVICE\n)\n\n\nmodel = DDP(\n\n    model,\n\n    device_ids=[\n        LOCAL_RANK\n    ],\n\n    output_device=(\n        LOCAL_RANK\n    ),\n\n    broadcast_buffers=False\n)\n\n\n# ============================================================\n# 12. THRESHOLDS\n# ============================================================\n\ndef apply_thresholds(\n    scores,\n    thresholds\n):\n\n    return np.digitize(\n        scores,\n        thresholds\n    )\n\n\ndef qwk(\n    truth,\n    pred\n):\n\n    return cohen_kappa_score(\n        truth,\n        pred,\n        weights=\"quadratic\"\n    )\n\n\ndef threshold_score(\n    thresholds,\n    truth,\n    scores\n):\n\n    thresholds = np.asarray(\n        thresholds\n    )\n\n    if (\n        not\n        np.isfinite(\n            thresholds\n        ).all()\n    ):\n\n        return -1.0\n\n    if np.any(\n        np.diff(\n            thresholds\n        )\n        <=\n        0.02\n    ):\n\n        return -1.0\n\n    pred = apply_thresholds(\n        scores,\n        thresholds\n    )\n\n    return qwk(\n        truth,\n        pred\n    )\n\n\ndef optimize_thresholds(\n    truth,\n    scores\n):\n\n    truth = np.asarray(\n        truth,\n        dtype=np.int64\n    )\n\n    scores = np.asarray(\n        scores,\n        dtype=np.float64\n    )\n\n    assert np.isfinite(\n        scores\n    ).all()\n\n    starting_points = [\n\n        [\n            0.50,\n            1.50,\n            2.50,\n            3.50\n        ],\n\n        [\n            0.65,\n            1.30,\n            2.10,\n            2.90\n        ],\n\n        [\n            0.70,\n            1.40,\n            2.20,\n            3.00\n        ],\n\n        [\n            0.80,\n            1.60,\n            2.40,\n            3.20\n        ]\n    ]\n\n    best_thresholds = None\n    best_score = -1.0\n\n    for start in starting_points:\n\n        result = minimize(\n\n            lambda x:\n                -threshold_score(\n                    x,\n                    truth,\n                    scores\n                ),\n\n            x0=np.asarray(\n                start,\n                dtype=np.float64\n            ),\n\n            method=\"Nelder-Mead\",\n\n            options={\n                \"maxiter\":\n                    700,\n\n                \"xatol\":\n                    1e-4,\n\n                \"fatol\":\n                    1e-5\n            }\n        )\n\n        thresholds = (\n            result.x\n        )\n\n        score = threshold_score(\n            thresholds,\n            truth,\n            scores\n        )\n\n        if score > best_score:\n\n            best_score = score\n\n            best_thresholds = (\n                thresholds.copy()\n            )\n\n    return (\n        best_thresholds,\n        best_score\n    )\n\n\n# ============================================================\n# 13. FULL FP32 VALIDATION\n# ============================================================\n\ndef validate_fp32(\n    image_size\n):\n\n    module = (\n        model.module\n    )\n\n    module.eval()\n\n    dataset = DRDataset(\n\n        val_df,\n\n        val_transform(\n            image_size\n        ),\n\n        validation=True\n    )\n\n    loader = DataLoader(\n\n        dataset,\n\n        batch_size=12,\n\n        shuffle=False,\n\n        num_workers=4,\n\n        pin_memory=True,\n\n        persistent_workers=True\n    )\n\n    rows = []\n\n    mse_sum = 0.0\n    n = 0\n\n\n    with torch.no_grad():\n\n        for batch in tqdm(\n\n            loader,\n\n            desc=\"Validation FP32\",\n\n            disable=(\n                RANK != 0\n            )\n        ):\n\n            images = (\n\n                batch[\"image\"]\n                .to(\n                    DEVICE,\n                    non_blocking=True\n                )\n                .float()\n            )\n\n            targets = (\n\n                batch[\"target\"]\n                .to(\n                    DEVICE,\n                    non_blocking=True\n                )\n                .float()\n            )\n\n            # aucune AMP en validation\n            predictions = (\n\n                module(\n                    images\n                )\n                .squeeze(1)\n                .float()\n            )\n\n\n            if not torch.isfinite(\n                predictions\n            ).all():\n\n                raise RuntimeError(\n                    \"NaN/Inf durant validation FP32\"\n                )\n\n\n            mse_sum += (\n\n                (\n                    predictions\n                    -\n                    targets\n                )\n                .pow(2)\n                .sum()\n                .item()\n            )\n\n            n += images.size(0)\n\n\n            scores_cpu = (\n                predictions\n                .cpu()\n                .numpy()\n            )\n\n            targets_cpu = (\n                targets\n                .cpu()\n                .numpy()\n                .astype(int)\n            )\n\n\n            for i in range(\n                len(scores_cpu)\n            ):\n\n                rows.append({\n\n                    \"image\":\n                        batch[\n                            \"image_id\"\n                        ][i],\n\n                    \"patient_id\":\n                        str(\n                            batch[\n                                \"patient_id\"\n                            ][i]\n                        ),\n\n                    \"eye\":\n                        batch[\n                            \"eye\"\n                        ][i],\n\n                    \"target\":\n                        int(\n                            targets_cpu[i]\n                        ),\n\n                    \"score\":\n                        float(\n                            scores_cpu[i]\n                        )\n                })\n\n\n    prediction_df = pd.DataFrame(\n        rows\n    )\n\n\n    calibration = prediction_df[\n        prediction_df[\n            \"patient_id\"\n        ].isin(\n            CALIBRATION_IDS\n        )\n    ]\n\n\n    evaluation = prediction_df[\n        prediction_df[\n            \"patient_id\"\n        ].isin(\n            EVALUATION_IDS\n        )\n    ]\n\n\n    thresholds, calibration_qwk = (\n        optimize_thresholds(\n\n            calibration[\n                \"target\"\n            ].values,\n\n            calibration[\n                \"score\"\n            ].values\n        )\n    )\n\n\n    evaluation_pred = (\n        apply_thresholds(\n\n            evaluation[\n                \"score\"\n            ].values,\n\n            thresholds\n        )\n    )\n\n\n    evaluation_qwk = qwk(\n\n        evaluation[\n            \"target\"\n        ].values,\n\n        evaluation_pred\n    )\n\n\n    fixed_pred = apply_thresholds(\n\n        evaluation[\n            \"score\"\n        ].values,\n\n        [\n            0.5,\n            1.5,\n            2.5,\n            3.5\n        ]\n    )\n\n\n    fixed_qwk = qwk(\n\n        evaluation[\n            \"target\"\n        ].values,\n\n        fixed_pred\n    )\n\n\n    mae = np.mean(\n\n        np.abs(\n\n            evaluation[\n                \"target\"\n            ].values\n\n            -\n\n            evaluation[\n                \"score\"\n            ].values\n        )\n    )\n\n\n    return {\n\n        \"val_mse\":\n            mse_sum / n,\n\n        \"mae\":\n            float(mae),\n\n        \"fixed_qwk\":\n            float(fixed_qwk),\n\n        \"calibration_qwk\":\n            float(\n                calibration_qwk\n            ),\n\n        \"evaluation_qwk\":\n            float(\n                evaluation_qwk\n            ),\n\n        \"thresholds\":\n            thresholds,\n\n        \"predictions\":\n            prediction_df\n    }\n\n\n# ============================================================\n# 14. OPTIMIZER\n# ============================================================\n\ndef create_optimizer(\n    backbone_lr,\n    head_lr\n):\n\n    return torch.optim.AdamW(\n\n        [\n\n            {\n                \"params\":\n                    model\n                    .module\n                    .features\n                    .parameters(),\n\n                \"lr\":\n                    backbone_lr\n            },\n\n            {\n                \"params\":\n                    model\n                    .module\n                    .classifier\n                    .parameters(),\n\n                \"lr\":\n                    head_lr\n            }\n        ],\n\n        weight_decay=1e-4\n    )\n\n\n# ============================================================\n# 15. TRAIN ONE PHASE\n# ============================================================\n\nhistory = []\n\nbest_evaluation_qwk = -1.0\n\nglobal_epoch = 0\n\n\nfor phase in PHASES:\n\n    image_size = (\n        phase[\n            \"image_size\"\n        ]\n    )\n\n    batch_size = (\n        phase[\n            \"batch_size\"\n        ]\n    )\n\n    phase_epochs = (\n        phase[\n            \"epochs\"\n        ]\n    )\n\n\n    dataset = DRDataset(\n\n        train_df,\n\n        train_transform(\n            image_size\n        ),\n\n        validation=False\n    )\n\n\n    optimizer = create_optimizer(\n\n        phase[\n            \"backbone_lr\"\n        ],\n\n        phase[\n            \"head_lr\"\n        ]\n    )\n\n\n    for phase_epoch in range(\n        phase_epochs\n    ):\n\n        global_epoch += 1\n\n        alpha = (\n            phase[\n                \"sampling_alphas\"\n            ][\n                phase_epoch\n            ]\n        )\n\n\n        sampler = DistributedWeightedSampler(\n\n            labels=(\n                train_df[\n                    \"level\"\n                ].values\n            ),\n\n            alpha=alpha,\n\n            num_replicas=(\n                WORLD_SIZE\n            ),\n\n            rank=RANK,\n\n            seed=SEED\n        )\n\n\n        sampler.set_epoch(\n            global_epoch\n        )\n\n\n        loader = DataLoader(\n\n            dataset,\n\n            batch_size=batch_size,\n\n            sampler=sampler,\n\n            num_workers=4,\n\n            pin_memory=True,\n\n            persistent_workers=True,\n\n            drop_last=True\n        )\n\n\n        # OneCycle :\n        # warmup puis décroissance\n        scheduler = (\n            torch.optim.lr_scheduler\n            .OneCycleLR(\n\n                optimizer,\n\n                max_lr=[\n\n                    phase[\n                        \"backbone_lr\"\n                    ],\n\n                    phase[\n                        \"head_lr\"\n                    ]\n                ],\n\n                epochs=1,\n\n                steps_per_epoch=(\n                    len(loader)\n                ),\n\n                pct_start=0.10,\n\n                anneal_strategy=\"cos\",\n\n                div_factor=10.0,\n\n                final_div_factor=100.0\n            )\n        )\n\n\n        scaler = torch.amp.GradScaler(\n            \"cuda\",\n            enabled=True\n        )\n\n\n        model.train()\n\n\n        train_loss_sum = 0.0\n        train_count = 0\n\n        fp32_fallback_batches = 0\n\n        start_time = time.time()\n\n\n        progress = tqdm(\n\n            loader,\n\n            desc=(\n                f\"{image_size}px \"\n                f\"epoch \"\n                f\"{phase_epoch+1}/\"\n                f\"{phase_epochs}\"\n            ),\n\n            disable=(\n                RANK != 0\n            )\n        )\n\n\n        for images, targets in progress:\n\n            images = images.to(\n                DEVICE,\n                non_blocking=True\n            )\n\n            targets = (\n\n                targets\n                .to(\n                    DEVICE,\n                    non_blocking=True\n                )\n                .float()\n            )\n\n\n            optimizer.zero_grad(\n                set_to_none=True\n            )\n\n\n            # ----------------------------------------\n            # FIRST TRY FP16\n            # ----------------------------------------\n\n            with torch.autocast(\n\n                device_type=\"cuda\",\n\n                dtype=torch.float16\n            ):\n\n                predictions = (\n\n                    model(\n                        images\n                    )\n                    .squeeze(1)\n                )\n\n\n            local_bad = torch.tensor(\n\n                [\n                    0\n                    if\n                    torch.isfinite(\n                        predictions\n                    ).all()\n                    else\n                    1\n                ],\n\n                device=DEVICE,\n                dtype=torch.int32\n            )\n\n\n            # si UN GPU obtient NaN,\n            # les DEUX repassent le batch en FP32\n            dist.all_reduce(\n\n                local_bad,\n\n                op=dist.ReduceOp.MAX\n            )\n\n\n            use_fp32 = bool(\n                local_bad.item()\n            )\n\n\n            if use_fp32:\n\n                fp32_fallback_batches += 1\n\n                with torch.autocast(\n\n                    device_type=\"cuda\",\n\n                    enabled=False\n                ):\n\n                    predictions = (\n\n                        model(\n                            images.float()\n                        )\n                        .squeeze(1)\n                        .float()\n                    )\n\n                    loss = torch.mean(\n\n                        (\n                            predictions\n                            -\n                            targets\n                        ) ** 2\n                    )\n\n\n                loss.backward()\n\n\n                torch.nn.utils.clip_grad_norm_(\n\n                    model.parameters(),\n\n                    max_norm=1.0\n                )\n\n\n                optimizer.step()\n\n\n            else:\n\n                # MSE calculée en FP32\n                loss = torch.mean(\n\n                    (\n                        predictions.float()\n                        -\n                        targets\n                    ) ** 2\n                )\n\n\n                scaler.scale(\n                    loss\n                ).backward()\n\n\n                scaler.unscale_(\n                    optimizer\n                )\n\n\n                torch.nn.utils.clip_grad_norm_(\n\n                    model.parameters(),\n\n                    max_norm=1.0\n                )\n\n\n                scaler.step(\n                    optimizer\n                )\n\n\n                scaler.update()\n\n\n            scheduler.step()\n\n\n            batch_n = (\n                images.size(0)\n            )\n\n            train_loss_sum += (\n                loss.item()\n                *\n                batch_n\n            )\n\n            train_count += (\n                batch_n\n            )\n\n\n            if RANK == 0:\n\n                progress.set_postfix({\n\n                    \"MSE\":\n                        f\"{loss.item():.4f}\",\n\n                    \"alpha\":\n                        f\"{alpha:.2f}\",\n\n                    \"FP32\":\n                        fp32_fallback_batches\n                })\n\n\n        # ====================================================\n        # FIN TRAIN : attendre les 2 GPU\n        # ====================================================\n\n        dist.barrier()\n\n\n        # ====================================================\n        # VALIDATION UNIQUEMENT RANK 0\n        # ====================================================\n\n        if RANK == 0:\n\n            metrics = validate_fp32(\n                image_size\n            )\n\n\n            duration = (\n                time.time()\n                -\n                start_time\n            ) / 60\n\n\n            train_mse = (\n                train_loss_sum\n                /\n                train_count\n            )\n\n\n            current_qwk = (\n                metrics[\n                    \"evaluation_qwk\"\n                ]\n            )\n\n\n            print()\n            print(\"=\" * 72)\n\n            print(\n                f\"RESULTATS \"\n                f\"{image_size}px \"\n                f\"EPOCH \"\n                f\"{phase_epoch+1}\"\n            )\n\n            print(\"=\" * 72)\n\n            print(\n                \"Train MSE       :\",\n                round(\n                    train_mse,\n                    4\n                )\n            )\n\n            print(\n                \"Validation MSE  :\",\n                round(\n                    metrics[\n                        \"val_mse\"\n                    ],\n                    4\n                )\n            )\n\n            print(\n                \"MAE evaluation :\",\n                round(\n                    metrics[\n                        \"mae\"\n                    ],\n                    4\n                )\n            )\n\n            print(\n                \"QWK fixe       :\",\n                round(\n                    metrics[\n                        \"fixed_qwk\"\n                    ],\n                    4\n                )\n            )\n\n            print(\n                \"QWK calibration:\",\n                round(\n                    metrics[\n                        \"calibration_qwk\"\n                    ],\n                    4\n                )\n            )\n\n            print(\n                \"QWK evaluation :\",\n                round(\n                    current_qwk,\n                    4\n                )\n            )\n\n            print(\n                \"Seuils         :\",\n                np.round(\n                    metrics[\n                        \"thresholds\"\n                    ],\n                    4\n                )\n            )\n\n            print(\n                \"Fallback FP32  :\",\n                fp32_fallback_batches\n            )\n\n            print(\n                \"Durée          :\",\n                round(\n                    duration,\n                    1\n                ),\n                \"min\"\n            )\n\n\n            history.append({\n\n                \"global_epoch\":\n                    global_epoch,\n\n                \"resolution\":\n                    image_size,\n\n                \"phase_epoch\":\n                    phase_epoch + 1,\n\n                \"sampling_alpha\":\n                    alpha,\n\n                \"train_mse\":\n                    train_mse,\n\n                \"val_mse\":\n                    metrics[\n                        \"val_mse\"\n                    ],\n\n                \"mae\":\n                    metrics[\n                        \"mae\"\n                    ],\n\n                \"fixed_qwk\":\n                    metrics[\n                        \"fixed_qwk\"\n                    ],\n\n                \"calibration_qwk\":\n                    metrics[\n                        \"calibration_qwk\"\n                    ],\n\n                \"evaluation_qwk\":\n                    current_qwk,\n\n                \"fp32_fallback\":\n                    fp32_fallback_batches,\n\n                \"duration_minutes\":\n                    duration\n            })\n\n\n            pd.DataFrame(\n                history\n            ).to_csv(\n\n                HISTORY_PATH,\n\n                index=False\n            )\n\n\n            checkpoint = {\n\n                \"model_state_dict\":\n                    model\n                    .module\n                    .state_dict(),\n\n                \"global_epoch\":\n                    global_epoch,\n\n                \"resolution\":\n                    image_size,\n\n                \"phase_epoch\":\n                    phase_epoch + 1,\n\n                \"sampling_alpha\":\n                    alpha,\n\n                \"evaluation_qwk\":\n                    current_qwk,\n\n                \"thresholds\":\n                    metrics[\n                        \"thresholds\"\n                    ]\n            }\n\n\n            # latest\n            torch.save(\n\n                checkpoint,\n\n                LATEST_MODEL_PATH\n            )\n\n\n            # chaque époque\n            torch.save(\n\n                checkpoint,\n\n                CHECKPOINT_DIR\n                /\n                (\n                    f\"effb4_\"\n                    f\"{image_size}_\"\n                    f\"epoch\"\n                    f\"{phase_epoch+1}\"\n                    f\".pth\"\n                )\n            )\n\n\n            # meilleur\n            if (\n                current_qwk\n                >\n                best_evaluation_qwk\n            ):\n\n                best_evaluation_qwk = (\n                    current_qwk\n                )\n\n                torch.save(\n\n                    checkpoint,\n\n                    BEST_MODEL_PATH\n                )\n\n\n                metrics[\n                    \"predictions\"\n                ].to_csv(\n\n                    WORK_DIR\n                    /\n                    \"best_validation_predictions.csv\",\n\n                    index=False\n                )\n\n\n                np.save(\n\n                    WORK_DIR\n                    /\n                    \"best_thresholds.npy\",\n\n                    metrics[\n                        \"thresholds\"\n                    ]\n                )\n\n\n                print(\n                    \"✅ NOUVEAU \"\n                    \"MEILLEUR MODELE\"\n                )\n\n\n            print(\n                \"MEILLEUR QWK :\",\n                round(\n                    best_evaluation_qwk,\n                    4\n                )\n            )\n\n\n        # validation terminée\n        dist.barrier()\n\n        gc.collect()\n\n        torch.cuda.empty_cache()\n\n\n# ============================================================\n# 16. FINAL TTA\n# ============================================================\n\ndist.barrier()\n\n\nif RANK == 0:\n\n    print()\n    print(\"=\" * 72)\n    print(\"FINAL TTA + FUSION DES DEUX YEUX\")\n    print(\"=\" * 72)\n\n\n    checkpoint = torch.load(\n\n        BEST_MODEL_PATH,\n\n        map_location=DEVICE,\n\n        weights_only=False\n    )\n\n\n    model.module.load_state_dict(\n        checkpoint[\n            \"model_state_dict\"\n        ]\n    )\n\n\n    model.module.eval()\n\n\n    dataset = DRDataset(\n\n        val_df,\n\n        val_transform(\n            512\n        ),\n\n        validation=True\n    )\n\n\n    loader = DataLoader(\n\n        dataset,\n\n        batch_size=10,\n\n        shuffle=False,\n\n        num_workers=4,\n\n        pin_memory=True,\n\n        persistent_workers=True\n    )\n\n\n    final_rows = []\n\n\n    with torch.no_grad():\n\n        for batch in tqdm(\n            loader,\n            desc=\"TTA x4\"\n        ):\n\n            images = (\n\n                batch[\"image\"]\n                .to(\n                    DEVICE,\n                    non_blocking=True\n                )\n                .float()\n            )\n\n\n            # FP32 volontaire :\n            # dernière évaluation fiable\n            p0 = (\n\n                model.module(\n                    images\n                )\n                .squeeze(1)\n                .float()\n            )\n\n\n            p1 = (\n\n                model.module(\n                    torch.flip(\n                        images,\n                        dims=[3]\n                    )\n                )\n                .squeeze(1)\n                .float()\n            )\n\n\n            p2 = (\n\n                model.module(\n                    torch.flip(\n                        images,\n                        dims=[2]\n                    )\n                )\n                .squeeze(1)\n                .float()\n            )\n\n\n            p3 = (\n\n                model.module(\n                    torch.flip(\n                        images,\n                        dims=[\n                            2,\n                            3\n                        ]\n                    )\n                )\n                .squeeze(1)\n                .float()\n            )\n\n\n            scores = (\n                p0\n                +\n                p1\n                +\n                p2\n                +\n                p3\n            ) / 4.0\n\n\n            assert torch.isfinite(\n                scores\n            ).all()\n\n\n            scores = (\n                scores\n                .cpu()\n                .numpy()\n            )\n\n\n            targets = (\n\n                batch[\n                    \"target\"\n                ]\n                .numpy()\n                .astype(int)\n            )\n\n\n            for i in range(\n                len(scores)\n            ):\n\n                final_rows.append({\n\n                    \"image\":\n                        batch[\n                            \"image_id\"\n                        ][i],\n\n                    \"patient_id\":\n                        str(\n                            batch[\n                                \"patient_id\"\n                            ][i]\n                        ),\n\n                    \"eye\":\n                        batch[\n                            \"eye\"\n                        ][i],\n\n                    \"target\":\n                        int(\n                            targets[i]\n                        ),\n\n                    \"score_tta\":\n                        float(\n                            scores[i]\n                        )\n                })\n\n\n    tta_df = pd.DataFrame(\n        final_rows\n    )\n\n\n    # ========================================================\n    # 17. EYE FUSION\n    #\n    # own_score * (1-alpha)\n    # + other_eye * alpha\n    #\n    # alpha choisi uniquement sur calibration\n    # ========================================================\n\n    other_eye_scores = (\n\n        tta_df[\n            [\n                \"patient_id\",\n                \"eye\",\n                \"score_tta\"\n            ]\n        ]\n        .copy()\n    )\n\n\n    other_eye_scores[\n        \"eye\"\n    ] = (\n        other_eye_scores[\n            \"eye\"\n        ]\n        .map({\n\n            \"left\":\n                \"right\",\n\n            \"right\":\n                \"left\"\n        })\n    )\n\n\n    other_eye_scores = (\n        other_eye_scores\n        .rename(\n            columns={\n                \"score_tta\":\n                    \"other_score\"\n            }\n        )\n    )\n\n\n    tta_df = tta_df.merge(\n\n        other_eye_scores,\n\n        on=[\n            \"patient_id\",\n            \"eye\"\n        ],\n\n        how=\"left\"\n    )\n\n\n    # normalement chaque patient a 2 yeux\n    tta_df[\n        \"other_score\"\n    ] = (\n        tta_df[\n            \"other_score\"\n        ]\n        .fillna(\n            tta_df[\n                \"score_tta\"\n            ]\n        )\n    )\n\n\n    calibration = tta_df[\n        tta_df[\n            \"patient_id\"\n        ].isin(\n            CALIBRATION_IDS\n        )\n    ].copy()\n\n\n    evaluation = tta_df[\n        tta_df[\n            \"patient_id\"\n        ].isin(\n            EVALUATION_IDS\n        )\n    ].copy()\n\n\n    best_alpha = 0.0\n    best_thresholds = None\n    best_calibration_qwk = -1.0\n\n\n    for alpha in np.arange(\n        0.00,\n        0.501,\n        0.025\n    ):\n\n        calibration_score = (\n\n            (\n                1.0\n                -\n                alpha\n            )\n            *\n            calibration[\n                \"score_tta\"\n            ].values\n\n            +\n\n            alpha\n            *\n            calibration[\n                \"other_score\"\n            ].values\n        )\n\n\n        thresholds, score = (\n            optimize_thresholds(\n\n                calibration[\n                    \"target\"\n                ].values,\n\n                calibration_score\n            )\n        )\n\n\n        if score > best_calibration_qwk:\n\n            best_calibration_qwk = (\n                score\n            )\n\n            best_alpha = float(\n                alpha\n            )\n\n            best_thresholds = (\n                thresholds.copy()\n            )\n\n\n    # ========================================================\n    # 18. HONEST EVALUATION\n    # ========================================================\n\n    evaluation_score = (\n\n        (\n            1.0\n            -\n            best_alpha\n        )\n        *\n        evaluation[\n            \"score_tta\"\n        ].values\n\n        +\n\n        best_alpha\n        *\n        evaluation[\n            \"other_score\"\n        ].values\n    )\n\n\n    evaluation_pred = (\n        apply_thresholds(\n\n            evaluation_score,\n\n            best_thresholds\n        )\n    )\n\n\n    final_qwk = qwk(\n\n        evaluation[\n            \"target\"\n        ].values,\n\n        evaluation_pred\n    )\n\n\n    # sans eye fusion pour mesurer le gain\n    no_eye_thresholds, _ = (\n        optimize_thresholds(\n\n            calibration[\n                \"target\"\n            ].values,\n\n            calibration[\n                \"score_tta\"\n            ].values\n        )\n    )\n\n\n    no_eye_pred = apply_thresholds(\n\n        evaluation[\n            \"score_tta\"\n        ].values,\n\n        no_eye_thresholds\n    )\n\n\n    no_eye_qwk = qwk(\n\n        evaluation[\n            \"target\"\n        ].values,\n\n        no_eye_pred\n    )\n\n\n    print()\n    print(\"=\" * 72)\n    print(\"RESULTAT FINAL MODEL A\")\n    print(\"=\" * 72)\n\n\n    print(\n        \"Best training QWK :\",\n        round(\n            best_evaluation_qwk,\n            4\n        )\n    )\n\n\n    print(\n        \"TTA sans eye fusion :\",\n        round(\n            no_eye_qwk,\n            4\n        )\n    )\n\n\n    print(\n        \"Alpha autre oeil     :\",\n        round(\n            best_alpha,\n            3\n        )\n    )\n\n\n    print(\n        \"Thresholds finals    :\",\n        np.round(\n            best_thresholds,\n            4\n        )\n    )\n\n\n    print(\n        \"Calibration QWK      :\",\n        round(\n            best_calibration_qwk,\n            4\n        )\n    )\n\n\n    print(\n        \"✅ FINAL HONEST QWK   :\",\n        round(\n            final_qwk,\n            4\n        )\n    )\n\n\n    # sauvegarde score fusionné\n    tta_df[\n        \"score_eye_fused\"\n    ] = (\n\n        (\n            1.0\n            -\n            best_alpha\n        )\n        *\n        tta_df[\n            \"score_tta\"\n        ]\n\n        +\n\n        best_alpha\n        *\n        tta_df[\n            \"other_score\"\n        ]\n    )\n\n\n    tta_df.to_csv(\n\n        FINAL_TTA_PATH,\n\n        index=False\n    )\n\n\n    final_config = {\n\n        \"best_training_qwk\":\n            float(\n                best_evaluation_qwk\n            ),\n\n        \"tta_qwk_without_eye\":\n            float(\n                no_eye_qwk\n            ),\n\n        \"eye_alpha\":\n            float(\n                best_alpha\n            ),\n\n        \"thresholds\":\n            [\n                float(x)\n                for x\n                in best_thresholds\n            ],\n\n        \"calibration_qwk\":\n            float(\n                best_calibration_qwk\n            ),\n\n        \"final_evaluation_qwk\":\n            float(\n                final_qwk\n            )\n    }\n\n\n    with open(\n\n        WORK_DIR\n        /\n        \"efficientnet_b4_final_config.json\",\n\n        \"w\"\n    ) as f:\n\n        json.dump(\n            final_config,\n            f,\n            indent=2\n        )\n\n\n    # ========================================================\n    # 19. PACKAGE IMPORTANT FILES\n    # ========================================================\n\n    ZIP_PATH = Path(\n        \"/kaggle/working/\"\n        \"RETINOPATHY_MODEL_A_FINAL.zip\"\n    )\n\n\n    important_files = [\n\n        BEST_MODEL_PATH,\n\n        HISTORY_PATH,\n\n        FINAL_TTA_PATH,\n\n        WORK_DIR\n        /\n        \"efficientnet_b4_final_config.json\",\n\n        WORK_DIR\n        /\n        \"train_split.csv\",\n\n        WORK_DIR\n        /\n        \"val_split.csv\"\n    ]\n\n\n    with zipfile.ZipFile(\n\n        ZIP_PATH,\n\n        \"w\",\n\n        compression=(\n            zipfile.ZIP_DEFLATED\n        )\n    ) as zf:\n\n        for file in important_files:\n\n            if file.exists():\n\n                zf.write(\n                    file,\n                    arcname=file.name\n                )\n\n\n    print()\n    print(\n        \"✅ Archive finale :\",\n        ZIP_PATH\n    )\n\n\ndist.barrier()\n\ndist.destroy_process_group()\n'''\n\n\nSCRIPT_PATH.write_text(\n    script\n)\n\nprint(\n    \"✅ Script créé :\",\n    SCRIPT_PATH\n)\n\nprint(\n    \"Taille :\",\n    round(\n        SCRIPT_PATH.stat().st_size\n        / 1024,\n        1\n    ),\n    \"Ko\"\n)\n\nprint()\nprint(\n    \"🚀 Lancement sur les 2 Tesla T4...\"\n)\nprint()\n\n# ============================================================\n# LANCEMENT DDP\n# ============================================================\n\n!OMP_NUM_THREADS=1 torchrun \\\n    --standalone \\\n    --nproc_per_node=2 \\\n    /kaggle/working/train_effb4_ddp.py","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T17:57:01.206099Z","iopub.execute_input":"2026-08-11T17:57:01.206506Z","iopub.status.idle":"2026-08-11T19:30:06.438594Z","shell.execute_reply.started":"2026-08-11T17:57:01.206476Z","shell.execute_reply":"2026-08-11T19:30:06.437797Z"}},"outputs":[],"execution_count":null}]}