{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"datasetVersion","sourceId":532013,"datasetId":253160,"databundleVersionId":548363}],"dockerImageVersionId":31287,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# -*- coding: utf-8 -*-\n\"\"\"\n# 🔬 Diabetic Retinopathy Detection — EyePACS Dataset\n### EfficientNet-B4 | 13-Stage Preprocessing | Kaggle Notebook (GPU)\n\nThis notebook runs the full DR detection pipeline:\n1. Setup & verify GPU\n2. Dataset is directly available (no download needed!)\n3. Explore class distribution\n4. Run 13-stage preprocessing pipeline\n5. Train EfficientNet-B4 with class-balanced strategy\n6. Evaluate with confusion matrix & classification report\n\n**How to run on Kaggle:**\n1. Go to kaggle.com/c/diabetic-retinopathy-detection → Code → New Notebook\n2. In the notebook, click + Add Data → Competition Data → select “Diabetic Retinopathy Detection”\n3. Go to Settings → Accelerator → GPU P100\n4. Paste this script and run all cells\n\"\"\"","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-02-25T22:54:57.945100Z","iopub.execute_input":"2026-02-25T22:54:57.945458Z","iopub.status.idle":"2026-02-25T22:54:57.951194Z","shell.execute_reply.started":"2026-02-25T22:54:57.945431Z","shell.execute_reply":"2026-02-25T22:54:57.950598Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# CELL 1: Check GPU & Install Dependencies\n# ============================================================================\nimport os\nimport shutil\nimport torch\n\nprint(f\"🔧 PyTorch version: {torch.__version__}\")\nprint(f\"🎮 CUDA available: {torch.cuda.is_available()}\")\nif torch.cuda.is_available():\n    print(f\"🎮 GPU: {torch.cuda.get_device_name(0)}\")\n    print(f\"💾 GPU Memory: {torch.cuda.get_device_properties(0).total_memory / 1e9:.1f} GB\")\nelse:\n    print(\"⚠️  No GPU! Go to Settings → Accelerator → GPU P100\")\n\n# Install extra deps if needed (most are pre-installed on Kaggle)\nimport subprocess\nsubprocess.run([\"pip\", \"install\", \"-q\", \"scikit-image\"], check=True)\n\n# ── Cleanup existing corrupted cache from previous run ─────────────\nif os.path.exists(\"/kaggle/working/tensor_cache\"):\n    print(\"🧹 Cleaning up old cache to free space...\")\n    shutil.rmtree(\"/kaggle/working/tensor_cache\", ignore_errors=True)\nos.makedirs(\"/kaggle/working/tensor_cache\", exist_ok=True)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-02-25T22:54:57.953285Z","iopub.execute_input":"2026-02-25T22:54:57.953996Z","iopub.status.idle":"2026-02-25T22:55:01.011949Z","shell.execute_reply.started":"2026-02-25T22:54:57.953962Z","shell.execute_reply":"2026-02-25T22:55:01.011239Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# CELL 2: Configure Paths & Find Dataset\n# ============================================================================\n# ┌─────────────────────────────────────────────────────────────────────┐\n# │  ⚠️  '0 images found' issue?                                       │\n# │  The default competition data is stored in .zip files.            │\n# │  You should add an UNZIPPED version of the dataset instead.        │\n# │                                                                   │\n# │  HOW TO FIX:                                                      │\n# │  1. Click '+ Add Data' in the right sidebar                       │\n# │  2. Search for: \"resized-2015-2019-blindness-detection-images\"    │\n# │  3. Add the dataset by Benjamin Warner (already unzipped!)        │\n# └─────────────────────────────────────────────────────────────────────┘\n\nimport os\nimport glob\nfrom pathlib import Path\n\n# Paths discovery function (finds images wherever they are in /kaggle/input)\ndef find_dataset_paths():\n    input_root = \"/kaggle/input\"\n    img_dir, labels_csv, count = None, None, 0\n    \n    # ── FIRST: Search for images (.jpeg, .jpg, .png) ─────────────────────\n    # We look for the folder that has the MOST images (likely the train folder)\n    all_imgs = []\n    for ext in [\"**/*.jpeg\", \"**/*.jpg\", \"**/*.png\"]:\n        all_imgs.extend(glob.glob(f\"{input_root}/{ext}\", recursive=True))\n    \n    if all_imgs:\n        path_counts = {}\n        for p in all_imgs:\n            d = str(Path(p).parent)\n            path_counts[d] = path_counts.get(d, 0) + 1\n        \n        # ── PRIORITIZE 'TRAIN' FOLDERS ────────────────────────────────────\n        # Select folder with 'train' in name AND reasonably large count\n        train_folders = {d: c for d, c in path_counts.items() if \"train\" in d.lower()}\n        if train_folders:\n            img_dir = max(train_folders, key=train_folders.get)\n            count = train_folders[img_dir]\n        else:\n            # Fallback to absolute max (e.g., if it's the test set)\n            img_dir = max(path_counts, key=path_counts.get)\n            count = path_counts[img_dir]\n        \n        # ── SECOND: Find corresponding CSV labels ────────────────────────\n        # Search for CSVs anywhere in the input folder\n        all_csvs = glob.glob(f\"{input_root}/**/*.csv\", recursive=True)\n        \n        # Try to find a CSV that matches the image folder name (e.g., '15' or 'train')\n        folder_name = os.path.basename(img_dir).lower()\n        best_csv = None\n        \n        # Priorities: \n        # 1. contains 'trainLabels' and '15'\n        # 2. contains 'trainLabels'\n        # 3. any csv with 'label'\n        for csv_path in all_csvs:\n            name = os.path.basename(csv_path).lower()\n            if \"trainlabels\" in name and \"15\" in name:\n                best_csv = csv_path\n                break\n            elif \"trainlabels\" in name:\n                best_csv = csv_path\n        \n        if not best_csv and all_csvs:\n            best_csv = all_csvs[0]\n            \n        labels_csv = best_csv\n    \n    return img_dir, labels_csv, count\n\n# Kaggle default competition mount point\nKAGGLE_INPUT = \"/kaggle/input/diabetic-retinopathy-detection\"\n\nprint(\"🔍 Searching for images and labels in /kaggle/input...\")\nTRAIN_IMAGES_DIR, TRAIN_LABELS_CSV, IMG_COUNT = find_dataset_paths()\n\n# ── MANUAL OVERRIDE (Use this if script still fails) ──────────────────────\n# Confirmed paths for the Benjamin Warner dataset:\nif not IMG_COUNT or \"test\" in str(TRAIN_IMAGES_DIR).lower():\n    TRAIN_IMAGES_DIR = \"/kaggle/input/resized-2015-2019-blindness-detection-images/resized train 15/resized train 15\"\n    TRAIN_LABELS_CSV = \"/kaggle/input/resized-2015-2019-blindness-detection-images/labels/trainLabels15.csv\"\n    if os.path.exists(TRAIN_IMAGES_DIR) and os.path.exists(TRAIN_LABELS_CSV):\n         IMG_COUNT = len(list(Path(TRAIN_IMAGES_DIR).glob(\"*.jpeg\")))\n# ──────────────────────────────────────────────────────────────────────────\n\n# Output directory (auto-saved to your Kaggle account)\nOUTPUT_DIR = \"/kaggle/working/outputs\"\nos.makedirs(OUTPUT_DIR, exist_ok=True)\n\nif IMG_COUNT > 0 and TRAIN_LABELS_CSV:\n    print(f\"✅ FOUND DATASET!\")\n    print(f\"📂 Train images:  {IMG_COUNT:,} files in {TRAIN_IMAGES_DIR}\")\n    print(f\"📋 Labels CSV:    {TRAIN_LABELS_CSV}\")\n    print(f\"💾 Output dir:    {OUTPUT_DIR}\")\nelse:\n    print(\"\\n❌ NO DATASET FOUND!\")\n    print(\"   The images were not found in the typical folders.\")\n    print(\"\\n📁 DIRECTORY DEBUG (Items in /kaggle/input):\")\n    # Deeper walk for debugging\n    for root, dirs, files in os.walk(\"/kaggle/input\"):\n        level = root.replace(\"/kaggle/input\", \"\").count(os.sep)\n        if level <= 3:\n            indent = \"  \" * level\n            print(f\"{indent}📁 {os.path.basename(root)}/\")\n            if files: print(f\"{indent}  📄 {files[0]} (+ {len(files)-1} more)\")\n        if level > 3: continue\n    \n    print(\"\\n💡 HOW TO FIX:\")\n    print(\"   Look at the debug list above. Find the folder with '.jpeg' files.\")\n    print(\"   Find the 'trainLabels.csv' file.\")\n    print(\"   Manually paste those paths into the # MANUAL OVERRIDE section in this cell.\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-02-25T22:55:01.013253Z","iopub.execute_input":"2026-02-25T22:55:01.013506Z","iopub.status.idle":"2026-02-25T22:55:30.449984Z","shell.execute_reply.started":"2026-02-25T22:55:01.013484Z","shell.execute_reply":"2026-02-25T22:55:30.449281Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# CELL 3: Setup Project Source Code\n# ============================================================================\nimport sys\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom pathlib import Path\nfrom typing import Dict, List, Optional, Tuple\nfrom collections import Counter\nimport math\nimport time\nimport random\nimport logging\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import DataLoader, Dataset, WeightedRandomSampler\nfrom torchvision import transforms, models\nfrom PIL import Image\nfrom skimage.filters import frangi\nfrom tqdm import tqdm\n\n# Set seeds for reproducibility\nSEED = 42\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\nif torch.cuda.is_available():\n    torch.cuda.manual_seed_all(SEED)\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"[INFO] Running on: {DEVICE}\")\n\n# ── Constants ─────────────────────────────────────────────────────────────\nIMG_SIZE        = 380\nIMAGENET_MEAN   = [0.485, 0.456, 0.406]\nIMAGENET_STD    = [0.229, 0.224, 0.225]\nNUM_CLASSES     = 5\nDR_GRADE_NAMES  = {0: \"No DR\", 1: \"Mild\", 2: \"Moderate\", 3: \"Severe\", 4: \"Proliferative DR\"}\n\n# ── Global normalisation transform (defined once) ─────────────────────────\nnormalize_transform = transforms.Compose([\n    transforms.ToTensor(),\n    transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD)\n])\n\nprint(\"✅ All imports and constants loaded!\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-02-25T22:55:30.450926Z","iopub.execute_input":"2026-02-25T22:55:30.451133Z","iopub.status.idle":"2026-02-25T22:55:30.460296Z","shell.execute_reply.started":"2026-02-25T22:55:30.451114Z","shell.execute_reply":"2026-02-25T22:55:30.459606Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# CELL 4: Preprocessing Pipeline (13 Stages)\n# ============================================================================\n\ndef load_image(image_path: str) -> np.ndarray:\n    \"\"\"Load an image from disk and convert BGR → RGB.\"\"\"\n    img_bgr = cv2.imread(str(image_path))\n    if img_bgr is None:\n        raise ValueError(f\"Failed to decode image: {image_path}\")\n    return cv2.cvtColor(img_bgr, cv2.COLOR_BGR2RGB)\n\n\ndef crop_black_borders(img: np.ndarray, threshold: int = 10) -> np.ndarray:\n    \"\"\"Remove uninformative black borders around the fundus.\"\"\"\n    gray = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY)\n    _, mask = cv2.threshold(gray, threshold, 255, cv2.THRESH_BINARY)\n    kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (25, 25))\n    mask = cv2.morphologyEx(mask, cv2.MORPH_CLOSE, kernel)\n    contours, _ = cv2.findContours(mask, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n    if not contours:\n        return img\n    largest = max(contours, key=cv2.contourArea)\n    x, y, w, h = cv2.boundingRect(largest)\n    margin = int(min(w, h) * 0.02)\n    x, y = max(0, x - margin), max(0, y - margin)\n    w = min(img.shape[1] - x, w + 2 * margin)\n    h = min(img.shape[0] - y, h + 2 * margin)\n    return img[y:y+h, x:x+w]\n\n\ndef retina_centering(img: np.ndarray, output_size: int = IMG_SIZE) -> np.ndarray:\n    \"\"\"Centre the fundus on a square canvas, then resize.\"\"\"\n    h, w = img.shape[:2]\n    max_dim = max(h, w)\n    canvas = np.zeros((max_dim, max_dim, 3), dtype=img.dtype)\n    y_off, x_off = (max_dim - h) // 2, (max_dim - w) // 2\n    canvas[y_off:y_off+h, x_off:x_off+w] = img\n    return cv2.resize(canvas, (output_size, output_size), interpolation=cv2.INTER_AREA)\n\n\ndef apply_circular_mask(img: np.ndarray) -> np.ndarray:\n    \"\"\"Zero-out camera corners outside the retinal disc.\"\"\"\n    h, w = img.shape[:2]\n    cx, cy = w // 2, h // 2\n    radius = int(min(h, w) * 0.48)\n    mask = np.zeros((h, w), dtype=np.uint8)\n    cv2.circle(mask, (cx, cy), radius, 255, -1)\n    result = img.copy()\n    result[mask == 0] = 0\n    return result\n\n\ndef remove_glare(img: np.ndarray, threshold: int = 245, inpaint_radius: int = 10) -> np.ndarray:\n    \"\"\"Detect and inpaint lens reflections / glare spots.\"\"\"\n    gray = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY)\n    _, glare_mask = cv2.threshold(gray, threshold, 255, cv2.THRESH_BINARY)\n    kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (15, 15))\n    glare_mask = cv2.dilate(glare_mask, kernel, iterations=1)\n    glare_fraction = glare_mask.sum() / (255.0 * glare_mask.size)\n    if 0.001 < glare_fraction < 0.15:\n        img_bgr = cv2.cvtColor(img, cv2.COLOR_RGB2BGR)\n        inpainted = cv2.inpaint(img_bgr, glare_mask, inpaint_radius, cv2.INPAINT_TELEA)\n        return cv2.cvtColor(inpainted, cv2.COLOR_BGR2RGB)\n    return img\n\n\ndef ben_graham_preprocess(img: np.ndarray, sigma: int = 10) -> np.ndarray:\n    \"\"\"Ben Graham's preprocessing — Kaggle DR competition winner technique.\"\"\"\n    return cv2.addWeighted(img, 4, cv2.GaussianBlur(img, (0, 0), sigma), -4, 128)\n\n\ndef apply_clahe(img: np.ndarray, clip_limit: float = 3.0,\n                tile_grid_size: Tuple[int, int] = (8, 8)) -> np.ndarray:\n    \"\"\"CLAHE on L-channel of LAB colour space.\"\"\"\n    lab = cv2.cvtColor(img, cv2.COLOR_RGB2LAB)\n    l, a, b = cv2.split(lab)\n    clahe = cv2.createCLAHE(clipLimit=clip_limit, tileGridSize=tile_grid_size)\n    l_enhanced = clahe.apply(l)\n    return cv2.cvtColor(cv2.merge([l_enhanced, a, b]), cv2.COLOR_LAB2RGB)\n\n\ndef enhance_green_channel(img: np.ndarray, gamma: float = 0.8) -> np.ndarray:\n    \"\"\"Gamma correction on green channel for vessel contrast.\"\"\"\n    result = img.copy().astype(np.float32)\n    green = result[:, :, 1] / 255.0\n    result[:, :, 1] = (np.power(np.clip(green, 0.0, 1.0), gamma) * 255).clip(0, 255)\n    return result.astype(np.uint8)\n\n\ndef vessel_enhancement(img: np.ndarray, scale_range=(1, 3),\n                        scale_step=1, beta1=0.5, beta2=15) -> np.ndarray:\n    \"\"\"Frangi filter on green channel for vessel enhancement.\"\"\"\n    green = img[:, :, 1].astype(np.float32) / 255.0\n    start, stop = scale_range\n    sigmas = list(np.arange(start, stop + scale_step, scale_step))\n    vessel_map = frangi(green, sigmas=sigmas, beta=beta1, gamma=beta2)\n    if vessel_map.max() > 0:\n        vessel_map = vessel_map / vessel_map.max()\n    enhanced_green = np.clip(green + 0.15 * vessel_map, 0.0, 1.0)\n    result = img.copy()\n    result[:, :, 1] = (enhanced_green * 255).astype(np.uint8)\n    return result\n\n\ndef enhance_red_channel(img: np.ndarray, alpha: float = 1.3, beta: int = 0) -> np.ndarray:\n    \"\"\"Boost red channel for micro-aneurysms and haemorrhages.\"\"\"\n    result = img.copy()\n    result[:, :, 0] = np.clip(alpha * img[:, :, 0].astype(np.float32) + beta, 0, 255).astype(np.uint8)\n    return result\n\n\ndef _estimate_noise_level(img: np.ndarray) -> float:\n    \"\"\"Estimate image noise using Laplacian MAD method.\"\"\"\n    gray = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY) if img.ndim == 3 else img\n    laplacian = cv2.Laplacian(gray, cv2.CV_64F)\n    return float(np.median(np.abs(laplacian)) * 1.4826)\n\n\ndef adaptive_denoise(img: np.ndarray, fast: bool = False) -> np.ndarray:\n    \"\"\"Noise-aware denoising — adjusts strength based on estimated noise level.\"\"\"\n    noise_sigma = _estimate_noise_level(img)\n    if fast:\n        if noise_sigma > 15:\n            return cv2.bilateralFilter(img, d=9, sigmaColor=90, sigmaSpace=90)\n        elif noise_sigma > 8:\n            return cv2.bilateralFilter(img, d=9, sigmaColor=75, sigmaSpace=75)\n        else:\n            return cv2.bilateralFilter(img, d=5, sigmaColor=50, sigmaSpace=50)\n    if noise_sigma > 15:\n        h, hColor = 5, 5\n    elif noise_sigma > 8:\n        h, hColor = 3, 3\n    else:\n        h, hColor = 2, 2\n    return cv2.fastNlMeansDenoisingColored(img, None, h=h, hColor=hColor,\n                                            templateWindowSize=7, searchWindowSize=21)\n\n\ndef resize_image(img: np.ndarray, size: int = IMG_SIZE) -> np.ndarray:\n    \"\"\"Resize to target square dimensions.\"\"\"\n    return cv2.resize(img, (size, size), interpolation=cv2.INTER_LANCZOS4)\n\n\ndef normalize_for_imagenet(img: np.ndarray) -> torch.Tensor:\n    \"\"\"Convert to normalised PyTorch tensor (ImageNet stats).\"\"\"\n    pil_img = Image.fromarray(img.astype(np.uint8))\n    transform = transforms.Compose([\n        transforms.ToTensor(),\n        transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),\n    ])\n    return transform(pil_img)\n\n\ndef preprocess_image(image_path: str, apply_vessel: bool = True,\n                      fast_denoise: bool = False) -> Tuple[np.ndarray, Dict]:\n    \"\"\"Full 13-stage preprocessing pipeline. Returns UINT8 image + stages.\"\"\"\n    stages = {}\n\n    img = load_image(image_path)\n    stages[\"1_original\"] = img.copy()\n\n    img = crop_black_borders(img)\n    stages[\"2_cropped\"] = img.copy()\n\n    img = retina_centering(img)\n    stages[\"3_centered\"] = img.copy()\n\n    img = apply_circular_mask(img)\n    stages[\"4_circular_mask\"] = img.copy()\n\n    img = remove_glare(img)\n    stages[\"5_glare_removed\"] = img.copy()\n\n    img = ben_graham_preprocess(img)\n    stages[\"6_ben_graham\"] = img.copy()\n\n    img = apply_clahe(img)\n    stages[\"7_clahe\"] = img.copy()\n\n    img = enhance_green_channel(img)\n    stages[\"8_green_enhanced\"] = img.copy()\n\n    if apply_vessel:\n        img = vessel_enhancement(img)\n        stages[\"9_vessel\"] = img.copy()\n\n    img = enhance_red_channel(img)\n    stages[\"10_red_enhanced\"] = img.copy()\n\n    img = adaptive_denoise(img, fast=fast_denoise)\n    stages[\"11_denoised\"] = img.copy()\n\n    img = resize_image(img, size=IMG_SIZE)\n    stages[\"12_resized\"] = img.copy()\n\n    # We return the uint8 image for caching convenience\n    return img, stages\n\n\nprint(\"✅ 13-stage preprocessing pipeline loaded!\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-02-25T22:55:30.461788Z","iopub.execute_input":"2026-02-25T22:55:30.462285Z","iopub.status.idle":"2026-02-25T22:55:30.488679Z","shell.execute_reply.started":"2026-02-25T22:55:30.462263Z","shell.execute_reply":"2026-02-25T22:55:30.488050Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# CELL 5: Visualize Preprocessing Pipeline (Test on 1 image)\n# ============================================================================\n\n# Find a sample image to test (search for multiple extensions)\nsample_images = []\nfor ext in [\"**/*.jpeg\", \"**/*.jpg\", \"**/*.png\"]:\n    sample_images.extend(glob.glob(f\"{TRAIN_IMAGES_DIR}/{ext}\", recursive=True))\n\nif sample_images:\n    test_img_path = str(sorted(sample_images)[0])\n    print(f\"🖼️  Testing on: {Path(test_img_path).name}\")\n    print(f\"📂 Path: {test_img_path}\")\n\n    img_uint8, stages = preprocess_image(test_img_path, apply_vessel=True, fast_denoise=True)\n    tensor = normalize_for_imagenet(img_uint8)\n    print(f\"✅ Output tensor shape: {tensor.shape}\")\n\n    # Show key stages\n    key_stages = [\"1_original\", \"4_circular_mask\", \"5_glare_removed\",\n                  \"6_ben_graham\", \"7_clahe\", \"12_resized\"]\n    available = [s for s in key_stages if s in stages]\n\n    fig, axes = plt.subplots(1, len(available), figsize=(24, 5))\n    for ax, key in zip(axes, available):\n        ax.imshow(stages[key])\n        ax.set_title(key.replace(\"_\", \" \").title(), fontsize=10, fontweight='bold')\n        ax.axis('off')\n    fig.suptitle(\"Preprocessing Pipeline Stages\", fontsize=14, fontweight='bold')\n    plt.tight_layout()\n    plt.show()\nelse:\n    print(f\"⚠️  No images found in {TRAIN_IMAGES_DIR}\")\n    print(\"   Make sure Cell 2 found the correct images folder.\")\n\n# ============================================================================\n# CELL 6: Explore Dataset Distribution\n# ============================================================================\n\ndf_labels = pd.read_csv(TRAIN_LABELS_CSV)\nprint(f\"📊 Total images: {len(df_labels):,}\")\nprint(f\"\\n{'='*52}\")\nprint(\"  DR Class Distribution\")\nprint(f\"{'='*52}\")\n\ncounts = df_labels['level'].value_counts().sort_index()\nfor grade, count in counts.items():\n    pct = count / len(df_labels) * 100\n    bar = \"█\" * int(pct / 2)\n    name = DR_GRADE_NAMES.get(grade, f\"Grade {grade}\")\n    print(f\"  Grade {grade} | {name:<18s} | {count:>6,} ({pct:>5.1f}%) {bar}\")\nprint(f\"{'='*52}\")\n\n# Plot distribution\nfig, ax = plt.subplots(figsize=(10, 5))\ncolors = ['#2ecc71', '#f1c40f', '#e67e22', '#e74c3c', '#8e44ad']\nbars = ax.bar(range(NUM_CLASSES), [counts.get(i, 0) for i in range(NUM_CLASSES)],\n              color=colors, edgecolor='white', linewidth=1.5)\nax.set_xticks(range(NUM_CLASSES))\nax.set_xticklabels([f\"Grade {i}\\n{DR_GRADE_NAMES[i]}\" for i in range(NUM_CLASSES)])\nax.set_ylabel(\"Count\", fontweight='bold')\nax.set_title(\"EyePACS DR Class Distribution (Imbalanced!)\", fontsize=14, fontweight='bold')\nfor bar, count in zip(bars, [counts.get(i, 0) for i in range(NUM_CLASSES)]):\n    ax.text(bar.get_x() + bar.get_width()/2, bar.get_height() + 200,\n            f'{count:,}', ha='center', fontweight='bold')\nax.grid(axis='y', alpha=0.3)\nplt.tight_layout()\nplt.show()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-02-25T22:55:30.489647Z","iopub.execute_input":"2026-02-25T22:55:30.489875Z","iopub.status.idle":"2026-02-25T22:55:32.112918Z","shell.execute_reply.started":"2026-02-25T22:55:30.489854Z","shell.execute_reply":"2026-02-25T22:55:32.112279Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# CELL 7: Build EfficientNet-B4 Model\n# ============================================================================\n\ndef build_efficientnet_b4(num_classes=NUM_CLASSES, pretrained=True, freeze_backbone=False):\n    \"\"\"Build EfficientNet-B4 with upgraded classifier head.\"\"\"\n    if pretrained:\n        weights = models.EfficientNet_B4_Weights.IMAGENET1K_V1\n        model = models.efficientnet_b4(weights=weights)\n        print(\"✅ Loaded EfficientNet-B4 with ImageNet pretrained weights\")\n    else:\n        model = models.efficientnet_b4(weights=None)\n\n    if freeze_backbone:\n        for param in model.features.parameters():\n            param.requires_grad = False\n        print(\"🔒 Backbone frozen — only classifier head is trainable\")\n\n    # Upgraded classifier head (512-unit intermediate layer)\n    in_features = model.classifier[1].in_features\n    model.classifier = nn.Sequential(\n        nn.Dropout(p=0.3),\n        nn.Linear(in_features, 512),\n        nn.BatchNorm1d(512),\n        nn.ReLU(inplace=True),\n        nn.Dropout(p=0.3),\n        nn.Linear(512, num_classes),\n    )\n\n    trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    total = sum(p.numel() for p in model.parameters())\n    print(f\"📊 Model params: {total:,} total, {trainable:,} trainable\")\n    return model\n\n\n# Test model build\nmodel = build_efficientnet_b4(pretrained=True, freeze_backbone=True)\nmodel = model.to(DEVICE)\nprint(f\"✅ Model on {DEVICE}\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-02-25T22:55:32.113782Z","iopub.execute_input":"2026-02-25T22:55:32.114063Z","iopub.status.idle":"2026-02-25T22:55:32.527897Z","shell.execute_reply.started":"2026-02-25T22:55:32.114041Z","shell.execute_reply":"2026-02-25T22:55:32.527230Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# CELL 8: Dataset, Augmentation & Sampler\n# ============================================================================\n\n# ── Data Augmentation (DR-specific, applied on PIL images) ────────────────\ntrain_augmentation = transforms.Compose([\n    transforms.RandomHorizontalFlip(p=0.5),\n    transforms.RandomVerticalFlip(p=0.5),\n    transforms.RandomApply([transforms.RandomRotation(degrees=15)], p=0.5),\n    transforms.RandomApply([\n        transforms.ColorJitter(brightness=0.1, contrast=0.1, saturation=0.05, hue=0.02)\n    ], p=0.3),\n    transforms.RandomApply([\n        transforms.RandomAffine(degrees=0, translate=(0.05, 0.05), scale=(0.95, 1.05))\n    ], p=0.3),\n])\n\n# ── PyTorch Dataset (with disk caching, robust error handling) ────────────\nCACHE_DIR = \"/kaggle/working/tensor_cache\"\nos.makedirs(CACHE_DIR, exist_ok=True)\n\nclass DRImageDataset(Dataset):\n    \"\"\"\n    PyTorch Dataset with uint8 caching.\n    Augmentations are applied on PIL images before normalisation.\n    \"\"\"\n\n    def __init__(self, image_paths, label_map, apply_vessel=False,\n                 augmentation=None, fast_denoise=True, max_retries=3):\n        self.label_map = label_map\n        self.apply_vessel = apply_vessel\n        self.augmentation = augmentation\n        self.fast_denoise = fast_denoise\n        self.max_retries = max_retries\n        self.min_free_gb = 2.0  # Stop caching if disk space is below 2GB\n\n        # Filter to images that have labels\n        self.valid = [p for p in image_paths if p.stem in label_map]\n        if len(self.valid) < len(image_paths):\n            print(f\"  ⚠️ Skipped {len(image_paths) - len(self.valid)} images without labels\")\n\n    def __len__(self):\n        return len(self.valid)\n\n    def _get_free_disk_gb(self):\n        stat = os.statvfs(\"/kaggle/working\")\n        return (stat.f_bavail * stat.f_frsize) / (1024**3)\n\n    def _load_image_safe(self, path):\n        \"\"\"Load image with retries, return uint8 array or raise exception.\"\"\"\n        for attempt in range(self.max_retries):\n            try:\n                img_uint8, _ = preprocess_image(\n                    str(path),\n                    apply_vessel=self.apply_vessel,\n                    fast_denoise=self.fast_denoise\n                )\n                return img_uint8\n            except Exception as e:\n                if attempt == self.max_retries - 1:\n                    raise\n                # Try a different random image as fallback\n                path = random.choice(self.valid)\n        # Should never reach here\n        raise RuntimeError(\"Failed to load image after retries\")\n\n    def __getitem__(self, idx):\n        path = self.valid[idx]\n        cache_path = os.path.join(CACHE_DIR, f\"{path.stem}.npy\")\n\n        try:\n            # 1. Try loading from cache\n            if os.path.exists(cache_path):\n                img_uint8 = np.load(cache_path)\n            else:\n                # 2. Preprocess and save to cache if disk space permits\n                img_uint8, _ = preprocess_image(\n                    str(path),\n                    apply_vessel=self.apply_vessel,\n                    fast_denoise=self.fast_denoise\n                )\n                if self._get_free_disk_gb() > self.min_free_gb:\n                    np.save(cache_path, img_uint8)\n\n        except Exception as e:\n            # 3. If any error (corrupted file, disk full, etc.), retry with another image\n            new_idx = random.randint(0, len(self.valid) - 1)\n            return self.__getitem__(new_idx)\n\n        # 4. Convert to PIL, apply augmentation, then normalise\n        img_pil = Image.fromarray(img_uint8)\n        if self.augmentation is not None:\n            img_pil = self.augmentation(img_pil)\n\n        tensor = normalize_transform(img_pil)   # uses global normalisation\n        label = self.label_map[path.stem]\n\n        return tensor, label\n\n    def get_labels(self):\n        return [self.label_map[p.stem] for p in self.valid]\n\n\n# ── Helper functions ─────────────────────────────────────────────────────\ndef compute_class_weights(labels, device):\n    \"\"\"Inverse-frequency class weights for CrossEntropyLoss.\"\"\"\n    counts = Counter(labels)\n    total = sum(counts.values())\n    weights = [total / (NUM_CLASSES * counts.get(c, 1)) for c in range(NUM_CLASSES)]\n    print(f\"⚖️  Class weights: {[f'{w:.3f}' for w in weights]}\")\n    return torch.tensor(weights, dtype=torch.float32).to(device)\n\n\ndef build_weighted_sampler(labels):\n    \"\"\"WeightedRandomSampler to oversample minority classes.\"\"\"\n    counts = Counter(labels)\n    class_weights = {c: 1.0 / counts.get(c, 1) for c in range(NUM_CLASSES)}\n    sample_weights = [class_weights[label] for label in labels]\n    return WeightedRandomSampler(sample_weights, num_samples=len(labels), replacement=True)\n\n\nprint(\"✅ Dataset, augmentation, and sampler ready!\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-02-25T22:55:32.528920Z","iopub.execute_input":"2026-02-25T22:55:32.529223Z","iopub.status.idle":"2026-02-25T22:55:32.544687Z","shell.execute_reply.started":"2026-02-25T22:55:32.529190Z","shell.execute_reply":"2026-02-25T22:55:32.544002Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# CELL 9: Training Configuration\n# ============================================================================\n# ┌─────────────────────────────────────────────────────────────────────┐\n# │  ⚠️  ADJUST THESE HYPERPARAMETERS AS NEEDED                       │\n# └─────────────────────────────────────────────────────────────────────┘\n\n# Training hyperparameters\nEPOCHS            = 30        # Total training epochs\nBATCH_SIZE        = 16        # Reduce to 8 if GPU OOM\nLEARNING_RATE     = 1e-4      # Peak LR (after warm-up)\nWARMUP_EPOCHS     = 2         # Linear LR warm-up\nLABEL_SMOOTHING   = 0.1       # Reduces overconfidence\nGRAD_CLIP         = 1.0       # Max gradient norm\nVAL_SPLIT         = 0.2       # 20% for validation\nFREEZE_BACKBONE   = True      # True for Phase 1 (transfer learning)\nAPPLY_VESSEL      = False     # Set True for better accuracy (slower)\nFAST_DENOISE      = True      # True for faster batch processing\n\nprint(f\"\"\"\n{'='*55}\n  TRAINING CONFIGURATION\n{'='*55}\n  Epochs:           {EPOCHS}\n  Batch Size:       {BATCH_SIZE}\n  Learning Rate:    {LEARNING_RATE}\n  Warm-up Epochs:   {WARMUP_EPOCHS}\n  Label Smoothing:  {LABEL_SMOOTHING}\n  Grad Clip:        {GRAD_CLIP}\n  Val Split:        {VAL_SPLIT}\n  Freeze Backbone:  {FREEZE_BACKBONE}\n  Vessel Enhance:   {APPLY_VESSEL}\n  Fast Denoise:     {FAST_DENOISE}\n  Device:           {DEVICE}\n{'='*55}\n\"\"\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-02-25T22:55:32.545625Z","iopub.execute_input":"2026-02-25T22:55:32.546239Z","iopub.status.idle":"2026-02-25T22:55:32.561227Z","shell.execute_reply.started":"2026-02-25T22:55:32.546204Z","shell.execute_reply":"2026-02-25T22:55:32.560508Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# CELL 10: Train the Model 🏋️ (with stratified split + mixed precision)\n# ============================================================================\nfrom sklearn.model_selection import train_test_split\n\n# ── Prepare data ──────────────────────────────────────────────────────────\nSUPPORTED_EXTENSIONS = ('.jpg', '.jpeg', '.png', '.tiff', '.bmp')\nall_paths = sorted([\n    p for p in Path(TRAIN_IMAGES_DIR).iterdir()\n    if p.is_file() and p.suffix.lower() in SUPPORTED_EXTENSIONS\n])\nprint(f\"📂 Found {len(all_paths):,} images in {TRAIN_IMAGES_DIR}\")\n\n# Build label map from CSV\nlabel_map = dict(zip(df_labels['image'].astype(str), df_labels['level'].astype(int)))\nprint(f\"📋 Loaded {len(label_map):,} labels\")\n\n# Filter paths that have labels\nfiltered_paths = [p for p in all_paths if p.stem in label_map]\nfiltered_labels = [label_map[p.stem] for p in filtered_paths]\n\n# Stratified train/val split\ntrain_paths, val_paths, train_labels, val_labels = train_test_split(\n    filtered_paths, filtered_labels,\n    test_size=VAL_SPLIT,\n    random_state=SEED,\n    stratify=filtered_labels\n)\nprint(f\"📊 Split: {len(train_paths):,} train, {len(val_paths):,} val\")\n\n# Create datasets\ntrain_ds = DRImageDataset(train_paths, label_map, APPLY_VESSEL, train_augmentation, FAST_DENOISE)\nval_ds = DRImageDataset(val_paths, label_map, APPLY_VESSEL, None, FAST_DENOISE)\n\n# Weighted sampler (only for train)\nsampler = build_weighted_sampler(train_labels)\n\ntrain_loader = DataLoader(train_ds, batch_size=BATCH_SIZE, sampler=sampler,\n                          num_workers=2, pin_memory=True)\nval_loader = DataLoader(val_ds, batch_size=BATCH_SIZE, shuffle=False,\n                        num_workers=2, pin_memory=True)\n\n# ── Build model ───────────────────────────────────────────────────────────\nmodel = build_efficientnet_b4(pretrained=True, freeze_backbone=FREEZE_BACKBONE)\nmodel = model.to(DEVICE)\n\n# ── Loss, Optimizer, Scheduler ────────────────────────────────────────────\nclass_weights = compute_class_weights(train_labels, DEVICE)\ncriterion = nn.CrossEntropyLoss(weight=class_weights, label_smoothing=LABEL_SMOOTHING)\n\noptimizer = torch.optim.AdamW(\n    filter(lambda p: p.requires_grad, model.parameters()),\n    lr=LEARNING_RATE, weight_decay=1e-4\n)\n\ndef lr_lambda(epoch):\n    if epoch < WARMUP_EPOCHS:\n        return (epoch + 1) / WARMUP_EPOCHS\n    progress = (epoch - WARMUP_EPOCHS) / max(1, EPOCHS - WARMUP_EPOCHS)\n    return 0.5 * (1 + math.cos(math.pi * progress))\n\nscheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)\n\n# ── Mixed Precision Scaler (new torch.amp API) ────────────────────────────\nscaler = torch.amp.GradScaler('cuda')\n\n# ── Training Loop ─────────────────────────────────────────────────────────\nhistory = {\"train_loss\": [], \"train_acc\": [], \"val_loss\": [], \"val_acc\": [], \"lr\": []}\nbest_val_acc = 0.0\nbest_model_path = f\"{OUTPUT_DIR}/best_model.pth\"\n\nprint(f\"\\n🚀 Starting training for {EPOCHS} epochs with mixed precision...\\n\")\n\nfor epoch in range(1, EPOCHS + 1):\n    t0 = time.time()\n    current_lr = optimizer.param_groups[0][\"lr\"]\n\n    # ── Train ─────────────────────────────────────────────────────────────\n    model.train()\n    running_loss, correct, total = 0.0, 0, 0\n\n    pbar = tqdm(train_loader, desc=f\"Epoch {epoch}/{EPOCHS} [train]\", leave=True)\n    for inputs, labels in pbar:\n        inputs, labels = inputs.to(DEVICE), labels.to(DEVICE)\n\n        optimizer.zero_grad()\n\n        # Mixed precision forward pass (new API)\n        with torch.amp.autocast('cuda'):\n            outputs = model(inputs)\n            loss = criterion(outputs, labels)\n\n        # Backward pass with gradient scaling\n        scaler.scale(loss).backward()\n\n        # Gradient clipping (if enabled)\n        if GRAD_CLIP > 0:\n            scaler.unscale_(optimizer)  # Unscale gradients before clipping\n            nn.utils.clip_grad_norm_(model.parameters(), GRAD_CLIP)\n\n        # Optimizer step with scaler\n        scaler.step(optimizer)\n        scaler.update()\n\n        running_loss += loss.item() * inputs.size(0)\n        _, preds = outputs.max(1)\n        correct += (preds == labels).sum().item()\n        total += labels.size(0)\n\n        pbar.set_postfix(loss=f\"{loss.item():.4f}\", acc=f\"{correct/max(total,1):.4f}\")\n\n    train_loss = running_loss / max(total, 1)\n    train_acc = correct / max(total, 1)\n\n    # ── Validate ──────────────────────────────────────────────────────────\n    model.eval()\n    val_loss_sum, val_correct, val_total = 0.0, 0, 0\n\n    with torch.no_grad():\n        vbar = tqdm(val_loader, desc=f\"Epoch {epoch}/{EPOCHS} [val]\", leave=False)\n        for inputs, labels in vbar:\n            inputs, labels = inputs.to(DEVICE), labels.to(DEVICE)\n\n            # No need for autocast during validation (faster without it)\n            outputs = model(inputs)\n            loss = criterion(outputs, labels)\n\n            val_loss_sum += loss.item() * inputs.size(0)\n            _, preds = outputs.max(1)\n            val_correct += (preds == labels).sum().item()\n            val_total += labels.size(0)\n\n    val_loss = val_loss_sum / max(val_total, 1)\n    val_acc = val_correct / max(val_total, 1)\n\n    scheduler.step()\n    dt = time.time() - t0\n\n    history[\"train_loss\"].append(train_loss)\n    history[\"train_acc\"].append(train_acc)\n    history[\"val_loss\"].append(val_loss)\n    history[\"val_acc\"].append(val_acc)\n    history[\"lr\"].append(current_lr)\n\n    print(\n        f\"Epoch {epoch:>3d}/{EPOCHS} │ \"\n        f\"Train Loss: {train_loss:.4f}  Acc: {train_acc:.4f} │ \"\n        f\"Val Loss: {val_loss:.4f}  Acc: {val_acc:.4f} │ \"\n        f\"LR: {current_lr:.6f} │ {dt:.1f}s\"\n    )\n\n    # Save best model\n    if val_acc >= best_val_acc:\n        best_val_acc = val_acc\n        torch.save(model.state_dict(), best_model_path)\n        print(f\"  💾 Best model saved! (Val Acc: {val_acc:.4f})\")\n\n# Save final model\nfinal_model_path = f\"{OUTPUT_DIR}/final_model.pth\"\ntorch.save(model.state_dict(), final_model_path)\nprint(f\"\\n{'='*55}\")\nprint(f\"  ✅ Training complete!\")\nprint(f\"  🏆 Best validation accuracy: {best_val_acc:.4f}\")\nprint(f\"  💾 Best model: {best_model_path}\")\nprint(f\"  💾 Final model: {final_model_path}\")\nprint(f\"{'='*55}\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-02-25T22:58:04.032930Z","iopub.execute_input":"2026-02-25T22:58:04.033255Z","iopub.status.idle":"2026-02-26T01:09:57.894952Z","shell.execute_reply.started":"2026-02-25T22:58:04.033222Z","shell.execute_reply":"2026-02-26T01:09:57.893455Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# CELL: Resume Training from Saved Model (Phase 2 - Full Fine-tuning with Focal Loss)\n# ============================================================================\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport math\nimport time\nfrom tqdm import tqdm\nfrom sklearn.metrics import cohen_kappa_score\n\n# ── Focal Loss Implementation ──────────────────────────────────────────────\nclass FocalLoss(nn.Module):\n    def __init__(self, alpha=None, gamma=2.0, reduction='mean'):\n        \"\"\"\n        alpha: class weights (tensor). Can be None or a tensor of shape (num_classes).\n        gamma: focusing parameter.\n        reduction: 'mean' or 'sum'.\n        \"\"\"\n        super(FocalLoss, self).__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        self.reduction = reduction\n\n    def forward(self, inputs, targets):\n        \"\"\"\n        inputs: raw logits (shape: N x C)\n        targets: ground truth labels (shape: N)\n        \"\"\"\n        ce_loss = F.cross_entropy(inputs, targets, reduction='none')  # shape: N\n        pt = torch.exp(-ce_loss)  # probability of true class\n        focal_loss = (1 - pt) ** self.gamma * ce_loss\n\n        if self.alpha is not None:\n            # alpha_t = alpha[targets]\n            alpha_t = self.alpha[targets]\n            focal_loss = alpha_t * focal_loss\n\n        if self.reduction == 'mean':\n            return focal_loss.mean()\n        elif self.reduction == 'sum':\n            return focal_loss.sum()\n        else:\n            return focal_loss\n\n# ── Load the previously saved model ────────────────────────────────────────\ncheckpoint_path = \"/kaggle/working/outputs/best_model.pth\"  # or final_model.pth\nmodel = build_efficientnet_b4(pretrained=False, freeze_backbone=False)\nmodel.load_state_dict(torch.load(checkpoint_path, map_location=DEVICE))\nmodel = model.to(DEVICE)\n\n# ── Ensure all layers are trainable ────────────────────────────────────────\nfor param in model.parameters():\n    param.requires_grad = True\n\n# ── Class weights for focal loss (optional, but recommended) ───────────────\n# Recompute from train_labels (still in memory)\nif 'train_labels' not in globals():\n    # Fallback: compute from the dataset if needed (should exist)\n    train_labels = train_ds.get_labels()  # if train_ds still exists\nclass_weights = compute_class_weights(train_labels, DEVICE)  # returns tensor on DEVICE\n\n# ── Focal Loss criterion (instead of weighted CE) ──────────────────────────\ncriterion = FocalLoss(alpha=class_weights, gamma=2.0, reduction='mean')\nprint(\"✅ Using Focal Loss with gamma=2.0 and class weights\")\n\n# ── New training hyperparameters for fine-tuning ───────────────────────────\nPHASE2_EPOCHS = 20          # Number of additional epochs\nPHASE2_LR = 1e-5            # Lower learning rate for fine-tuning\nPHASE2_WARMUP = 2           # Optional warmup epochs\n\n# ── New optimizer (all parameters now trainable) ───────────────────────────\noptimizer = torch.optim.AdamW(model.parameters(), lr=PHASE2_LR, weight_decay=1e-4)\n\n# ── Cosine decay scheduler (with optional warmup) ──────────────────────────\ndef lr_lambda(epoch):\n    if epoch < PHASE2_WARMUP:\n        return (epoch + 1) / PHASE2_WARMUP\n    progress = (epoch - PHASE2_WARMUP) / max(1, PHASE2_EPOCHS - PHASE2_WARMUP)\n    return 0.5 * (1 + math.cos(math.pi * progress))\n\nscheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)\n\n# ── Mixed precision scaler ─────────────────────────────────────────────────\nscaler = torch.amp.GradScaler('cuda')\n\n# ── Continue tracking history (optional) ───────────────────────────────────\nhistory_phase2 = {\"train_loss\": [], \"train_acc\": [], \"val_loss\": [], \"val_acc\": [], \"val_kappa\": [], \"lr\": []}\nbest_val_kappa = -1.0   # Track best kappa for this phase\n\nprint(f\"\\n🚀 Resuming fine‑tuning from {checkpoint_path} for {PHASE2_EPOCHS} epochs with Focal Loss...\\n\")\n\nfor epoch in range(1, PHASE2_EPOCHS + 1):\n    t0 = time.time()\n    current_lr = optimizer.param_groups[0][\"lr\"]\n\n    # ── Train ─────────────────────────────────────────────────────────────\n    model.train()\n    running_loss, correct, total = 0.0, 0, 0\n    pbar = tqdm(train_loader, desc=f\"Phase2 Epoch {epoch}/{PHASE2_EPOCHS} [train]\", leave=True)\n    for inputs, labels in pbar:\n        inputs, labels = inputs.to(DEVICE), labels.to(DEVICE)\n        optimizer.zero_grad()\n        with torch.amp.autocast('cuda'):\n            outputs = model(inputs)\n            loss = criterion(outputs, labels)\n        scaler.scale(loss).backward()\n        if GRAD_CLIP > 0:\n            scaler.unscale_(optimizer)\n            nn.utils.clip_grad_norm_(model.parameters(), GRAD_CLIP)\n        scaler.step(optimizer)\n        scaler.update()\n        running_loss += loss.item() * inputs.size(0)\n        _, preds = outputs.max(1)\n        correct += (preds == labels).sum().item()\n        total += labels.size(0)\n        pbar.set_postfix(loss=f\"{loss.item():.4f}\", acc=f\"{correct/max(total,1):.4f}\")\n    train_loss = running_loss / max(total, 1)\n    train_acc = correct / max(total, 1)\n\n    # ── Validate ──────────────────────────────────────────────────────────\n    model.eval()\n    val_loss_sum, val_correct, val_total = 0.0, 0, 0\n    all_preds, all_labels = [], []\n    with torch.no_grad():\n        vbar = tqdm(val_loader, desc=f\"Phase2 Epoch {epoch}/{PHASE2_EPOCHS} [val]\", leave=False)\n        for inputs, labels in vbar:\n            inputs, labels = inputs.to(DEVICE), labels.to(DEVICE)\n            outputs = model(inputs)\n            loss = criterion(outputs, labels)\n            val_loss_sum += loss.item() * inputs.size(0)\n            _, preds = outputs.max(1)\n            val_correct += (preds == labels).sum().item()\n            val_total += labels.size(0)\n            all_preds.extend(preds.cpu().numpy())\n            all_labels.extend(labels.cpu().numpy())\n    val_loss = val_loss_sum / max(val_total, 1)\n    val_acc = val_correct / max(val_total, 1)\n    val_kappa = cohen_kappa_score(all_labels, all_preds, weights='quadratic')\n\n    scheduler.step()\n    dt = time.time() - t0\n\n    history_phase2[\"train_loss\"].append(train_loss)\n    history_phase2[\"train_acc\"].append(train_acc)\n    history_phase2[\"val_loss\"].append(val_loss)\n    history_phase2[\"val_acc\"].append(val_acc)\n    history_phase2[\"val_kappa\"].append(val_kappa)\n    history_phase2[\"lr\"].append(current_lr)\n\n    print(\n        f\"Phase2 Epoch {epoch:>3d}/{PHASE2_EPOCHS} │ \"\n        f\"Train Loss: {train_loss:.4f}  Acc: {train_acc:.4f} │ \"\n        f\"Val Loss: {val_loss:.4f}  Acc: {val_acc:.4f}  Kappa: {val_kappa:.4f} │ \"\n        f\"LR: {current_lr:.6f} │ {dt:.1f}s\"\n    )\n\n    # Save best model for this phase (based on kappa)\n    if val_kappa > best_val_kappa:\n        best_val_kappa = val_kappa\n        torch.save(model.state_dict(), \"/kaggle/working/outputs/best_model_phase2.pth\")\n        print(f\"  💾 New best model saved! (Kappa: {val_kappa:.4f})\")\n\nprint(f\"\\n✅ Phase 2 complete. Best kappa: {best_val_kappa:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T01:23:44.111174Z","iopub.execute_input":"2026-02-26T01:23:44.112050Z","iopub.status.idle":"2026-02-26T01:25:07.723162Z","shell.execute_reply.started":"2026-02-26T01:23:44.112017Z","shell.execute_reply":"2026-02-26T01:25:07.722177Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from google.colab import files\nfiles.download(\"/kaggle/working/outputs/best_model.pth\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T01:26:21.376951Z","iopub.execute_input":"2026-02-26T01:26:21.377683Z","iopub.status.idle":"2026-02-26T01:26:21.383490Z","shell.execute_reply.started":"2026-02-26T01:26:21.377649Z","shell.execute_reply":"2026-02-26T01:26:21.382876Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# CELL 11: Plot Training History\n# ============================================================================\n\nfig, axes = plt.subplots(1, 3, figsize=(20, 5))\nepochs_range = range(1, EPOCHS + 1)\n\n# Loss\naxes[0].plot(epochs_range, history[\"train_loss\"], 'b-o', label='Train Loss', markersize=3)\naxes[0].plot(epochs_range, history[\"val_loss\"], 'r-o', label='Val Loss', markersize=3)\naxes[0].set_xlabel(\"Epoch\")\naxes[0].set_ylabel(\"Loss\")\naxes[0].set_title(\"Loss\", fontweight='bold')\naxes[0].legend()\naxes[0].grid(alpha=0.3)\n\n# Accuracy\naxes[1].plot(epochs_range, history[\"train_acc\"], 'b-o', label='Train Acc', markersize=3)\naxes[1].plot(epochs_range, history[\"val_acc\"], 'r-o', label='Val Acc', markersize=3)\naxes[1].set_xlabel(\"Epoch\")\naxes[1].set_ylabel(\"Accuracy\")\naxes[1].set_title(\"Accuracy\", fontweight='bold')\naxes[1].legend()\naxes[1].grid(alpha=0.3)\naxes[1].set_ylim(0, 1.05)\n\n# LR\naxes[2].plot(epochs_range, history[\"lr\"], 'g-o', markersize=3)\naxes[2].set_xlabel(\"Epoch\")\naxes[2].set_ylabel(\"Learning Rate\")\naxes[2].set_title(\"Learning Rate Schedule\", fontweight='bold')\naxes[2].grid(alpha=0.3)\n\nfig.suptitle(\"Training History\", fontsize=14, fontweight='bold')\nplt.tight_layout()\nplt.savefig(f\"{OUTPUT_DIR}/training_history.png\", dpi=150, bbox_inches='tight')\nplt.show()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-02-26T01:18:42.619984Z","iopub.execute_input":"2026-02-26T01:18:42.620285Z","iopub.status.idle":"2026-02-26T01:18:43.544173Z","shell.execute_reply.started":"2026-02-26T01:18:42.620259Z","shell.execute_reply":"2026-02-26T01:18:43.543550Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# CELL 12: Evaluation — Confusion Matrix & Classification Report\n# ============================================================================\n\nfrom sklearn.metrics import (\n    accuracy_score, classification_report, confusion_matrix,\n    cohen_kappa_score\n)\n\n# Load best model\nmodel_eval = build_efficientnet_b4(pretrained=False, freeze_backbone=False)\nmodel_eval.load_state_dict(torch.load(best_model_path, map_location=DEVICE, weights_only=True))\nmodel_eval = model_eval.to(DEVICE)\nmodel_eval.eval()\n\n# Predict on validation set\nprint(\"🔍 Running predictions on validation set...\")\ny_true, y_pred = [], []\n\nwith torch.no_grad():\n    for inputs, labels in tqdm(val_loader, desc=\"Evaluating\"):\n        inputs = inputs.to(DEVICE)\n        outputs = model_eval(inputs)\n        _, preds = outputs.max(1)\n        y_true.extend(labels.cpu().tolist())\n        y_pred.extend(preds.cpu().tolist())\n\n# ── Metrics ───────────────────────────────────────────────────────────────\nacc = accuracy_score(y_true, y_pred)\nkappa = cohen_kappa_score(y_true, y_pred, weights=\"quadratic\")\nclass_names = [DR_GRADE_NAMES[i] for i in range(NUM_CLASSES)]\n\nprint(f\"\\n{'='*60}\")\nprint(f\"  EVALUATION RESULTS\")\nprint(f\"{'='*60}\")\nprint(f\"  Overall Accuracy     : {acc:.4f}  ({acc*100:.1f}%)\")\nprint(f\"  Quadratic Kappa (κ)  : {kappa:.4f}\")\nprint(f\"-\"*60)\nprint(\"  Classification Report:\")\nprint(classification_report(y_true, y_pred, labels=list(range(5)),\n                            target_names=class_names, zero_division=0))\nprint(f\"{'='*60}\")\n\n# ── Confusion Matrix ─────────────────────────────────────────────────────\ncm = confusion_matrix(y_true, y_pred, labels=list(range(5)))\n\nfig, axes = plt.subplots(1, 2, figsize=(18, 7))\n\nfor ax_idx, (normalize, title, fmt) in enumerate([\n    (False, \"Confusion Matrix (Counts)\", \"d\"),\n    (True, \"Confusion Matrix (Normalized)\", \".2f\"),\n]):\n    ax = axes[ax_idx]\n    if normalize:\n        row_sums = cm.sum(axis=1, keepdims=True)\n        row_sums[row_sums == 0] = 1\n        cm_display = cm.astype(float) / row_sums\n    else:\n        cm_display = cm\n\n    im = ax.imshow(cm_display, interpolation='nearest', cmap='Blues')\n    ax.figure.colorbar(im, ax=ax, fraction=0.046, pad=0.04)\n    ax.set(xticks=np.arange(5), yticks=np.arange(5),\n           xticklabels=class_names, yticklabels=class_names,\n           xlabel='Predicted', ylabel='True')\n    ax.set_title(title, fontsize=13, fontweight='bold', pad=15)\n    plt.setp(ax.get_xticklabels(), rotation=45, ha='right')\n\n    thresh = cm_display.max() / 2.0\n    for i in range(5):\n        for j in range(5):\n            val = cm_display[i, j]\n            text = f\"{val:{fmt}}\" if isinstance(val, (int, np.integer)) else f\"{val:.2f}\"\n            ax.text(j, i, text, ha='center', va='center', fontsize=11,\n                    fontweight='bold', color='white' if val > thresh else 'black')\n\nplt.suptitle(f\"DR Classification — Acc: {acc*100:.1f}% | κ: {kappa:.4f}\",\n             fontsize=14, fontweight='bold')\nplt.tight_layout()\nplt.savefig(f\"{OUTPUT_DIR}/confusion_matrix.png\", dpi=150, bbox_inches='tight')\nplt.show()\n\nprint(f\"\\n📊 Confusion matrix saved to: {OUTPUT_DIR}/confusion_matrix.png\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-02-26T01:10:22.881569Z","iopub.execute_input":"2026-02-26T01:10:22.881893Z","iopub.status.idle":"2026-02-26T01:11:03.768689Z","shell.execute_reply.started":"2026-02-26T01:10:22.881868Z","shell.execute_reply":"2026-02-26T01:11:03.767898Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# CELL 13: Phase 2 — Unfreeze Backbone & Fine-tune\n# ============================================================================\nprint(\"🔓 Phase 2: Unfreezing backbone for full fine-tuning...\")\n\n# Unfreeze all layers\nfor param in model.parameters():\n    param.requires_grad = True\n\n# Recompute number of trainable params\ntrainable = sum(p.numel() for p in model.parameters() if p.requires_grad)\ntotal = sum(p.numel() for p in model.parameters())\nprint(f\"📊 Model params: {total:,} total, {trainable:,} trainable (all layers now trainable)\")\n\n# Lower learning rate for fine-tuning\nPHASE2_EPOCHS = 15\nPHASE2_LR = 1e-5\n\noptimizer2 = torch.optim.AdamW(model.parameters(), lr=PHASE2_LR, weight_decay=1e-4)\n\ndef lr_lambda2(epoch):\n    progress = epoch / max(1, PHASE2_EPOCHS)\n    return 0.5 * (1 + math.cos(math.pi * progress))\n\nscheduler2 = torch.optim.lr_scheduler.LambdaLR(optimizer2, lr_lambda2)\n\n# New scaler for Phase 2 (optional, but keep for consistency)\nscaler2 = torch.amp.GradScaler('cuda')\n\nprint(f\"🚀 Starting Phase 2 for {PHASE2_EPOCHS} epochs with LR={PHASE2_LR}...\")\n\n# Continue training loop (use same history dict to extend plots)\nfor epoch in range(1, PHASE2_EPOCHS + 1):\n    t0 = time.time()\n    current_lr = optimizer2.param_groups[0][\"lr\"]\n\n    # ── Train ─────────────────────────────────────────────────────────────\n    model.train()\n    running_loss, correct, total = 0.0, 0, 0\n\n    pbar = tqdm(train_loader, desc=f\"Phase2 Epoch {epoch}/{PHASE2_EPOCHS} [train]\", leave=True)\n    for inputs, labels in pbar:\n        inputs, labels = inputs.to(DEVICE), labels.to(DEVICE)\n\n        optimizer2.zero_grad()\n\n        with torch.amp.autocast('cuda'):\n            outputs = model(inputs)\n            loss = criterion(outputs, labels)\n\n        scaler2.scale(loss).backward()\n\n        if GRAD_CLIP > 0:\n            scaler2.unscale_(optimizer2)\n            nn.utils.clip_grad_norm_(model.parameters(), GRAD_CLIP)\n\n        scaler2.step(optimizer2)\n        scaler2.update()\n\n        running_loss += loss.item() * inputs.size(0)\n        _, preds = outputs.max(1)\n        correct += (preds == labels).sum().item()\n        total += labels.size(0)\n\n        pbar.set_postfix(loss=f\"{loss.item():.4f}\", acc=f\"{correct/max(total,1):.4f}\")\n\n    train_loss = running_loss / max(total, 1)\n    train_acc = correct / max(total, 1)\n\n    # ── Validate ──────────────────────────────────────────────────────────\n    model.eval()\n    val_loss_sum, val_correct, val_total = 0.0, 0, 0\n\n    with torch.no_grad():\n        vbar = tqdm(val_loader, desc=f\"Phase2 Epoch {epoch}/{PHASE2_EPOCHS} [val]\", leave=False)\n        for inputs, labels in vbar:\n            inputs, labels = inputs.to(DEVICE), labels.to(DEVICE)\n            outputs = model(inputs)\n            loss = criterion(outputs, labels)\n            val_loss_sum += loss.item() * inputs.size(0)\n            _, preds = outputs.max(1)\n            val_correct += (preds == labels).sum().item()\n            val_total += labels.size(0)\n\n    val_loss = val_loss_sum / max(val_total, 1)\n    val_acc = val_correct / max(val_total, 1)\n\n    scheduler2.step()\n    dt = time.time() - t0\n\n    # Append to existing history\n    history[\"train_loss\"].append(train_loss)\n    history[\"train_acc\"].append(train_acc)\n    history[\"val_loss\"].append(val_loss)\n    history[\"val_acc\"].append(val_acc)\n    history[\"lr\"].append(current_lr)\n\n    print(\n        f\"Phase2 Epoch {epoch:>3d}/{PHASE2_EPOCHS} │ \"\n        f\"Train Loss: {train_loss:.4f}  Acc: {train_acc:.4f} │ \"\n        f\"Val Loss: {val_loss:.4f}  Acc: {val_acc:.4f} │ \"\n        f\"LR: {current_lr:.6f} │ {dt:.1f}s\"\n    )\n\n    # Save best model (overall best so far)\n    if val_acc >= best_val_acc:\n        best_val_acc = val_acc\n        torch.save(model.state_dict(), best_model_path)\n        print(f\"  💾 New best model saved! (Val Acc: {val_acc:.4f})\")\n\n# Save final Phase 2 model\nphase2_model_path = f\"{OUTPUT_DIR}/phase2_model.pth\"\ntorch.save(model.state_dict(), phase2_model_path)\nprint(f\"\\n✅ Phase 2 complete. Best val acc: {best_val_acc:.4f}\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-02-25T22:57:55.688115Z","iopub.status.idle":"2026-02-25T22:57:55.688505Z","shell.execute_reply.started":"2026-02-25T22:57:55.688310Z","shell.execute_reply":"2026-02-25T22:57:55.688333Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# CELL 14: Download Trained Model\n# ============================================================================\n# On Kaggle, files in /kaggle/working/ are saved as notebook Output.\n# You can download them from the Output tab, OR use this cell:\n\nimport shutil\n\n# Copy models to /kaggle/working/ root (easier to find in Output tab)\nshutil.copy2(best_model_path, \"/kaggle/working/best_model.pth\")\nshutil.copy2(final_model_path, \"/kaggle/working/final_model.pth\")\n\nprint(f\"\"\"\n{'='*60}\n  📁 ALL OUTPUT FILES (saved to Kaggle Output)\n{'='*60}\n  Best model:        /kaggle/working/best_model.pth\n  Final model:       /kaggle/working/final_model.pth\n  Training history:  {OUTPUT_DIR}/training_history.png\n  Confusion matrix:  {OUTPUT_DIR}/confusion_matrix.png\n{'='*60}\n\n  🔽 HOW TO DOWNLOAD YOUR MODEL:\n  1. After notebook finishes, click the \"Output\" tab on the right\n  2. You'll see best_model.pth and final_model.pth listed\n  3. Click the \"⋮\" menu → Download\n\n  💡 Next steps:\n  1. Check the confusion matrix — do all 5 classes have predictions?\n  2. If Grade 3/4 recall is low, run Phase 2 (Cell 15)\n  3. Try APPLY_VESSEL = True for better accuracy (slower training)\n\"\"\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-02-25T22:57:55.689482Z","iopub.status.idle":"2026-02-25T22:57:55.690099Z","shell.execute_reply.started":"2026-02-25T22:57:55.689894Z","shell.execute_reply":"2026-02-25T22:57:55.689919Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# CELL 15: Test on EyePACS Test Set (Generate Submission)\n# ============================================================================\n# This cell predicts on the official EyePACS test images and creates a\n# submission CSV (same format as the Kaggle competition).\n#\n# Make sure you added the test data:\n#   + Add Data → Competition Data → Diabetic Retinopathy Detection\n# The test images will be at: /kaggle/input/diabetic-retinopathy-detection/test/\n\nTEST_IMAGES_DIR = f\"{KAGGLE_INPUT}/test\"\ntest_paths = sorted(Path(TEST_IMAGES_DIR).glob(\"*.jpeg\"))\nprint(f\"🔍 Found {len(test_paths):,} test images\")\n\nif len(test_paths) > 0:\n    # Load best model\n    model_test = build_efficientnet_b4(pretrained=False, freeze_backbone=False)\n    model_test.load_state_dict(torch.load(best_model_path, map_location=DEVICE, weights_only=True))\n    model_test = model_test.to(DEVICE)\n    model_test.eval()\n\n    test_results = []\n\n    print(\"🧪 Running predictions on test set...\")\n    for i, img_path in enumerate(tqdm(test_paths, desc=\"Testing\")):\n        try:\n            img_uint8, _ = preprocess_image(str(img_path), apply_vessel=APPLY_VESSEL,\n                                           fast_denoise=True)\n            tensor = normalize_for_imagenet(img_uint8)\n            input_tensor = tensor.unsqueeze(0).to(DEVICE)\n\n            with torch.no_grad():\n                outputs = model_test(input_tensor)\n                probabilities = torch.softmax(outputs, dim=1)\n                predicted_grade = outputs.argmax(dim=1).item()\n                confidence = probabilities[0, predicted_grade].item()\n\n            test_results.append({\n                \"image\": img_path.stem,\n                \"level\": predicted_grade,\n                \"confidence\": confidence,\n            })\n        except Exception as e:\n            test_results.append({\"image\": img_path.stem, \"level\": 0, \"confidence\": 0.0})\n\n    # Create submission CSV\n    submission_df = pd.DataFrame(test_results)\n    submission_path = f\"{OUTPUT_DIR}/submission.csv\"\n    submission_df[[\"image\", \"level\"]].to_csv(submission_path, index=False)\n    shutil.copy2(submission_path, \"/kaggle/working/submission.csv\")\n\n    print(f\"\\n✅ Submission saved: /kaggle/working/submission.csv\")\n    print(f\"📊 Prediction distribution:\")\n    for grade in range(NUM_CLASSES):\n        count = (submission_df[\"level\"] == grade).sum()\n        pct = count / len(submission_df) * 100\n        print(f\"   Grade {grade} ({DR_GRADE_NAMES[grade]:<18s}): {count:>5,} ({pct:.1f}%)\")\n\n    # Show sample predictions\n    print(f\"\\n📋 Sample predictions (first 10):\")\n    print(submission_df.head(10).to_string(index=False))\nelse:\n    print(\"⚠️  No test images found. Add test data from the competition.\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-02-25T22:57:55.691594Z","iopub.status.idle":"2026-02-25T22:57:55.691999Z","shell.execute_reply.started":"2026-02-25T22:57:55.691791Z","shell.execute_reply":"2026-02-25T22:57:55.691828Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# CELL 16: Test on a Single Image (Interactive)\n# ============================================================================\n# ... (update to use the new normalisation) ...\ndef predict_single_image(image_path, model_path=None):\n    \"\"\"\n    Predict DR grade for a single retinal image.\n    \"\"\"\n    if model_path is None:\n        model_path = best_model_path\n\n    # Load model\n    model = build_efficientnet_b4(pretrained=False, freeze_backbone=False)\n    model.load_state_dict(torch.load(model_path, map_location=DEVICE, weights_only=True))\n    model = model.to(DEVICE)\n    model.eval()\n\n    # Preprocess (uint8)\n    img_uint8, stages = preprocess_image(str(image_path), apply_vessel=True, fast_denoise=True)\n    # Convert to PIL and normalise (no augmentation)\n    img_pil = Image.fromarray(img_uint8)\n    tensor = normalize_transform(img_pil).unsqueeze(0).to(DEVICE)\n\n    # Predict\n    with torch.no_grad():\n        outputs = model(tensor)\n        probabilities = torch.softmax(outputs, dim=1)[0]\n        predicted_grade = outputs.argmax(dim=1).item()\n\n    # Display results (same as before)\n    # ... (plotting code unchanged) ...\n    return predicted_grade, probabilities.cpu().numpy()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-02-25T22:57:55.693577Z","iopub.status.idle":"2026-02-25T22:57:55.694008Z","shell.execute_reply.started":"2026-02-25T22:57:55.693773Z","shell.execute_reply":"2026-02-25T22:57:55.693795Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── Example: Test on a sample image ──────────────────────────────────────\n# Uncomment one of these to test:\n\n# Option 1: Test on a training image\n# grade, probs = predict_single_image(f\"{TRAIN_IMAGES_DIR}/10_left.jpeg\")\n\n# Option 2: Test on a test image\n# grade, probs = predict_single_image(f\"{TEST_IMAGES_DIR}/44353_right.jpeg\")\n\nprint(\"✅ predict_single_image() function ready!\")\nprint(\"   Usage: grade, probs = predict_single_image('path/to/image.jpeg')\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-02-25T22:57:55.696106Z","iopub.status.idle":"2026-02-25T22:57:55.696463Z","shell.execute_reply.started":"2026-02-25T22:57:55.696282Z","shell.execute_reply":"2026-02-25T22:57:55.696303Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}