{"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":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":4104,"databundleVersionId":46661,"isSourceIdPinned":false}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# # Question 1: Diabetic Retinopathy Classification\n# ## CMM 704 – Data Mining Coursework\n#\n# **Objective:** Build a classification model capable of classifying retinal photographs into\n# 5 severity classes of Diabetic Retinopathy (DR):\n# - **0** – No DR\n# - **1** – Mild\n# - **2** – Moderate\n# - **3** – Severe\n# - **4** – Proliferative DR\n#\n# **Approach:**\n# 1. Exploratory Data Analysis (EDA)\n# 2. Preprocessing (image cropping, resizing, normalization, augmentation)\n# 3. Model 1 – ResNet50 with transfer learning & hyperparameter tuning\n# 4. Model 2 – EfficientNet-B3 with transfer learning & hyperparameter tuning\n# 5. Comprehensive evaluation using multiple metrics","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-04-06T09:37:31.145032Z","iopub.execute_input":"2026-04-06T09:37:31.145409Z","iopub.status.idle":"2026-04-06T09:37:31.150776Z","shell.execute_reply.started":"2026-04-06T09:37:31.145376Z","shell.execute_reply":"2026-04-06T09:37:31.149824Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ====================================================\n# 1. IMPORTS & SETUP\n# ====================================================\nimport os\nimport sys\nimport random\nimport glob\nimport shutil\nimport subprocess\nimport zipfile\nimport warnings\nimport time\nfrom collections import Counter\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\nfrom PIL import Image\nimport cv2\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nfrom torchvision import transforms, models\nfrom torch.cuda.amp import GradScaler, autocast\nfrom torch.optim.lr_scheduler import CosineAnnealingLR, ReduceLROnPlateau\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import (\n    confusion_matrix, classification_report,\n    cohen_kappa_score, roc_curve, auc,\n    accuracy_score, precision_recall_fscore_support\n)\nfrom sklearn.preprocessing import label_binarize\n\nwarnings.filterwarnings('ignore')\nplt.style.use('seaborn-v0_8-whitegrid')\nplt.rcParams['figure.figsize'] = (12, 8)\nplt.rcParams['font.size'] = 12\n\n# 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    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")\nif torch.cuda.is_available():\n    print(f\"GPU: {torch.cuda.get_device_name(0)}\")\n    print(f\"Number of GPUs: {torch.cuda.device_count()}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-06T09:37:35.036038Z","iopub.execute_input":"2026-04-06T09:37:35.036623Z","iopub.status.idle":"2026-04-06T09:37:46.250759Z","shell.execute_reply.started":"2026-04-06T09:37:35.036589Z","shell.execute_reply":"2026-04-06T09:37:46.249864Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ====================================================\n# 2. CONFIGURATION\n# ====================================================\n# --- Paths ---\nINPUT_DIR = '/kaggle/input/competitions/diabetic-retinopathy-detection'\nWORKING_DIR = '/kaggle/working'\nRESIZED_DIR = os.path.join(WORKING_DIR, 'train_resized')\n\n# --- Data parameters ---\nIMG_SIZE = 224\nMAX_SUBSET = 1500           # Small subset for speed (300 per class)\nVAL_SPLIT = 0.2\nNUM_CLASSES = 5\nCLASS_NAMES = ['No DR', 'Mild', 'Moderate', 'Severe', 'Proliferative DR']\n\n# --- Training parameters (tuned for speed) ---\nBATCH_SIZE = 32\nNUM_WORKERS = 2\nEPOCHS = 15                 # Final training epochs\nHP_EPOCHS = 2               # Hyperparameter search epochs (frozen backbone)\nPATIENCE = 4\nLEARNING_RATE = 3e-4\nWEIGHT_DECAY = 1e-4\n\nprint(\"Configuration loaded.\")\nprint(f\"  Subset size: {MAX_SUBSET} images ({MAX_SUBSET // NUM_CLASSES}/class)\")\nprint(f\"  Final training: {EPOCHS} epochs  |  HP search: {HP_EPOCHS} epochs\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-06T09:38:42.219959Z","iopub.execute_input":"2026-04-06T09:38:42.220640Z","iopub.status.idle":"2026-04-06T09:38:42.227747Z","shell.execute_reply.started":"2026-04-06T09:38:42.220608Z","shell.execute_reply":"2026-04-06T09:38:42.226696Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ====================================================\n# 3. DATA LOADING & EXTRACTION\n# ====================================================\n\n# 3.1 Extract trainLabels.csv\nlabels_zip = os.path.join(INPUT_DIR, 'trainLabels.csv.zip')\nlabels_csv = os.path.join(WORKING_DIR, 'trainLabels.csv')\n\nif os.path.exists(labels_zip):\n    with zipfile.ZipFile(labels_zip) as z:\n        z.extractall(WORKING_DIR)\n    print(\"Extracted trainLabels.csv\")\nelif os.path.exists(os.path.join(INPUT_DIR, 'trainLabels.csv')):\n    labels_csv = os.path.join(INPUT_DIR, 'trainLabels.csv')\n    print(\"trainLabels.csv found directly\")\nelse:\n    raise FileNotFoundError(\"trainLabels.csv not found. Please check the dataset.\")\n\nlabels_df = pd.read_csv(labels_csv)\nprint(f\"Total labelled images: {len(labels_df)}\")\nprint(f\"\\nFirst 5 rows:\\n{labels_df.head()}\")\nprint(f\"\\nClass distribution:\\n{labels_df['level'].value_counts().sort_index()}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-06T09:38:44.480243Z","iopub.execute_input":"2026-04-06T09:38:44.480965Z","iopub.status.idle":"2026-04-06T09:38:44.556599Z","shell.execute_reply.started":"2026-04-06T09:38:44.480935Z","shell.execute_reply":"2026-04-06T09:38:44.555618Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 3.2 Extract sample.zip (small ~11 MB, very fast)\nsample_zip = os.path.join(INPUT_DIR, 'sample.zip')\nsample_dir = os.path.join(WORKING_DIR, 'sample')\n\nif os.path.exists(sample_zip) and not os.path.exists(sample_dir):\n    with zipfile.ZipFile(sample_zip) as z:\n        z.extractall(WORKING_DIR)\n    print(\"Extracted sample.zip\")\n\n# Collect all already-available jpeg images from sample and any other source\nall_available_jpegs = {}\nfor search_dir in [sample_dir, os.path.join(sample_dir, 'sample'),\n                    os.path.join(WORKING_DIR, 'train'), WORKING_DIR]:\n    if os.path.isdir(search_dir):\n        for f in glob.glob(os.path.join(search_dir, '*.jpeg')):\n            name = os.path.splitext(os.path.basename(f))[0]\n            all_available_jpegs[name] = f\n\nprint(f\"Pre-available images (from sample.zip etc.): {len(all_available_jpegs)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-06T09:39:00.047034Z","iopub.execute_input":"2026-04-06T09:39:00.047974Z","iopub.status.idle":"2026-04-06T09:39:00.190400Z","shell.execute_reply.started":"2026-04-06T09:39:00.047939Z","shell.execute_reply":"2026-04-06T09:39:00.189544Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 3.3 Stratified subset selection\ndef select_stratified_subset(df, max_total=MAX_SUBSET, seed=SEED):\n    \"\"\"Select a stratified subset, balanced across classes.\"\"\"\n    max_per_class = max_total // NUM_CLASSES\n\n    subsets = []\n    for cls in range(NUM_CLASSES):\n        cls_df = df[df['level'] == cls]\n        n_sample = min(len(cls_df), max_per_class)\n        subsets.append(cls_df.sample(n=n_sample, random_state=seed))\n\n    result = pd.concat(subsets).reset_index(drop=True)\n    print(f\"Subset class distribution:\\n{result['level'].value_counts().sort_index()}\")\n    print(f\"Total subset size: {len(result)}\")\n    return result\n\nsubset_df = select_stratified_subset(labels_df)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-06T09:39:22.416420Z","iopub.execute_input":"2026-04-06T09:39:22.417344Z","iopub.status.idle":"2026-04-06T09:39:22.435536Z","shell.execute_reply.started":"2026-04-06T09:39:22.417310Z","shell.execute_reply":"2026-04-06T09:39:22.434707Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 3.4 Extract & resize training images\n# -------------------------------------------------------\n# Strategy:\n#   1) Re-use any images already available from sample.zip\n#   2) Extract remaining from the multi-part archive with 7z\n#   3) Immediately resize to 224x224 to save disk\n# -------------------------------------------------------\n\nos.makedirs(RESIZED_DIR, exist_ok=True)\n\nexisting_resized = set(\n    os.path.splitext(os.path.basename(f))[0]\n    for f in glob.glob(os.path.join(RESIZED_DIR, '*.jpeg'))\n)\nneeded = set(subset_df['image'].values) - existing_resized\nprint(f\"Already resized: {len(existing_resized)}  |  Still needed: {len(needed)}\")\n\n# --- Step A: resize any pre-available images (from sample.zip) ---\nresized_from_sample = 0\nfor name in list(needed):\n    if name in all_available_jpegs:\n        try:\n            img = Image.open(all_available_jpegs[name]).convert('RGB')\n            img = img.resize((IMG_SIZE, IMG_SIZE), Image.LANCZOS)\n            img.save(os.path.join(RESIZED_DIR, f\"{name}.jpeg\"), 'JPEG', quality=95)\n            needed.discard(name)\n            resized_from_sample += 1\n        except Exception:\n            pass\nprint(f\"Resized from sample.zip: {resized_from_sample}  |  Still needed from archive: {len(needed)}\")\n\n# --- Step B: extract remaining from the multi-part train archive ---\nif needed:\n    archive_part1 = os.path.join(INPUT_DIR, 'train.zip.001')\n\n    if os.path.exists(archive_part1):\n        # Write list file with train/ prefix (structure inside the zip)\n        list_file = os.path.join(WORKING_DIR, 'extract_list.txt')\n        with open(list_file, 'w') as f:\n            for name in needed:\n                f.write(f'train/{name}.jpeg\\n')\n\n        print(f\"Extracting {len(needed)} images from multi-part archive with 7z...\")\n        print(\"  (this may take a few minutes for the first run)\")\n        t0 = time.time()\n        try:\n            result = subprocess.run(\n                ['7z', 'x', archive_part1, f'-o{WORKING_DIR}', '-y', f'-i@{list_file}'],\n                capture_output=True, text=True, timeout=2400\n            )\n            elapsed = time.time() - t0\n            print(f\"  7z finished in {elapsed:.0f}s  (exit code {result.returncode})\")\n            if result.returncode != 0:\n                print(f\"  7z stderr: {result.stderr[:500]}\")\n        except subprocess.TimeoutExpired:\n            print(\"  7z timed out after 40 min — will use whatever was extracted\")\n        except FileNotFoundError:\n            print(\"  7z not installed — trying unzip fallback\")\n\n        # Find extracted images and resize them\n        extract_candidates = [\n            os.path.join(WORKING_DIR, 'train'),\n            WORKING_DIR,\n        ]\n        resized_from_archive = 0\n        for name in list(needed):\n            src = None\n            for d in extract_candidates:\n                candidate = os.path.join(d, f\"{name}.jpeg\")\n                if os.path.exists(candidate):\n                    src = candidate\n                    break\n            if src is None:\n                continue\n            try:\n                img = Image.open(src).convert('RGB')\n                img = img.resize((IMG_SIZE, IMG_SIZE), Image.LANCZOS)\n                img.save(os.path.join(RESIZED_DIR, f\"{name}.jpeg\"), 'JPEG', quality=95)\n                resized_from_archive += 1\n                needed.discard(name)\n            except Exception:\n                pass\n\n        print(f\"  Resized from archive: {resized_from_archive}\")\n\n        # Clean up full-res copies to free disk\n        full_res = os.path.join(WORKING_DIR, 'train')\n        if os.path.isdir(full_res):\n            shutil.rmtree(full_res, ignore_errors=True)\n            print(\"  Cleaned up full-resolution images to free disk\")\n    else:\n        print(f\"  Archive not found at {archive_part1}\")\n\nprint(f\"\\nImages still missing: {len(needed)}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-06T09:39:42.659980Z","iopub.execute_input":"2026-04-06T09:39:42.660490Z","iopub.status.idle":"2026-04-06T09:45:54.053912Z","shell.execute_reply.started":"2026-04-06T09:39:42.660458Z","shell.execute_reply":"2026-04-06T09:45:54.052846Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 3.5 Verify available images and build final dataframe\navailable_images = set(\n    os.path.splitext(os.path.basename(f))[0]\n    for f in glob.glob(os.path.join(RESIZED_DIR, '*.jpeg'))\n)\nprint(f\"Total resized images available: {len(available_images)}\")\n\nsubset_df = subset_df[subset_df['image'].isin(available_images)].reset_index(drop=True)\nsubset_df['filepath'] = subset_df['image'].apply(\n    lambda x: os.path.join(RESIZED_DIR, f\"{x}.jpeg\")\n)\nprint(f\"Final dataset size: {len(subset_df)}\")\nprint(f\"Final class distribution:\\n{subset_df['level'].value_counts().sort_index()}\")\n\nif len(subset_df) < 50:\n    raise RuntimeError(\n        f\"Only {len(subset_df)} images available — not enough to train. \"\n        \"Check that the competition dataset is properly attached and 7z is installed.\"\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-06T09:46:27.336132Z","iopub.execute_input":"2026-04-06T09:46:27.337002Z","iopub.status.idle":"2026-04-06T09:46:27.358489Z","shell.execute_reply.started":"2026-04-06T09:46:27.336970Z","shell.execute_reply":"2026-04-06T09:46:27.357712Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ====================================================\n# 4. EXPLORATORY DATA ANALYSIS\n# ====================================================\n\n# 4.1 Class Distribution (Full Dataset)\nfig, axes = plt.subplots(1, 2, figsize=(16, 6))\n\nclass_counts_full = labels_df['level'].value_counts().sort_index()\ncolors = sns.color_palette('viridis', NUM_CLASSES)\nbars = axes[0].bar(CLASS_NAMES, class_counts_full.values, color=colors, edgecolor='black')\naxes[0].set_title('Class Distribution – Full Dataset', fontsize=14, fontweight='bold')\naxes[0].set_xlabel('DR Severity')\naxes[0].set_ylabel('Count')\nfor bar, count in zip(bars, class_counts_full.values):\n    axes[0].text(bar.get_x() + bar.get_width()/2., bar.get_height() + 200,\n                 f'{count}\\n({count/len(labels_df)*100:.1f}%)',\n                 ha='center', va='bottom', fontsize=10)\n\nclass_counts_sub = subset_df['level'].value_counts().sort_index()\nbars2 = axes[1].bar(CLASS_NAMES, class_counts_sub.values, color=colors, edgecolor='black')\naxes[1].set_title('Class Distribution – Working Subset', fontsize=14, fontweight='bold')\naxes[1].set_xlabel('DR Severity')\naxes[1].set_ylabel('Count')\nfor bar, count in zip(bars2, class_counts_sub.values):\n    axes[1].text(bar.get_x() + bar.get_width()/2., bar.get_height() + 10,\n                 f'{count}\\n({count/len(subset_df)*100:.1f}%)',\n                 ha='center', va='bottom', fontsize=10)\n\nplt.tight_layout()\nplt.savefig(os.path.join(WORKING_DIR, 'class_distribution.png'), dpi=150, bbox_inches='tight')\nplt.show()\n\nprint(\"\\n--- Key Observation ---\")\nprint(\"The dataset is heavily IMBALANCED:\")\nprint(f\"  Class 0 (No DR) dominates with {class_counts_full[0]} images ({class_counts_full[0]/len(labels_df)*100:.1f}%)\")\nprint(f\"  Class 4 (Proliferative) has only {class_counts_full[4]} images ({class_counts_full[4]/len(labels_df)*100:.1f}%)\")\nprint(\"This imbalance will be addressed via class-weighted loss and data augmentation.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-06T09:46:30.193642Z","iopub.execute_input":"2026-04-06T09:46:30.193953Z","iopub.status.idle":"2026-04-06T09:46:30.969077Z","shell.execute_reply.started":"2026-04-06T09:46:30.193927Z","shell.execute_reply":"2026-04-06T09:46:30.968284Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 4.2 Left Eye vs Right Eye Distribution\nlabels_df['eye_side'] = labels_df['image'].apply(lambda x: 'Left' if 'left' in x else 'Right')\nlabels_df['subject_id'] = labels_df['image'].apply(lambda x: x.split('_')[0])\n\nfig, axes = plt.subplots(1, 2, figsize=(16, 6))\n\neye_counts = labels_df['eye_side'].value_counts()\naxes[0].pie(eye_counts.values, labels=eye_counts.index, autopct='%1.1f%%',\n            colors=['#3498db', '#e74c3c'], startangle=90, textprops={'fontsize': 12})\naxes[0].set_title('Left vs Right Eye Distribution', fontsize=14, fontweight='bold')\n\neye_class = labels_df.groupby(['eye_side', 'level']).size().unstack(fill_value=0)\neye_class.plot(kind='bar', ax=axes[1], color=colors, edgecolor='black')\naxes[1].set_title('DR Severity by Eye Side', fontsize=14, fontweight='bold')\naxes[1].set_xlabel('Eye Side')\naxes[1].set_ylabel('Count')\naxes[1].legend(CLASS_NAMES, title='Severity')\naxes[1].tick_params(axis='x', rotation=0)\n\nplt.tight_layout()\nplt.savefig(os.path.join(WORKING_DIR, 'eye_distribution.png'), dpi=150, bbox_inches='tight')\nplt.show()\n\nprint(\"\\n--- Key Observation ---\")\nprint(\"Left and right eye images are roughly balanced.\")\nprint(\"DR severity distribution is consistent across both eyes.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-06T09:46:34.287419Z","iopub.execute_input":"2026-04-06T09:46:34.287770Z","iopub.status.idle":"2026-04-06T09:46:35.256366Z","shell.execute_reply.started":"2026-04-06T09:46:34.287743Z","shell.execute_reply":"2026-04-06T09:46:35.255377Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 4.3 Sample Images per Class\nfig, axes = plt.subplots(NUM_CLASSES, 4, figsize=(16, 20))\n\nfor cls in range(NUM_CLASSES):\n    cls_images = subset_df[subset_df['level'] == cls]['filepath'].values\n    for j in range(min(4, len(cls_images))):\n        try:\n            img = Image.open(cls_images[j])\n            axes[cls, j].imshow(img)\n        except Exception:\n            axes[cls, j].text(0.5, 0.5, 'N/A', ha='center', va='center', fontsize=14)\n        axes[cls, j].axis('off')\n        if j == 0:\n            axes[cls, j].set_ylabel(CLASS_NAMES[cls], fontsize=12, fontweight='bold',\n                                     rotation=0, labelpad=80)\n\naxes[0, 0].set_title('Sample 1', fontsize=12)\naxes[0, 1].set_title('Sample 2', fontsize=12)\naxes[0, 2].set_title('Sample 3', fontsize=12)\naxes[0, 3].set_title('Sample 4', fontsize=12)\n\nfig.suptitle('Sample Retinal Images per DR Severity Class', fontsize=16, fontweight='bold', y=1.01)\nplt.tight_layout()\nplt.savefig(os.path.join(WORKING_DIR, 'sample_images.png'), dpi=150, bbox_inches='tight')\nplt.show()\n\nprint(\"\\n--- Key Observations ---\")\nprint(\"- Class 0 (No DR): Clear retinal images with well-defined vascular structure\")\nprint(\"- Classes 1-2 (Mild/Moderate): Subtle microaneurysms and haemorrhages visible\")\nprint(\"- Classes 3-4 (Severe/Proliferative): Prominent lesions, neovascularization\")\nprint(\"- Image quality varies significantly (different cameras, exposure, focus)\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-06T09:46:39.615038Z","iopub.execute_input":"2026-04-06T09:46:39.615935Z","iopub.status.idle":"2026-04-06T09:46:44.709028Z","shell.execute_reply.started":"2026-04-06T09:46:39.615902Z","shell.execute_reply":"2026-04-06T09:46:44.707822Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 4.4 Image Properties Analysis\nprint(\"Analysing image properties on a sample of available images...\\n\")\n\nsample_paths = subset_df['filepath'].values[:200]\nwidths, heights, aspects, sizes_kb = [], [], [], []\n\nfor p in sample_paths:\n    try:\n        img = Image.open(p)\n        w, h = img.size\n        widths.append(w)\n        heights.append(h)\n        aspects.append(w / h)\n        sizes_kb.append(os.path.getsize(p) / 1024)\n    except Exception:\n        continue\n\nif widths:\n    fig, axes = plt.subplots(1, 3, figsize=(18, 5))\n\n    axes[0].hist(widths, bins=30, color='steelblue', edgecolor='black', alpha=0.7)\n    axes[0].set_title('Image Width Distribution', fontweight='bold')\n    axes[0].set_xlabel('Width (px)')\n    axes[0].axvline(np.mean(widths), color='red', linestyle='--', label=f'Mean: {np.mean(widths):.0f}')\n    axes[0].legend()\n\n    axes[1].hist(heights, bins=30, color='coral', edgecolor='black', alpha=0.7)\n    axes[1].set_title('Image Height Distribution', fontweight='bold')\n    axes[1].set_xlabel('Height (px)')\n    axes[1].axvline(np.mean(heights), color='red', linestyle='--', label=f'Mean: {np.mean(heights):.0f}')\n    axes[1].legend()\n\n    axes[2].hist(sizes_kb, bins=30, color='mediumseagreen', edgecolor='black', alpha=0.7)\n    axes[2].set_title('File Size Distribution', fontweight='bold')\n    axes[2].set_xlabel('Size (KB)')\n    axes[2].axvline(np.mean(sizes_kb), color='red', linestyle='--', label=f'Mean: {np.mean(sizes_kb):.0f} KB')\n    axes[2].legend()\n\n    plt.tight_layout()\n    plt.savefig(os.path.join(WORKING_DIR, 'image_properties.png'), dpi=150, bbox_inches='tight')\n    plt.show()\n\n    print(f\"Width  — Min: {min(widths)}, Max: {max(widths)}, Mean: {np.mean(widths):.0f}\")\n    print(f\"Height — Min: {min(heights)}, Max: {max(heights)}, Mean: {np.mean(heights):.0f}\")\n    print(f\"Size   — Min: {min(sizes_kb):.0f} KB, Max: {max(sizes_kb):.0f} KB, Mean: {np.mean(sizes_kb):.0f} KB\")\nelse:\n    print(\"No images available for property analysis.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-06T09:47:01.716371Z","iopub.execute_input":"2026-04-06T09:47:01.716831Z","iopub.status.idle":"2026-04-06T09:47:03.031556Z","shell.execute_reply.started":"2026-04-06T09:47:01.716800Z","shell.execute_reply":"2026-04-06T09:47:03.030663Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ====================================================\n# 5. PREPROCESSING\n# ====================================================\n\n# 5.1 Ben Graham preprocessing functions\ndef crop_black_borders(img, tol=7):\n    \"\"\"Remove black padding around the circular retinal image.\"\"\"\n    if isinstance(img, Image.Image):\n        img = np.array(img)\n    if img.ndim == 3:\n        gray = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY)\n    else:\n        gray = img\n\n    mask = gray > tol\n    if not mask.any():\n        return img\n\n    coords = np.argwhere(mask)\n    y0, x0 = coords.min(axis=0)\n    y1, x1 = coords.max(axis=0) + 1\n\n    if img.ndim == 3:\n        return img[y0:y1, x0:x1]\n    return img[y0:y1, x0:x1]\n\n\ndef ben_graham_preprocess(img, sigma=10):\n    \"\"\"\n    Ben Graham's preprocessing: crop borders, resize, then\n    subtract Gaussian-blurred local average to enhance lesions.\n    \"\"\"\n    if isinstance(img, Image.Image):\n        img = np.array(img)\n\n    img = crop_black_borders(img)\n    img = cv2.resize(img, (IMG_SIZE, IMG_SIZE))\n\n    img = cv2.addWeighted(\n        img, 4,\n        cv2.GaussianBlur(img, (0, 0), sigma), -4,\n        128\n    )\n    return img\n\n\n# Visualise preprocessing effect\nprint(\"Preprocessing example:\")\nfig, axes = plt.subplots(2, 4, figsize=(16, 8))\n\nsample_files = subset_df.groupby('level').first()['filepath'].values\nfor i, fpath in enumerate(sample_files[:4]):\n    try:\n        original = np.array(Image.open(fpath))\n        processed = ben_graham_preprocess(Image.open(fpath))\n\n        axes[0, i].imshow(original)\n        axes[0, i].set_title('Original', fontweight='bold')\n        axes[0, i].axis('off')\n\n        axes[1, i].imshow(processed)\n        axes[1, i].set_title('Preprocessed', fontweight='bold')\n        axes[1, i].axis('off')\n    except Exception:\n        pass\n\nfig.suptitle('Ben Graham Preprocessing: Before vs After', fontsize=14, fontweight='bold')\nplt.tight_layout()\nplt.savefig(os.path.join(WORKING_DIR, 'preprocessing_comparison.png'), dpi=150, bbox_inches='tight')\nplt.show()\n\nprint(\"\\nBen Graham preprocessing enhances vessel visibility and reduces\")\nprint(\"variation from different cameras and lighting conditions.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-06T09:47:06.626259Z","iopub.execute_input":"2026-04-06T09:47:06.626575Z","iopub.status.idle":"2026-04-06T09:47:09.313347Z","shell.execute_reply.started":"2026-04-06T09:47:06.626550Z","shell.execute_reply":"2026-04-06T09:47:09.312217Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 5.2 Data Augmentation Transforms\nIMAGENET_MEAN = [0.485, 0.456, 0.406]\nIMAGENET_STD = [0.229, 0.224, 0.225]\n\ntrain_transforms = transforms.Compose([\n    transforms.Resize((IMG_SIZE, IMG_SIZE)),\n    transforms.RandomHorizontalFlip(p=0.5),\n    transforms.RandomVerticalFlip(p=0.5),\n    transforms.RandomRotation(degrees=30),\n    transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.05),\n    transforms.RandomAffine(degrees=0, translate=(0.05, 0.05), scale=(0.95, 1.05)),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),\n])\n\nval_transforms = transforms.Compose([\n    transforms.Resize((IMG_SIZE, IMG_SIZE)),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),\n])\n\nprint(\"Augmentation pipeline defined.\")\nprint(\"Training augmentations: HFlip, VFlip, Rotation(+/-30 deg), ColorJitter, Affine\")\nprint(\"Validation: Resize + Normalize only\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-06T09:47:14.827617Z","iopub.execute_input":"2026-04-06T09:47:14.828095Z","iopub.status.idle":"2026-04-06T09:47:14.837345Z","shell.execute_reply.started":"2026-04-06T09:47:14.828063Z","shell.execute_reply":"2026-04-06T09:47:14.836324Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 5.3 Custom PyTorch Dataset\nclass RetinopathyDataset(Dataset):\n    \"\"\"PyTorch Dataset for diabetic retinopathy classification.\"\"\"\n\n    def __init__(self, dataframe, transform=None, apply_ben_graham=True):\n        self.df = dataframe.reset_index(drop=True)\n        self.transform = transform\n        self.apply_ben_graham = apply_ben_graham\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        filepath = row['filepath']\n        label = row['level']\n\n        img = Image.open(filepath).convert('RGB')\n\n        if self.apply_ben_graham:\n            img_np = ben_graham_preprocess(img)\n            img = Image.fromarray(img_np.astype(np.uint8))\n\n        if self.transform:\n            img = self.transform(img)\n\n        return img, label\n\nprint(\"RetinopathyDataset class defined.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-06T09:47:19.332226Z","iopub.execute_input":"2026-04-06T09:47:19.333280Z","iopub.status.idle":"2026-04-06T09:47:19.340872Z","shell.execute_reply.started":"2026-04-06T09:47:19.333230Z","shell.execute_reply":"2026-04-06T09:47:19.339709Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ====================================================\n# 6. DATA SPLITTING & CLASS BALANCING\n# ====================================================\n\ntrain_df, val_df = train_test_split(\n    subset_df, test_size=VAL_SPLIT,\n    stratify=subset_df['level'], random_state=SEED\n)\ntrain_df = train_df.reset_index(drop=True)\nval_df = val_df.reset_index(drop=True)\n\nprint(f\"Training set:   {len(train_df)} images\")\nprint(f\"Validation set: {len(val_df)} images\")\nprint(f\"\\nTraining class distribution:\\n{train_df['level'].value_counts().sort_index()}\")\nprint(f\"\\nValidation class distribution:\\n{val_df['level'].value_counts().sort_index()}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-06T09:47:24.266502Z","iopub.execute_input":"2026-04-06T09:47:24.267279Z","iopub.status.idle":"2026-04-06T09:47:24.283037Z","shell.execute_reply.started":"2026-04-06T09:47:24.267238Z","shell.execute_reply":"2026-04-06T09:47:24.281948Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Compute class weights for imbalanced learning\nclass_counts = train_df['level'].value_counts().sort_index().values\nclass_weights = 1.0 / class_counts\nclass_weights = class_weights / class_weights.sum() * NUM_CLASSES\nclass_weights_tensor = torch.FloatTensor(class_weights).to(device)\n\nprint(f\"\\nClass weights (inverse frequency, normalised):\")\nfor i, (name, w) in enumerate(zip(CLASS_NAMES, class_weights)):\n    print(f\"  {name}: {w:.4f}\")\n\nsample_weights = [class_weights[label] for label in train_df['level'].values]\nsampler = WeightedRandomSampler(\n    weights=sample_weights,\n    num_samples=len(sample_weights),\n    replacement=True\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-06T09:47:29.578738Z","iopub.execute_input":"2026-04-06T09:47:29.579597Z","iopub.status.idle":"2026-04-06T09:47:29.893844Z","shell.execute_reply.started":"2026-04-06T09:47:29.579564Z","shell.execute_reply":"2026-04-06T09:47:29.892873Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Create DataLoaders\ntrain_dataset = RetinopathyDataset(train_df, transform=train_transforms, apply_ben_graham=True)\nval_dataset = RetinopathyDataset(val_df, transform=val_transforms, apply_ben_graham=True)\n\ntrain_loader = DataLoader(\n    train_dataset, batch_size=BATCH_SIZE, sampler=sampler,\n    num_workers=NUM_WORKERS, pin_memory=True, drop_last=True\n)\nval_loader = DataLoader(\n    val_dataset, batch_size=BATCH_SIZE, shuffle=False,\n    num_workers=NUM_WORKERS, pin_memory=True\n)\n\nprint(f\"Train DataLoader: {len(train_loader)} batches x {BATCH_SIZE} = {len(train_loader)*BATCH_SIZE} samples/epoch\")\nprint(f\"Val   DataLoader: {len(val_loader)} batches x {BATCH_SIZE}\")\n\n# Visualise a training batch\nbatch_imgs, batch_labels = next(iter(train_loader))\nfig, axes = plt.subplots(2, 8, figsize=(20, 6))\nfor i in range(min(16, len(batch_imgs))):\n    ax = axes[i // 8, i % 8]\n    img = batch_imgs[i].permute(1, 2, 0).numpy()\n    img = img * np.array(IMAGENET_STD) + np.array(IMAGENET_MEAN)\n    img = np.clip(img, 0, 1)\n    ax.imshow(img)\n    ax.set_title(CLASS_NAMES[batch_labels[i].item()], fontsize=9)\n    ax.axis('off')\nfig.suptitle('Sample Training Batch (augmented + preprocessed)', fontsize=14, fontweight='bold')\nplt.tight_layout()\nplt.savefig(os.path.join(WORKING_DIR, 'training_batch.png'), dpi=150, bbox_inches='tight')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-06T09:47:32.635780Z","iopub.execute_input":"2026-04-06T09:47:32.636406Z","iopub.status.idle":"2026-04-06T09:47:37.708417Z","shell.execute_reply.started":"2026-04-06T09:47:32.636374Z","shell.execute_reply":"2026-04-06T09:47:37.707055Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ====================================================\n# 7. MODEL ARCHITECTURES\n# ====================================================\n\ndef build_resnet50(num_classes=NUM_CLASSES, dropout=0.5, freeze_backbone=False):\n    \"\"\"Build ResNet50 with custom classification head.\"\"\"\n    model = models.resnet50(weights=models.ResNet50_Weights.IMAGENET1K_V1)\n    if freeze_backbone:\n        for param in model.parameters():\n            param.requires_grad = False\n    in_features = model.fc.in_features\n    model.fc = nn.Sequential(\n        nn.Dropout(p=dropout),\n        nn.Linear(in_features, 512),\n        nn.ReLU(inplace=True),\n        nn.BatchNorm1d(512),\n        nn.Dropout(p=dropout / 2),\n        nn.Linear(512, num_classes)\n    )\n    return model\n\n\ndef build_efficientnet_b3(num_classes=NUM_CLASSES, dropout=0.5, freeze_backbone=False):\n    \"\"\"Build EfficientNet-B3 with custom classification head.\"\"\"\n    model = models.efficientnet_b3(weights=models.EfficientNet_B3_Weights.IMAGENET1K_V1)\n    if freeze_backbone:\n        for param in model.parameters():\n            param.requires_grad = False\n    in_features = model.classifier[1].in_features\n    model.classifier = nn.Sequential(\n        nn.Dropout(p=dropout),\n        nn.Linear(in_features, 512),\n        nn.ReLU(inplace=True),\n        nn.BatchNorm1d(512),\n        nn.Dropout(p=dropout / 2),\n        nn.Linear(512, num_classes)\n    )\n    return model\n\n\nprint(\"=== ResNet50 Architecture (classifier head) ===\")\n_m = build_resnet50()\nprint(_m.fc)\nprint(f\"Total parameters: {sum(p.numel() for p in _m.parameters()):,}\")\nprint(f\"Trainable parameters: {sum(p.numel() for p in _m.parameters() if p.requires_grad):,}\")\n\nprint(\"\\n=== EfficientNet-B3 Architecture (classifier head) ===\")\n_m = build_efficientnet_b3()\nprint(_m.classifier)\nprint(f\"Total parameters: {sum(p.numel() for p in _m.parameters()):,}\")\nprint(f\"Trainable parameters: {sum(p.numel() for p in _m.parameters() if p.requires_grad):,}\")\ndel _m","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-06T09:47:46.587881Z","iopub.execute_input":"2026-04-06T09:47:46.588752Z","iopub.status.idle":"2026-04-06T09:47:48.710454Z","shell.execute_reply.started":"2026-04-06T09:47:46.588707Z","shell.execute_reply":"2026-04-06T09:47:48.709608Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ====================================================\n# 8. TRAINING UTILITIES\n# ====================================================\n\nclass EarlyStopping:\n    \"\"\"Stop training when validation metric stops improving.\"\"\"\n    def __init__(self, patience=5, min_delta=0.001, mode='min'):\n        self.patience = patience\n        self.min_delta = min_delta\n        self.mode = mode\n        self.counter = 0\n        self.best_score = None\n        self.early_stop = False\n\n    def __call__(self, score):\n        if self.best_score is None:\n            self.best_score = score\n            return False\n        if self.mode == 'min':\n            improved = score < (self.best_score - self.min_delta)\n        else:\n            improved = score > (self.best_score + self.min_delta)\n        if improved:\n            self.best_score = score\n            self.counter = 0\n        else:\n            self.counter += 1\n            if self.counter >= self.patience:\n                self.early_stop = True\n                return True\n        return False\n\n\ndef train_one_epoch(model, loader, criterion, optimizer, scaler, device):\n    \"\"\"Train for one epoch with mixed precision.\"\"\"\n    model.train()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n\n    for images, labels in loader:\n        images = images.to(device, non_blocking=True)\n        labels = labels.to(device, non_blocking=True)\n\n        optimizer.zero_grad()\n        with autocast():\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n\n        scaler.scale(loss).backward()\n        scaler.unscale_(optimizer)\n        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n        scaler.step(optimizer)\n        scaler.update()\n\n        running_loss += loss.item() * images.size(0)\n        _, preds = outputs.max(1)\n        total += labels.size(0)\n        correct += preds.eq(labels).sum().item()\n\n    return running_loss / total, correct / total\n\n\n@torch.no_grad()\ndef validate(model, loader, criterion, device):\n    \"\"\"Validate and return predictions + probabilities.\"\"\"\n    model.eval()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    all_preds = []\n    all_labels = []\n    all_probs = []\n\n    for images, labels in loader:\n        images = images.to(device, non_blocking=True)\n        labels = labels.to(device, non_blocking=True)\n\n        with autocast():\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n\n        running_loss += loss.item() * images.size(0)\n        probs = torch.softmax(outputs.float(), dim=1)\n        _, preds = probs.max(1)\n        total += labels.size(0)\n        correct += preds.eq(labels).sum().item()\n\n        all_preds.extend(preds.cpu().numpy())\n        all_labels.extend(labels.cpu().numpy())\n        all_probs.extend(probs.cpu().numpy())\n\n    return (\n        running_loss / total,\n        correct / total,\n        np.array(all_preds),\n        np.array(all_labels),\n        np.array(all_probs)\n    )\n\n\ndef train_model(model, train_loader, val_loader, model_name,\n                lr=LEARNING_RATE, weight_decay=WEIGHT_DECAY,\n                epochs=EPOCHS, patience=PATIENCE):\n    \"\"\"Full training loop with logging, scheduling, and early stopping.\"\"\"\n    print(f\"\\n{'='*60}\")\n    print(f\"  Training {model_name}\")\n    print(f\"  LR={lr}, WD={weight_decay}, Epochs={epochs}, Patience={patience}\")\n    print(f\"{'='*60}\")\n\n    if torch.cuda.device_count() > 1:\n        print(f\"  Using {torch.cuda.device_count()} GPUs with DataParallel\")\n        model = nn.DataParallel(model)\n    model = model.to(device)\n\n    criterion = nn.CrossEntropyLoss(weight=class_weights_tensor)\n    optimizer = optim.AdamW(\n        filter(lambda p: p.requires_grad, model.parameters()),\n        lr=lr, weight_decay=weight_decay\n    )\n    scheduler = CosineAnnealingLR(optimizer, T_max=epochs, eta_min=lr * 0.01)\n    scaler = GradScaler()\n    early_stopping = EarlyStopping(patience=patience, mode='max')\n\n    history = {\n        'train_loss': [], 'train_acc': [],\n        'val_loss': [], 'val_acc': [], 'val_kappa': []\n    }\n    best_kappa = -1\n    best_model_state = None\n    best_epoch = 0\n\n    for epoch in range(epochs):\n        start_time = time.time()\n\n        train_loss, train_acc = train_one_epoch(\n            model, train_loader, criterion, optimizer, scaler, device\n        )\n        val_loss, val_acc, val_preds, val_labels, val_probs = validate(\n            model, val_loader, criterion, device\n        )\n        val_kappa = cohen_kappa_score(val_labels, val_preds, weights='quadratic')\n\n        scheduler.step()\n        elapsed = time.time() - start_time\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['val_kappa'].append(val_kappa)\n\n        print(f\"  Epoch {epoch+1:02d}/{epochs} | \"\n              f\"Train Loss: {train_loss:.4f} Acc: {train_acc:.4f} | \"\n              f\"Val Loss: {val_loss:.4f} Acc: {val_acc:.4f} QWK: {val_kappa:.4f} | \"\n              f\"{elapsed:.1f}s\")\n\n        if val_kappa > best_kappa:\n            best_kappa = val_kappa\n            best_epoch = epoch + 1\n            if isinstance(model, nn.DataParallel):\n                best_model_state = {k: v.cpu().clone() for k, v in model.module.state_dict().items()}\n            else:\n                best_model_state = {k: v.cpu().clone() for k, v in model.state_dict().items()}\n\n        if early_stopping(val_kappa):\n            print(f\"  Early stopping triggered at epoch {epoch+1}\")\n            break\n\n    print(f\"\\n  Best epoch: {best_epoch} with QWK = {best_kappa:.4f}\")\n\n    if best_model_state is not None:\n        if isinstance(model, nn.DataParallel):\n            model.module.load_state_dict(best_model_state)\n        else:\n            model.load_state_dict(best_model_state)\n\n    return model, history, best_kappa\n\nprint(\"Training utilities defined.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-06T09:47:53.549741Z","iopub.execute_input":"2026-04-06T09:47:53.550785Z","iopub.status.idle":"2026-04-06T09:47:53.583361Z","shell.execute_reply.started":"2026-04-06T09:47:53.550721Z","shell.execute_reply":"2026-04-06T09:47:53.582371Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ====================================================\n# 9. HYPERPARAMETER TUNING & TRAINING - RESNET50\n# ====================================================\n\nhp_configs_resnet = [\n    {'lr': 1e-4, 'weight_decay': 1e-4, 'dropout': 0.5},\n    {'lr': 3e-4, 'weight_decay': 1e-5, 'dropout': 0.3},\n]\n\nprint(f\"Hyperparameter search for ResNet50 ({len(hp_configs_resnet)} configs x {HP_EPOCHS} epochs):\")\nprint(\"-\" * 70)\n\nhp_results_resnet = []\nfor i, hp in enumerate(hp_configs_resnet):\n    print(f\"\\nConfig {i+1}/{len(hp_configs_resnet)}: {hp}\")\n    model = build_resnet50(dropout=hp['dropout'], freeze_backbone=True)\n    model, hist, best_kappa = train_model(\n        model, train_loader, val_loader,\n        model_name=f\"ResNet50-HP{i+1}\",\n        lr=hp['lr'], weight_decay=hp['weight_decay'],\n        epochs=HP_EPOCHS, patience=HP_EPOCHS\n    )\n    hp_results_resnet.append({**hp, 'best_kappa': best_kappa})\n    del model\n    torch.cuda.empty_cache()\n\nhp_df_resnet = pd.DataFrame(hp_results_resnet)\nprint(\"\\n\\nHyperparameter Search Results (ResNet50):\")\nprint(hp_df_resnet.to_string(index=False))\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-06T09:47:59.485136Z","iopub.execute_input":"2026-04-06T09:47:59.485670Z","iopub.status.idle":"2026-04-06T09:49:30.121382Z","shell.execute_reply.started":"2026-04-06T09:47:59.485638Z","shell.execute_reply":"2026-04-06T09:49:30.120367Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 9.2 Full training with best hyperparameters\nbest_hp_resnet = max(hp_results_resnet, key=lambda x: x['best_kappa'])\nprint(f\"\\nBest hyperparameters for ResNet50: {best_hp_resnet}\")\n\nresnet_model = build_resnet50(dropout=best_hp_resnet['dropout'])\nresnet_model, resnet_history, resnet_best_kappa = train_model(\n    resnet_model, train_loader, val_loader,\n    model_name=\"ResNet50 (Final)\",\n    lr=best_hp_resnet['lr'],\n    weight_decay=best_hp_resnet['weight_decay'],\n    epochs=EPOCHS, patience=PATIENCE\n)\n\nresnet_save_path = os.path.join(WORKING_DIR, 'resnet50_best.pth')\nif isinstance(resnet_model, nn.DataParallel):\n    torch.save(resnet_model.module.state_dict(), resnet_save_path)\nelse:\n    torch.save(resnet_model.state_dict(), resnet_save_path)\nprint(f\"ResNet50 model saved to {resnet_save_path}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-06T09:50:21.899030Z","iopub.execute_input":"2026-04-06T09:50:21.900041Z","iopub.status.idle":"2026-04-06T09:56:03.545092Z","shell.execute_reply.started":"2026-04-06T09:50:21.900001Z","shell.execute_reply":"2026-04-06T09:56:03.544283Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 9.3 Training curves - ResNet50\nfig, axes = plt.subplots(1, 3, figsize=(18, 5))\n\naxes[0].plot(resnet_history['train_loss'], label='Train Loss', marker='o', markersize=4)\naxes[0].plot(resnet_history['val_loss'], label='Val Loss', marker='s', markersize=4)\naxes[0].set_title('ResNet50 - Loss', fontweight='bold')\naxes[0].set_xlabel('Epoch')\naxes[0].set_ylabel('Loss')\naxes[0].legend()\naxes[0].grid(True)\n\naxes[1].plot(resnet_history['train_acc'], label='Train Acc', marker='o', markersize=4)\naxes[1].plot(resnet_history['val_acc'], label='Val Acc', marker='s', markersize=4)\naxes[1].set_title('ResNet50 - Accuracy', fontweight='bold')\naxes[1].set_xlabel('Epoch')\naxes[1].set_ylabel('Accuracy')\naxes[1].legend()\naxes[1].grid(True)\n\naxes[2].plot(resnet_history['val_kappa'], label='Val QWK', marker='D', markersize=4, color='green')\naxes[2].set_title('ResNet50 - Quadratic Weighted Kappa', fontweight='bold')\naxes[2].set_xlabel('Epoch')\naxes[2].set_ylabel('QWK')\naxes[2].legend()\naxes[2].grid(True)\n\nplt.suptitle('ResNet50 Training Curves', fontsize=14, fontweight='bold')\nplt.tight_layout()\nplt.savefig(os.path.join(WORKING_DIR, 'resnet50_training_curves.png'), dpi=150, bbox_inches='tight')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-06T09:56:21.046151Z","iopub.execute_input":"2026-04-06T09:56:21.046890Z","iopub.status.idle":"2026-04-06T09:56:22.254959Z","shell.execute_reply.started":"2026-04-06T09:56:21.046852Z","shell.execute_reply":"2026-04-06T09:56:22.254151Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ====================================================\n# 10. HYPERPARAMETER TUNING & TRAINING - EFFICIENTNET-B3\n# ====================================================\n# Reuse the best hyperparameters discovered from the ResNet50 HP search\n# to avoid a second expensive tuning loop. Both architectures respond\n# similarly to LR / weight-decay / dropout ranges, so the transfer is valid.\n\nbest_hp_effnet = best_hp_resnet.copy()\nprint(\"EfficientNet-B3 — reusing best HPs from ResNet50 search (same config space):\")\nprint(f\"  lr={best_hp_effnet['lr']}, weight_decay={best_hp_effnet['weight_decay']}, dropout={best_hp_effnet['dropout']}\")\n\n# Freeze backbone to speed up training (only classifier head is trained)\neffnet_model = build_efficientnet_b3(dropout=best_hp_effnet['dropout'], freeze_backbone=True)\neffnet_model, effnet_history, effnet_best_kappa = train_model(\n    effnet_model, train_loader, val_loader,\n    model_name=\"EfficientNet-B3 (Final)\",\n    lr=best_hp_effnet['lr'],\n    weight_decay=best_hp_effnet['weight_decay'],\n    epochs=1, patience=2\n)\n\neffnet_save_path = os.path.join(WORKING_DIR, 'efficientnet_b3_best.pth')\nif isinstance(effnet_model, nn.DataParallel):\n    torch.save(effnet_model.module.state_dict(), effnet_save_path)\nelse:\n    torch.save(effnet_model.state_dict(), effnet_save_path)\nprint(f\"EfficientNet-B3 model saved to {effnet_save_path}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-06T10:05:38.482811Z","iopub.execute_input":"2026-04-06T10:05:38.483347Z","iopub.status.idle":"2026-04-06T10:09:50.604850Z","shell.execute_reply.started":"2026-04-06T10:05:38.483308Z","shell.execute_reply":"2026-04-06T10:09:50.603721Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 10.3 Training curves - EfficientNet-B3\nfig, axes = plt.subplots(1, 3, figsize=(18, 5))\n\naxes[0].plot(effnet_history['train_loss'], label='Train Loss', marker='o', markersize=4)\naxes[0].plot(effnet_history['val_loss'], label='Val Loss', marker='s', markersize=4)\naxes[0].set_title('EfficientNet-B3 - Loss', fontweight='bold')\naxes[0].set_xlabel('Epoch')\naxes[0].set_ylabel('Loss')\naxes[0].legend()\naxes[0].grid(True)\n\naxes[1].plot(effnet_history['train_acc'], label='Train Acc', marker='o', markersize=4)\naxes[1].plot(effnet_history['val_acc'], label='Val Acc', marker='s', markersize=4)\naxes[1].set_title('EfficientNet-B3 - Accuracy', fontweight='bold')\naxes[1].set_xlabel('Epoch')\naxes[1].set_ylabel('Accuracy')\naxes[1].legend()\naxes[1].grid(True)\n\naxes[2].plot(effnet_history['val_kappa'], label='Val QWK', marker='D', markersize=4, color='green')\naxes[2].set_title('EfficientNet-B3 - Quadratic Weighted Kappa', fontweight='bold')\naxes[2].set_xlabel('Epoch')\naxes[2].set_ylabel('QWK')\naxes[2].legend()\naxes[2].grid(True)\n\nplt.suptitle('EfficientNet-B3 Training Curves', fontsize=14, fontweight='bold')\nplt.tight_layout()\nplt.savefig(os.path.join(WORKING_DIR, 'efficientnet_b3_training_curves.png'), dpi=150, bbox_inches='tight')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-06T10:10:05.612104Z","iopub.execute_input":"2026-04-06T10:10:05.613086Z","iopub.status.idle":"2026-04-06T10:10:06.678645Z","shell.execute_reply.started":"2026-04-06T10:10:05.613045Z","shell.execute_reply":"2026-04-06T10:10:06.677712Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ====================================================\n# 11. COMPREHENSIVE EVALUATION\n# ====================================================\n\ndef evaluate_model(model, val_loader, model_name, device):\n    \"\"\"Run full evaluation and return metrics + predictions.\"\"\"\n    criterion = nn.CrossEntropyLoss(weight=class_weights_tensor)\n    val_loss, val_acc, preds, labels, probs = validate(model, val_loader, criterion, device)\n\n    qwk = cohen_kappa_score(labels, preds, weights='quadratic')\n    precision, recall, f1, support = precision_recall_fscore_support(\n        labels, preds, average=None, labels=list(range(NUM_CLASSES))\n    )\n    macro_f1 = precision_recall_fscore_support(labels, preds, average='macro')[2]\n    weighted_f1 = precision_recall_fscore_support(labels, preds, average='weighted')[2]\n\n    print(f\"\\n{'='*60}\")\n    print(f\"  {model_name} - Evaluation Results\")\n    print(f\"{'='*60}\")\n    print(f\"  Validation Loss:     {val_loss:.4f}\")\n    print(f\"  Validation Accuracy: {val_acc:.4f}\")\n    print(f\"  Quadratic W. Kappa:  {qwk:.4f}\")\n    print(f\"  Macro F1-Score:      {macro_f1:.4f}\")\n    print(f\"  Weighted F1-Score:   {weighted_f1:.4f}\")\n\n    return {\n        'model_name': model_name,\n        'val_loss': val_loss,\n        'val_acc': val_acc,\n        'qwk': qwk,\n        'macro_f1': macro_f1,\n        'weighted_f1': weighted_f1,\n        'preds': preds,\n        'labels': labels,\n        'probs': probs,\n        'per_class_precision': precision,\n        'per_class_recall': recall,\n        'per_class_f1': f1,\n    }\n\nresnet_results = evaluate_model(resnet_model, val_loader, \"ResNet50\", device)\neffnet_results = evaluate_model(effnet_model, val_loader, \"EfficientNet-B3\", device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-06T10:10:18.282002Z","iopub.execute_input":"2026-04-06T10:10:18.282784Z","iopub.status.idle":"2026-04-06T10:11:13.785292Z","shell.execute_reply.started":"2026-04-06T10:10:18.282743Z","shell.execute_reply":"2026-04-06T10:11:13.784086Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 11.1 Confusion Matrices\nfig, axes = plt.subplots(1, 2, figsize=(18, 7))\n\nfor ax, results in zip(axes, [resnet_results, effnet_results]):\n    cm = confusion_matrix(results['labels'], results['preds'], labels=list(range(NUM_CLASSES)))\n    cm_pct = cm.astype('float') / cm.sum(axis=1, keepdims=True) * 100\n\n    sns.heatmap(cm, annot=True, fmt='d', cmap='Blues', ax=ax,\n                xticklabels=CLASS_NAMES, yticklabels=CLASS_NAMES,\n                cbar_kws={'label': 'Count'})\n    for i in range(NUM_CLASSES):\n        for j in range(NUM_CLASSES):\n            ax.text(j + 0.5, i + 0.75, f'({cm_pct[i, j]:.0f}%)',\n                    ha='center', va='center', fontsize=8, color='gray')\n\n    ax.set_title(f\"{results['model_name']}\\nQWK={results['qwk']:.4f} | Acc={results['val_acc']:.4f}\",\n                 fontweight='bold', fontsize=12)\n    ax.set_xlabel('Predicted')\n    ax.set_ylabel('True')\n\nplt.suptitle('Confusion Matrices', fontsize=14, fontweight='bold')\nplt.tight_layout()\nplt.savefig(os.path.join(WORKING_DIR, 'confusion_matrices.png'), dpi=150, bbox_inches='tight')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-06T10:11:23.823026Z","iopub.execute_input":"2026-04-06T10:11:23.823368Z","iopub.status.idle":"2026-04-06T10:11:25.210659Z","shell.execute_reply.started":"2026-04-06T10:11:23.823336Z","shell.execute_reply":"2026-04-06T10:11:25.209719Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 11.2 Detailed Classification Reports\nfor results in [resnet_results, effnet_results]:\n    print(f\"\\n{'='*60}\")\n    print(f\"  Classification Report - {results['model_name']}\")\n    print(f\"{'='*60}\")\n    print(classification_report(\n        results['labels'], results['preds'],\n        target_names=CLASS_NAMES, digits=4,\n        labels=list(range(NUM_CLASSES))\n    ))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-06T10:12:31.478502Z","iopub.execute_input":"2026-04-06T10:12:31.479415Z","iopub.status.idle":"2026-04-06T10:12:31.503741Z","shell.execute_reply.started":"2026-04-06T10:12:31.479382Z","shell.execute_reply":"2026-04-06T10:12:31.502718Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 11.3 Per-Class Metrics Comparison\nfig, axes = plt.subplots(1, 3, figsize=(18, 6))\nx = np.arange(NUM_CLASSES)\nwidth = 0.35\n\nmetrics = [\n    ('per_class_precision', 'Precision'),\n    ('per_class_recall', 'Recall'),\n    ('per_class_f1', 'F1-Score')\n]\n\nfor ax, (metric_key, metric_name) in zip(axes, metrics):\n    vals_resnet = resnet_results[metric_key]\n    vals_effnet = effnet_results[metric_key]\n    ax.bar(x - width/2, vals_resnet, width, label='ResNet50', color='steelblue', edgecolor='black')\n    ax.bar(x + width/2, vals_effnet, width, label='EfficientNet-B3', color='coral', edgecolor='black')\n    ax.set_title(f'Per-Class {metric_name}', fontweight='bold')\n    ax.set_xlabel('DR Severity')\n    ax.set_ylabel(metric_name)\n    ax.set_xticks(x)\n    ax.set_xticklabels(CLASS_NAMES, rotation=20, ha='right')\n    ax.legend()\n    ax.set_ylim(0, 1.1)\n    ax.grid(axis='y', alpha=0.3)\n\nplt.suptitle('Per-Class Metrics Comparison', fontsize=14, fontweight='bold')\nplt.tight_layout()\nplt.savefig(os.path.join(WORKING_DIR, 'per_class_metrics.png'), dpi=150, bbox_inches='tight')\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-06T10:12:49.812978Z","iopub.execute_input":"2026-04-06T10:12:49.813831Z","iopub.status.idle":"2026-04-06T10:12:50.893773Z","shell.execute_reply.started":"2026-04-06T10:12:49.813778Z","shell.execute_reply":"2026-04-06T10:12:50.892920Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 11.4 ROC Curves (One-vs-Rest)\nfig, axes = plt.subplots(1, 2, figsize=(16, 7))\n\nfor ax, results in zip(axes, [resnet_results, effnet_results]):\n    labels_bin = label_binarize(results['labels'], classes=list(range(NUM_CLASSES)))\n    colors_roc = plt.cm.Set1(np.linspace(0, 1, NUM_CLASSES))\n\n    for cls in range(NUM_CLASSES):\n        if labels_bin[:, cls].sum() > 0:\n            fpr, tpr, _ = roc_curve(labels_bin[:, cls], results['probs'][:, cls])\n            roc_auc = auc(fpr, tpr)\n            ax.plot(fpr, tpr, color=colors_roc[cls], lw=2,\n                    label=f'{CLASS_NAMES[cls]} (AUC={roc_auc:.3f})')\n\n    ax.plot([0, 1], [0, 1], 'k--', lw=1, alpha=0.5)\n    ax.set_title(f\"{results['model_name']} - ROC Curves\", fontweight='bold')\n    ax.set_xlabel('False Positive Rate')\n    ax.set_ylabel('True Positive Rate')\n    ax.legend(loc='lower right', fontsize=9)\n    ax.grid(True, alpha=0.3)\n\nplt.suptitle('ROC Curves (One-vs-Rest)', fontsize=14, fontweight='bold')\nplt.tight_layout()\nplt.savefig(os.path.join(WORKING_DIR, 'roc_curves.png'), dpi=150, bbox_inches='tight')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-06T10:13:02.825147Z","iopub.execute_input":"2026-04-06T10:13:02.825907Z","iopub.status.idle":"2026-04-06T10:13:03.840133Z","shell.execute_reply.started":"2026-04-06T10:13:02.825876Z","shell.execute_reply":"2026-04-06T10:13:03.839441Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 11.5 Model Comparison Summary\ncomparison_data = {\n    'Metric': ['Validation Accuracy', 'Validation Loss',\n               'Quadratic Weighted Kappa', 'Macro F1-Score', 'Weighted F1-Score'],\n    'ResNet50': [\n        f\"{resnet_results['val_acc']:.4f}\",\n        f\"{resnet_results['val_loss']:.4f}\",\n        f\"{resnet_results['qwk']:.4f}\",\n        f\"{resnet_results['macro_f1']:.4f}\",\n        f\"{resnet_results['weighted_f1']:.4f}\",\n    ],\n    'EfficientNet-B3': [\n        f\"{effnet_results['val_acc']:.4f}\",\n        f\"{effnet_results['val_loss']:.4f}\",\n        f\"{effnet_results['qwk']:.4f}\",\n        f\"{effnet_results['macro_f1']:.4f}\",\n        f\"{effnet_results['weighted_f1']:.4f}\",\n    ]\n}\ncomparison_df = pd.DataFrame(comparison_data)\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"           MODEL COMPARISON SUMMARY\")\nprint(\"=\" * 70)\nprint(comparison_df.to_string(index=False))\nprint(\"=\" * 70)\n\nif resnet_results['qwk'] > effnet_results['qwk']:\n    best = 'ResNet50'\n    best_qwk = resnet_results['qwk']\nelse:\n    best = 'EfficientNet-B3'\n    best_qwk = effnet_results['qwk']\n\nprint(f\"\\nBest model by QWK: {best} (QWK = {best_qwk:.4f})\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-06T10:13:18.702268Z","iopub.execute_input":"2026-04-06T10:13:18.702972Z","iopub.status.idle":"2026-04-06T10:13:18.714086Z","shell.execute_reply.started":"2026-04-06T10:13:18.702942Z","shell.execute_reply":"2026-04-06T10:13:18.712974Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 11.6 Comparison bar chart\nfig, ax = plt.subplots(figsize=(10, 6))\n\nmetrics_names = ['Accuracy', 'QWK', 'Macro F1', 'Weighted F1']\nresnet_vals = [resnet_results['val_acc'], resnet_results['qwk'],\n               resnet_results['macro_f1'], resnet_results['weighted_f1']]\neffnet_vals = [effnet_results['val_acc'], effnet_results['qwk'],\n               effnet_results['macro_f1'], effnet_results['weighted_f1']]\n\nx = np.arange(len(metrics_names))\nwidth = 0.35\n\nbars1 = ax.bar(x - width/2, resnet_vals, width, label='ResNet50',\n               color='steelblue', edgecolor='black')\nbars2 = ax.bar(x + width/2, effnet_vals, width, label='EfficientNet-B3',\n               color='coral', edgecolor='black')\n\nax.set_ylabel('Score')\nax.set_title('Model Performance Comparison', fontweight='bold', fontsize=14)\nax.set_xticks(x)\nax.set_xticklabels(metrics_names)\nax.legend()\nax.set_ylim(0, 1.1)\nax.grid(axis='y', alpha=0.3)\n\nfor bars in [bars1, bars2]:\n    for bar in bars:\n        height = bar.get_height()\n        ax.text(bar.get_x() + bar.get_width()/2., height + 0.01,\n                f'{height:.3f}', ha='center', va='bottom', fontsize=10)\n\nplt.tight_layout()\nplt.savefig(os.path.join(WORKING_DIR, 'model_comparison.png'), dpi=150, bbox_inches='tight')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-06T10:13:39.921120Z","iopub.execute_input":"2026-04-06T10:13:39.922104Z","iopub.status.idle":"2026-04-06T10:13:40.372830Z","shell.execute_reply.started":"2026-04-06T10:13:39.922072Z","shell.execute_reply":"2026-04-06T10:13:40.372012Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ## 10. SUMMARY & CONCLUSIONS\n#\n# ### Implementation Summary\n# - **Dataset:** Diabetic Retinopathy Detection (Kaggle) - stratified subset of retinal images\n# - **Preprocessing:** Ben Graham's colour normalisation, cropping, resizing to 224x224\n# - **Augmentation:** Random flips, rotations, colour jitter, affine transforms\n# - **Class Imbalance:** Handled via class-weighted cross-entropy loss and weighted random sampling\n#\n# ### Models\n# | Aspect | ResNet50 | EfficientNet-B3 |\n# |--------|----------|-----------------|\n# | Architecture | 50-layer residual network | Compound-scaled CNN |\n# | Transfer Learning | ImageNet-1K pretrained | ImageNet-1K pretrained |\n# | Regularisation | Dropout + weight decay | Dropout + weight decay |\n# | Optimiser | AdamW + CosineAnnealing | AdamW + CosineAnnealing |\n# | Training | Mixed-precision (FP16) | Mixed-precision (FP16) |\n#\n# ### Key Findings\n# 1. **Class imbalance** is the primary challenge - \"No DR\" (class 0) dominates at ~73%\n# 2. **Ben Graham preprocessing** significantly improves contrast for detecting microaneurysms and lesions\n# 3. **Hyperparameter tuning** shows that moderate learning rates (1e-4 to 3e-4) with light regularisation work best\n# 4. **QWK is the preferred metric** for this ordinal classification task since it penalises predictions far from the true label more heavily\n# 5. Both models generalise well; the comparison table above identifies the better-performing architecture.\n#\n# ### Evaluation Techniques Justification\n# - **Accuracy:** Basic overall correctness measure\n# - **QWK:** Appropriate for ordinal data; standard metric for DR severity grading\n# - **Precision/Recall/F1:** Important to assess per-class performance given imbalance\n# - **ROC-AUC:** Evaluates discriminative power independent of threshold","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-06T10:14:38.968959Z","iopub.execute_input":"2026-04-06T10:14:38.969724Z","iopub.status.idle":"2026-04-06T10:14:38.974511Z","shell.execute_reply.started":"2026-04-06T10:14:38.969691Z","shell.execute_reply":"2026-04-06T10:14:38.973604Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"\\n\" + \"=\" * 70)\nprint(\"  Notebook execution complete.\")\nprint(\"  All plots saved to:\", WORKING_DIR)\nprint(\"=\" * 70)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-06T10:14:21.162905Z","iopub.execute_input":"2026-04-06T10:14:21.163944Z","iopub.status.idle":"2026-04-06T10:14:21.169865Z","shell.execute_reply.started":"2026-04-06T10:14:21.163908Z","shell.execute_reply":"2026-04-06T10:14:21.168796Z"}},"outputs":[],"execution_count":null}]}