{"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":14774,"databundleVersionId":875431},{"sourceType":"datasetVersion","sourceId":13092692,"datasetId":8292992,"databundleVersionId":13775082},{"sourceType":"datasetVersion","sourceId":5835229,"datasetId":3354256,"databundleVersionId":5912207},{"sourceType":"modelInstanceVersion","sourceId":834145,"databundleVersionId":16710449,"modelInstanceId":634527,"modelId":646511},{"sourceType":"modelInstanceVersion","sourceId":848382,"databundleVersionId":16921981,"modelInstanceId":645062,"modelId":657010},{"sourceType":"modelInstanceVersion","sourceId":838918,"databundleVersionId":16780547,"modelInstanceId":638137,"modelId":650123},{"sourceType":"modelInstanceVersion","sourceId":840281,"databundleVersionId":16800059,"modelInstanceId":639215,"modelId":651215},{"sourceType":"modelInstanceVersion","sourceId":841676,"databundleVersionId":16822087,"modelInstanceId":640327,"modelId":652342}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\n\n\nAPTOS_IMAGES = \"/kaggle/input/aptos2019-blindness-detection/train_images\"\n\n\n\nif not os.path.exists(APTOS_IMAGES):\n\n    for root, dirs, files in os.walk('/kaggle/input'):\n\n        if 'train_images' in dirs:\n\n            APTOS_IMAGES = os.path.join(root, 'train_images')\n\n            break\n\n\n\nprint(f\"Final Path: {APTOS_IMAGES}\")\n\nif os.path.exists(APTOS_IMAGES):\n\n    print(\"Success! Images found.\")\n\n    print(\"Sample files:\", os.listdir(APTOS_IMAGES)[:5])\n\nelse:\n\n    print(\"Still not found. Please check if the dataset is attached correctly.\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:33.121265Z","iopub.execute_input":"2026-05-01T11:21:33.121653Z","iopub.status.idle":"2026-05-01T11:21:33.133795Z","shell.execute_reply.started":"2026-05-01T11:21:33.121575Z","shell.execute_reply":"2026-05-01T11:21:33.133076Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\nfound_path = \"\"\nfor root, dirs, files in os.walk('/kaggle/input'):\n    if 'indian-diabetic-retinopathy-image-dataset' in root.lower() and 'images' in root.lower():\n        found_path = root\n        break\n\nif found_path:\n    IDRID_IMAGES = found_path\n    print(f\"Success! IDRID Path found: {IDRID_IMAGES}\")\n    print(\"Sample files:\", os.listdir(IDRID_IMAGES)[:3])\nelse:\n    \n    IDRID_IMAGES = \"/kaggle/input/indian-diabetic-retinopathy-image-d/images\"\n    print(\"Manual path set. Please check if it works.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:33.135123Z","iopub.execute_input":"2026-05-01T11:21:33.135434Z","iopub.status.idle":"2026-05-01T11:21:36.778941Z","shell.execute_reply.started":"2026-05-01T11:21:33.135410Z","shell.execute_reply":"2026-05-01T11:21:36.778273Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nprint(\"APTOS sample files:\")\nprint(os.listdir(APTOS_IMAGES)[:5])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:36.779723Z","iopub.execute_input":"2026-05-01T11:21:36.779934Z","iopub.status.idle":"2026-05-01T11:21:36.785769Z","shell.execute_reply.started":"2026-05-01T11:21:36.779912Z","shell.execute_reply":"2026-05-01T11:21:36.784970Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"PHASE 2 — STEP 2: EXPLORATORY DATA ANALYSIS (EDA)\n🎯 Goal of this step:\n\nWe will understand:\n\nClass distribution (VERY IMPORTANT for your research gap)\nDataset imbalance\nVisual differences (APTOS vs IDRiD)\nBasic statistics\n\nThis step is research foundation level.","metadata":{}},{"cell_type":"markdown","source":"STEP 2.1 — LOAD APTOS LABELS","metadata":{}},{"cell_type":"code","source":"import pandas as pd\n\naptos_labels = pd.read_csv(f\"{APTOS_PATH}/train.csv\")\naptos_labels.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:36.787174Z","iopub.execute_input":"2026-05-01T11:21:36.787364Z","iopub.status.idle":"2026-05-01T11:21:36.809621Z","shell.execute_reply.started":"2026-05-01T11:21:36.787344Z","shell.execute_reply":"2026-05-01T11:21:36.808830Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"STEP 2.2 — CHECK APTOS CLASS DISTRIBUTION","metadata":{}},{"cell_type":"code","source":"aptos_labels['diagnosis'].value_counts().sort_index()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:36.810358Z","iopub.execute_input":"2026-05-01T11:21:36.810632Z","iopub.status.idle":"2026-05-01T11:21:36.817130Z","shell.execute_reply.started":"2026-05-01T11:21:36.810575Z","shell.execute_reply":"2026-05-01T11:21:36.816344Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Class\tMeaning\n0\tNo DR\n1\tMild\n2\tModerate\n3\tSevere\n4\tProliferative","metadata":{}},{"cell_type":"markdown","source":"STEP 2.3 — PLOT APTOS DISTRIBUTION","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\naptos_labels['diagnosis'].value_counts().sort_index().plot(kind='bar')\nplt.title(\"APTOS Class Distribution\")\nplt.xlabel(\"DR Grade\")\nplt.ylabel(\"Number of Images\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:36.817989Z","iopub.execute_input":"2026-05-01T11:21:36.818252Z","iopub.status.idle":"2026-05-01T11:21:36.947085Z","shell.execute_reply.started":"2026-05-01T11:21:36.818231Z","shell.execute_reply":"2026-05-01T11:21:36.946475Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"STEP 2.4 — APTOS VISUALIZE","metadata":{}},{"cell_type":"code","source":"import cv2\nimport matplotlib.pyplot as plt\n\nsample_imgs = aptos_labels.sample(5)\nplt.figure(figsize=(15, 5))\n\nfor i, row in enumerate(sample_imgs.itertuples()):\n    img_path = f\"{APTOS_IMAGES}/{row.id_code}.png\"\n    img = cv2.imread(img_path)\n    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    \n    plt.subplot(1, 5, i+1)\n    plt.imshow(img)\n    plt.title(f\"Class: {row.diagnosis}\")\n    plt.axis(\"off\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:36.947896Z","iopub.execute_input":"2026-05-01T11:21:36.948192Z","iopub.status.idle":"2026-05-01T11:21:38.785618Z","shell.execute_reply.started":"2026-05-01T11:21:36.948165Z","shell.execute_reply":"2026-05-01T11:21:38.784860Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"TEP 2.5 — VISUALIZE SAMPLE IMAGES (VERY IMPORTANT)","metadata":{}},{"cell_type":"code","source":"import os\nimport random\nimport cv2\nimport matplotlib.pyplot as plt\n\nIDRID_TRAIN_PATH = os.path.join(IDRID_IMAGES, \"a. Training Set\")\n\nfiles = [f for f in os.listdir(IDRID_TRAIN_PATH) if f.endswith(('.jpg', '.png', '.jpeg'))]\n\nif len(files) > 0:\n    sample_files = random.sample(files, min(len(files), 5))\n    plt.figure(figsize=(12,6))\n\n    for i, f in enumerate(sample_files):\n        img_path = os.path.join(IDRID_TRAIN_PATH, f)\n        img = cv2.imread(img_path)\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        \n        plt.subplot(1, 5, i+1)\n        plt.imshow(img)\n        plt.title(\"IDRID Sample\")\n        plt.axis(\"off\")\n    plt.show()\nelse:\n    print(\"No images found in the specified folder!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:38.786770Z","iopub.execute_input":"2026-05-01T11:21:38.787185Z","iopub.status.idle":"2026-05-01T11:21:41.303442Z","shell.execute_reply.started":"2026-05-01T11:21:38.787149Z","shell.execute_reply":"2026-05-01T11:21:41.302627Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import cv2\nimport matplotlib.pyplot as plt\n\nsample_imgs = aptos_labels.sample(5)\n\nplt.figure(figsize=(12,6))\n\nfor i, row in enumerate(sample_imgs.itertuples()):\n    img_path = f\"{APTOS_IMAGES}/{row.id_code}.png\"\n    \n    img = cv2.imread(img_path)\n    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n\n    plt.subplot(1,5,i+1)\n    plt.imshow(img)\n    plt.title(f\"Class: {row.diagnosis}\")\n    plt.axis(\"off\")\n\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:41.306297Z","iopub.execute_input":"2026-05-01T11:21:41.306739Z","iopub.status.idle":"2026-05-01T11:21:41.938563Z","shell.execute_reply.started":"2026-05-01T11:21:41.306712Z","shell.execute_reply":"2026-05-01T11:21:41.937793Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"✔ APTOS analysis\nClass distribution ✔\nLabels loaded ✔\nDataset structure understood ✔\n✔ IDRiD analysis\n455 images confirmed ✔\nSample images visible ✔","metadata":{}},{"cell_type":"markdown","source":"CLASS IMBALANCE EXISTS\nClass\tCount\n0\t1805\n1\t370\n2\t999\n3\t193\n4\t295\n\n👉 Severe imbalance confirmed\n👉 This directly supports your Focal Loss + CBAM motivation","metadata":{}},{"cell_type":"markdown","source":"**PHASE 3 — PREPROCESSING PIPELINE (START)**\n Goal:\n\nConvert raw eye images into clean, model-ready inputs so EfficientNetB5 + CBAM can learn properly.\n\nThis is one of the most important phases for your research quality.","metadata":{}},{"cell_type":"markdown","source":"STEP 3.1 — CIRCULAR CROP (REMOVE BLACK BORDERS)\nWhy we do this:\n\nFundus images often have:\n\nblack circular borders\nuseless background noise\n\nWe remove it so model focuses only on retina.","metadata":{}},{"cell_type":"code","source":"import cv2\nimport numpy as np\n\ndef circular_crop(img):\n    gray = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)\n    \n    # threshold to find retina region\n    _, thresh = cv2.threshold(gray, 10, 255, cv2.THRESH_BINARY)\n    \n    contours, _ = cv2.findContours(thresh, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n    \n    if len(contours) == 0:\n        return img\n    \n    cnt = max(contours, key=cv2.contourArea)\n    x, y, w, h = cv2.boundingRect(cnt)\n    \n    cropped = img[y:y+h, x:x+w]\n    return cropped","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:41.939828Z","iopub.execute_input":"2026-05-01T11:21:41.940168Z","iopub.status.idle":"2026-05-01T11:21:41.945473Z","shell.execute_reply.started":"2026-05-01T11:21:41.940131Z","shell.execute_reply":"2026-05-01T11:21:41.944786Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"STEP 3.2 — BEN GRAHAM PREPROCESSING\nWhy:\n\nFix:\n\nlighting differences\ncamera variations\nbrightness inconsistency\n\nThis is VERY important for your multi-dataset fusion idea","metadata":{}},{"cell_type":"code","source":"def ben_graham_preprocess(img, sigmaX=10):\n    img = cv2.addWeighted(img, 4, cv2.GaussianBlur(img, (0,0), sigmaX), -4, 128)\n    return img","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:41.946353Z","iopub.execute_input":"2026-05-01T11:21:41.946573Z","iopub.status.idle":"2026-05-01T11:21:41.960507Z","shell.execute_reply.started":"2026-05-01T11:21:41.946550Z","shell.execute_reply":"2026-05-01T11:21:41.959986Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"STEP 3.3 — FULL PREPROCESS PIPELINE","metadata":{}},{"cell_type":"code","source":"IMG_SIZE = 456\n\ndef preprocess_image(path):\n    img = cv2.imread(path)\n    \n    # Step 1: circular crop\n    img = circular_crop(img)\n    \n    # Step 2: Ben Graham\n    img = ben_graham_preprocess(img)\n    \n    # Step 3: resize\n    img = cv2.resize(img, (IMG_SIZE, IMG_SIZE))\n    \n    # Step 4: convert RGB\n    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    \n    # Step 5: normalize (0–1)\n    img = img / 255.0\n    \n    return img","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:41.961366Z","iopub.execute_input":"2026-05-01T11:21:41.961651Z","iopub.status.idle":"2026-05-01T11:21:41.973575Z","shell.execute_reply.started":"2026-05-01T11:21:41.961617Z","shell.execute_reply":"2026-05-01T11:21:41.973037Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"STEP 3.4 — TEST PREPROCESSING","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\nsample = aptos_labels.iloc[0]\nimg_path = f\"{APTOS_IMAGES}/{sample.id_code}.png\"\n\noriginal = cv2.imread(img_path)\nprocessed = preprocess_image(img_path)\n\nplt.figure(figsize=(10,5))\n\nplt.subplot(1,2,1)\nplt.imshow(cv2.cvtColor(original, cv2.COLOR_BGR2RGB))\nplt.title(\"Original\")\nplt.axis(\"off\")\n\nplt.subplot(1,2,2)\nplt.imshow(processed)\nplt.title(\"Preprocessed\")\nplt.axis(\"off\")\n\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:41.974459Z","iopub.execute_input":"2026-05-01T11:21:41.974823Z","iopub.status.idle":"2026-05-01T11:21:42.978037Z","shell.execute_reply.started":"2026-05-01T11:21:41.974791Z","shell.execute_reply":"2026-05-01T11:21:42.977273Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"PROBLEM IN PREPROCESSING\n\nYour processed image is:\n\n❌ Too gray / too flat\n❌ Contrast is distorted\n❌ Retina color information is lost\n❌ Lesion visibility is reduced\n\nThis usually happens when:\n\n1. Over-processing (Ben Graham too strong)\nImage becomes “washed out”\n2. Incorrect normalization display\nYou might be visualizing normalized image directly\n3. Missing color space correction (LAB/CLAHE issue)","metadata":{}},{"cell_type":"markdown","source":"Reduce Ben Graham intensity\n\nYour formula is likely too strong.\n\n🔧 Fix 2: Prefer CLAHE instead","metadata":{}},{"cell_type":"code","source":"import cv2\nimport matplotlib.pyplot as plt\n\n# take one sample from dataset\nsample = aptos_labels.iloc[0]\n\nimg_path = f\"{APTOS_IMAGES}/{sample.id_code}.png\"\n\nprint(\"Image path:\", img_path)\n\nimg = cv2.imread(img_path)\n\nimport cv2\nimport matplotlib.pyplot as plt\n\nsample = aptos_labels.iloc[0]\nimg_path = f\"{APTOS_IMAGES}/{sample.id_code}.png\"\n\nimg = cv2.imread(img_path)\n\n# CLAHE on LAB\nimg_lab = cv2.cvtColor(img, cv2.COLOR_BGR2LAB)\n\nl, a, b = cv2.split(img_lab)\n\nclahe = cv2.createCLAHE(\n    clipLimit=2.0,\n    tileGridSize=(8,8)\n)\n\nl = clahe.apply(l)\n\nimg_lab = cv2.merge((l, a, b))\n\nimg_clahe = cv2.cvtColor(\n    img_lab,\n    cv2.COLOR_LAB2BGR\n)\n\n# Convert for display\nimg_rgb = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n\nimg_clahe_rgb = cv2.cvtColor(\n    img_clahe,\n    cv2.COLOR_BGR2RGB\n)\n\nplt.figure(figsize=(10,5))\n\nplt.subplot(1,2,1)\nplt.imshow(img_rgb)\nplt.title(\"Original\")\n\nplt.subplot(1,2,2)\nplt.imshow(img_clahe_rgb)\nplt.title(\"CLAHE Preprocessed\")\n\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:42.979131Z","iopub.execute_input":"2026-05-01T11:21:42.979382Z","iopub.status.idle":"2026-05-01T11:21:44.333804Z","shell.execute_reply.started":"2026-05-01T11:21:42.979358Z","shell.execute_reply":"2026-05-01T11:21:44.332909Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"CLAHE Result — Professional Evaluation\nLeft: Original\nRight: CLAHE Preprocessed\nWhat improved (GOOD):\n\n✔ Blood vessels are clearer\n✔ Lesions (yellow spots) are more visible\n✔ Contrast is better\n✔ Colors are preserved\n✔ No gray distortion\n✔ Retina structure intact\n\nThis is exactly what we want for Diabetic Retinopathy detection.\n\n🧠 Scientific Interpretation (for your thesis)\n\nYou can confidently say:\n\nCLAHE improves local contrast and enhances visibility of retinal lesions while preserving color information, making it suitable for multi-dataset training and generalization.\n\nThis statement is methodology-ready.","metadata":{}},{"cell_type":"markdown","source":"FREEZE DECISION (VERY IMPORTANT)\n\nWe now officially freeze preprocessing.\n\nFinal Preprocessing Pipeline:\n\n1. Circular Crop\n2. CLAHE (clipLimit=2.0, tileGridSize=(8,8))\n3. Resize to 456×456\n4. Normalize (ImageNet)","metadata":{}},{"cell_type":"markdown","source":"Next Step — Build Final Reusable Preprocessing Function\nThis function will be used everywhere:\n\ntraining\nvalidation\ncross-dataset testing\nGradCAM\ndeployment","metadata":{}},{"cell_type":"markdown","source":"STEP 3 — Final Production Preprocessing Function","metadata":{}},{"cell_type":"code","source":"import cv2\nimport numpy as np\n\nIMG_SIZE = 456\n\ndef circular_crop(img):\n    gray = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)\n\n    _, thresh = cv2.threshold(gray, 10, 255, cv2.THRESH_BINARY)\n\n    contours, _ = cv2.findContours(\n        thresh,\n        cv2.RETR_EXTERNAL,\n        cv2.CHAIN_APPROX_SIMPLE\n    )\n\n    if len(contours) == 0:\n        return img\n\n    cnt = max(contours, key=cv2.contourArea)\n\n    x, y, w, h = cv2.boundingRect(cnt)\n\n    cropped = img[y:y+h, x:x+w]\n\n    return cropped\n\n\ndef apply_clahe(img):\n\n    lab = cv2.cvtColor(img, cv2.COLOR_BGR2LAB)\n\n    l, a, b = cv2.split(lab)\n\n    clahe = cv2.createCLAHE(\n        clipLimit=2.0,\n        tileGridSize=(8,8)\n    )\n\n    l = clahe.apply(l)\n\n    lab = cv2.merge((l, a, b))\n\n    img = cv2.cvtColor(lab, cv2.COLOR_LAB2BGR)\n\n    return img\n\n\ndef preprocess_image(path):\n\n    img = cv2.imread(path)\n\n    img = circular_crop(img)\n\n    img = apply_clahe(img)\n\n    img = cv2.resize(\n        img,\n        (IMG_SIZE, IMG_SIZE)\n    )\n\n    img = cv2.cvtColor(\n        img,\n        cv2.COLOR_BGR2RGB\n    )\n\n    img = img / 255.0\n\n    return img","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:44.334899Z","iopub.execute_input":"2026-05-01T11:21:44.335208Z","iopub.status.idle":"2026-05-01T11:21:44.648212Z","shell.execute_reply.started":"2026-05-01T11:21:44.335181Z","shell.execute_reply":"2026-05-01T11:21:44.647659Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Test the Final Pipeline","metadata":{}},{"cell_type":"code","source":"sample = aptos_labels.iloc[0]\n\npath = f\"{APTOS_IMAGES}/{sample.id_code}.png\"\n\nprocessed = preprocess_image(path)\n\nplt.imshow(processed)\nplt.title(\"Final Preprocessed Image\")\nplt.axis(\"off\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:44.649033Z","iopub.execute_input":"2026-05-01T11:21:44.649247Z","iopub.status.idle":"2026-05-01T11:21:45.053230Z","shell.execute_reply.started":"2026-05-01T11:21:44.649226Z","shell.execute_reply":"2026-05-01T11:21:45.052478Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"final preprocessed image looks:\n\nclear\nnatural\ncolor preserved\nvessels visible\nlesions visible\nproperly resized to 456 × 456","metadata":{}},{"cell_type":"markdown","source":"Preprocessing Phase Status\n\nWe now freeze this preprocessing pipeline:\n\nCircular Crop\nCLAHE\nResize to 456×456\nScale to 0–1\n\nLater during training, we will add ImageNet normalization inside dataset transforms.","metadata":{}},{"cell_type":"markdown","source":"**NEXT PHASE: DATASET FUSION**","metadata":{}},{"cell_type":"markdown","source":"Now we will:\n\ncreate APTOS dataframe\nprepare IDRiD dataframe\nmerge both\nadd source column\ncheck class distribution","metadata":{}},{"cell_type":"markdown","source":"STEP 4.1 — APTOS DataFrame","metadata":{}},{"cell_type":"code","source":"aptos_df = aptos_labels.copy()\n\naptos_df[\"image_path\"] = aptos_df[\"id_code\"].apply(\n    lambda x: f\"{APTOS_IMAGES}/{x}.png\"\n)\n\naptos_df[\"label\"] = aptos_df[\"diagnosis\"]\naptos_df[\"source\"] = \"aptos\"\n\naptos_df = aptos_df[[\"image_path\", \"label\", \"source\"]]\n\naptos_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:45.054205Z","iopub.execute_input":"2026-05-01T11:21:45.054444Z","iopub.status.idle":"2026-05-01T11:21:45.067454Z","shell.execute_reply.started":"2026-05-01T11:21:45.054421Z","shell.execute_reply":"2026-05-01T11:21:45.066791Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"STEP 4.2 — Check APTOS DataFrame","metadata":{}},{"cell_type":"code","source":"print(\"APTOS shape:\", aptos_df.shape)\nprint(aptos_df[\"label\"].value_counts().sort_index())\nprint(aptos_df[\"source\"].value_counts())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:45.068364Z","iopub.execute_input":"2026-05-01T11:21:45.068649Z","iopub.status.idle":"2026-05-01T11:21:45.083485Z","shell.execute_reply.started":"2026-05-01T11:21:45.068618Z","shell.execute_reply":"2026-05-01T11:21:45.082791Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"STEP 4.6 — BUILD IDRiD DATAFRAMES","metadata":{}},{"cell_type":"code","source":"idrid_train_df = train_gt.copy()\nidrid_test_df = test_gt.copy()\n\nidrid_train_df[\"image_path\"] = idrid_train_df[\"Image name\"].apply(\n    lambda x: f\"{TRAIN_IMG_PATH}/{x}.jpg\"\n)\n\nidrid_test_df[\"image_path\"] = idrid_test_df[\"Image name\"].apply(\n    lambda x: f\"{TEST_IMG_PATH}/{x}.jpg\"\n)\n\nidrid_train_df[\"label\"] = idrid_train_df[\"Retinopathy grade\"]\nidrid_test_df[\"label\"] = idrid_test_df[\"Retinopathy grade\"]\n\nidrid_train_df[\"source\"] = \"idrid\"\nidrid_test_df[\"source\"] = \"idrid\"\n\nidrid_df = pd.concat([idrid_train_df, idrid_test_df], ignore_index=True)\n\nidrid_df = idrid_df[[\"image_path\", \"label\", \"source\"]]\n\nidrid_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:45.084467Z","iopub.execute_input":"2026-05-01T11:21:45.084785Z","iopub.status.idle":"2026-05-01T11:21:45.102401Z","shell.execute_reply.started":"2026-05-01T11:21:45.084750Z","shell.execute_reply":"2026-05-01T11:21:45.101750Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"STEP 4.7 — VERIFY IDRiD DATAFRAME","metadata":{}},{"cell_type":"code","source":"print(\"IDRiD shape:\", idrid_df.shape)\nprint(idrid_df[\"label\"].value_counts().sort_index())\nprint(idrid_df[\"source\"].value_counts())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:45.103282Z","iopub.execute_input":"2026-05-01T11:21:45.103593Z","iopub.status.idle":"2026-05-01T11:21:45.130423Z","shell.execute_reply.started":"2026-05-01T11:21:45.103556Z","shell.execute_reply":"2026-05-01T11:21:45.129605Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"IDRiD dataframe is ready, and its class distribution matches the expected one:\n\n0 → 168\n1 → 25\n2 → 168\n3 → 93\n4 → 62\n\nSo now we can do the main fusion step.","metadata":{}},{"cell_type":"markdown","source":"**STEP 4.8 — MERGE APTOS + IDRiD)**","metadata":{}},{"cell_type":"code","source":"combined_df = pd.concat([aptos_df, idrid_df], ignore_index=True)\n\nprint(\"Combined shape:\", combined_df.shape)\nprint(\"\\nCombined class distribution:\")\nprint(combined_df[\"label\"].value_counts().sort_index())\n\nprint(\"\\nSource distribution:\")\nprint(combined_df[\"source\"].value_counts())\n\ncombined_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:45.131480Z","iopub.execute_input":"2026-05-01T11:21:45.131771Z","iopub.status.idle":"2026-05-01T11:21:45.153786Z","shell.execute_reply.started":"2026-05-01T11:21:45.131729Z","shell.execute_reply":"2026-05-01T11:21:45.153117Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**STEP 4.9 — CHECK IMAGE PATH VALIDITY**","metadata":{}},{"cell_type":"code","source":"import os\n\nprint(\"Missing image paths:\",\n      combined_df[\"image_path\"].apply(lambda x: not os.path.exists(x)).sum())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:45.154680Z","iopub.execute_input":"2026-05-01T11:21:45.155018Z","iopub.status.idle":"2026-05-01T11:21:45.837210Z","shell.execute_reply.started":"2026-05-01T11:21:45.154986Z","shell.execute_reply":"2026-05-01T11:21:45.836458Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"fusion step is now fully successful.","metadata":{}},{"cell_type":"markdown","source":"Combined dataset\n4178 images\n5 DR classes\n2 sources: APTOS + IDRiD\n0 missing image paths\n\nThis means our core multi-dataset fusion foundation is ready.","metadata":{}},{"cell_type":"markdown","source":"successfully built a merged 5-class diabetic retinopathy dataset from APTOS and IDRiD, and all image paths are valid. This fused dataset will be used to study generalization across domains.","metadata":{}},{"cell_type":"markdown","source":"**NEXT STEP — CREATE 5 FOLDS**","metadata":{}},{"cell_type":"code","source":"from sklearn.model_selection import StratifiedKFold\n\ncombined_df = combined_df.copy()\ncombined_df[\"fold\"] = -1\n\nskf = StratifiedKFold(\n    n_splits=5,\n    shuffle=True,\n    random_state=42\n)\n\nfor fold, (_, val_idx) in enumerate(skf.split(combined_df, combined_df[\"label\"])):\n    combined_df.loc[val_idx, \"fold\"] = fold\n\nprint(combined_df[\"fold\"].value_counts().sort_index())\ncombined_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:45.838259Z","iopub.execute_input":"2026-05-01T11:21:45.838792Z","iopub.status.idle":"2026-05-01T11:21:45.857285Z","shell.execute_reply.started":"2026-05-01T11:21:45.838753Z","shell.execute_reply":"2026-05-01T11:21:45.856493Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"fold_dist = pd.crosstab(combined_df[\"fold\"], combined_df[\"label\"])\nprint(fold_dist)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:45.861761Z","iopub.execute_input":"2026-05-01T11:21:45.861965Z","iopub.status.idle":"2026-05-01T11:21:45.872992Z","shell.execute_reply.started":"2026-05-01T11:21:45.861945Z","shell.execute_reply":"2026-05-01T11:21:45.872420Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"We have successfully finished:\n\nPhase 1 — Data loading\nPhase 2 — EDA\nPhase 3 — Preprocessing\nPhase 4 — Dataset fusion + 5-fold split\n\nNow we move to the next major phase:","metadata":{}},{"cell_type":"markdown","source":"**PHASE 5 — BASELINE MODEL TRAINING**","metadata":{}},{"cell_type":"markdown","source":"We will first build a plain EfficientNetB5 baseline\nwithout CBAM.\n\nThis is important because later you must prove:\n\nCBAM improves over baseline\nfusion helps over single dataset\nfocal loss helps over normal loss\n\nSo baseline comes first.","metadata":{}},{"cell_type":"markdown","source":"STEP 5.1 — INSTALL / IMPORT LIBRARIES","metadata":{}},{"cell_type":"code","source":"import os\nimport cv2\nimport math\nimport time\nimport copy\nimport random\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\nfrom PIL import Image\nfrom sklearn.metrics import accuracy_score, f1_score, confusion_matrix, classification_report, cohen_kappa_score\nfrom sklearn.model_selection import StratifiedKFold\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\n\nimport torchvision\nfrom torchvision import transforms\nfrom torchvision.models import efficientnet_b5, EfficientNet_B5_Weights","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:45.873901Z","iopub.execute_input":"2026-05-01T11:21:45.874400Z","iopub.status.idle":"2026-05-01T11:21:45.880452Z","shell.execute_reply.started":"2026-05-01T11:21:45.874375Z","shell.execute_reply":"2026-05-01T11:21:45.879931Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"STEP 5.2 — DEVICE CHECK","metadata":{}},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Device:\", device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:45.881326Z","iopub.execute_input":"2026-05-01T11:21:45.881577Z","iopub.status.idle":"2026-05-01T11:21:45.899199Z","shell.execute_reply.started":"2026-05-01T11:21:45.881545Z","shell.execute_reply":"2026-05-01T11:21:45.898553Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"STEP 5.3 — SET RANDOM SEED","metadata":{}},{"cell_type":"code","source":"def set_seed(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n\nset_seed(42)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:45.900107Z","iopub.execute_input":"2026-05-01T11:21:45.900346Z","iopub.status.idle":"2026-05-01T11:21:45.913874Z","shell.execute_reply.started":"2026-05-01T11:21:45.900325Z","shell.execute_reply":"2026-05-01T11:21:45.913236Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"STEP 5.4 — TRAIN / VALID SPLIT USING ONE FOLD","metadata":{}},{"cell_type":"code","source":"train_df = combined_df[combined_df[\"fold\"] != 0].reset_index(drop=True)\nval_df   = combined_df[combined_df[\"fold\"] == 0].reset_index(drop=True)\n\nprint(\"Train shape:\", train_df.shape)\nprint(\"Val shape:\", val_df.shape)\n\nprint(\"\\nTrain class distribution:\")\nprint(train_df[\"label\"].value_counts().sort_index())\n\nprint(\"\\nVal class distribution:\")\nprint(val_df[\"label\"].value_counts().sort_index())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:45.914778Z","iopub.execute_input":"2026-05-01T11:21:45.915028Z","iopub.status.idle":"2026-05-01T11:21:45.934462Z","shell.execute_reply.started":"2026-05-01T11:21:45.914993Z","shell.execute_reply":"2026-05-01T11:21:45.933517Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"STEP 5.5 — FINAL PREPROCESS FUNCTION FOR OPENCV IMAGE LOADING","metadata":{}},{"cell_type":"code","source":"IMG_SIZE = 456\n\ndef circular_crop(img):\n    gray = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)\n    _, thresh = cv2.threshold(gray, 10, 255, cv2.THRESH_BINARY)\n    contours, _ = cv2.findContours(thresh, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n\n    if len(contours) == 0:\n        return img\n\n    cnt = max(contours, key=cv2.contourArea)\n    x, y, w, h = cv2.boundingRect(cnt)\n    cropped = img[y:y+h, x:x+w]\n    return cropped\n\n\ndef apply_clahe(img):\n    lab = cv2.cvtColor(img, cv2.COLOR_BGR2LAB)\n    l, a, b = cv2.split(lab)\n\n    clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8,8))\n    l = clahe.apply(l)\n\n    lab = cv2.merge((l, a, b))\n    img = cv2.cvtColor(lab, cv2.COLOR_LAB2BGR)\n    return img","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:45.935837Z","iopub.execute_input":"2026-05-01T11:21:45.936726Z","iopub.status.idle":"2026-05-01T11:21:45.946372Z","shell.execute_reply.started":"2026-05-01T11:21:45.936682Z","shell.execute_reply":"2026-05-01T11:21:45.945661Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"STEP 5.6 — DEFINE TRAIN / VALID TRANSFORMS\nNow we add augmentation for training and ImageNet normalization.","metadata":{}},{"cell_type":"code","source":"train_transform = transforms.Compose([\n    transforms.ToPILImage(),\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.1, hue=0.05),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225]\n    )\n])\n\nval_transform = transforms.Compose([\n    transforms.ToPILImage(),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225]\n    )\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:45.947880Z","iopub.execute_input":"2026-05-01T11:21:45.948879Z","iopub.status.idle":"2026-05-01T11:21:45.966796Z","shell.execute_reply.started":"2026-05-01T11:21:45.948841Z","shell.execute_reply":"2026-05-01T11:21:45.965536Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"STEP 5.7 — CREATE PYTORCH DATASET CLASS","metadata":{}},{"cell_type":"code","source":"class DRDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df.reset_index(drop=True)\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n\n        image_path = row[\"image_path\"]\n        label = int(row[\"label\"])\n\n        img = cv2.imread(image_path)\n        img = circular_crop(img)\n        img = apply_clahe(img)\n        img = cv2.resize(img, (IMG_SIZE, IMG_SIZE))\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n\n        if self.transform:\n            img = self.transform(img)\n\n        return img, label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:45.968050Z","iopub.execute_input":"2026-05-01T11:21:45.968340Z","iopub.status.idle":"2026-05-01T11:21:45.987422Z","shell.execute_reply.started":"2026-05-01T11:21:45.968309Z","shell.execute_reply":"2026-05-01T11:21:45.986907Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"STEP 5.8 — CREATE DATASETS AND DATALOADERS","metadata":{}},{"cell_type":"code","source":"train_dataset = DRDataset(train_df, transform=train_transform)\nval_dataset   = DRDataset(val_df, transform=val_transform)\n\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=8,\n    shuffle=True,\n    num_workers=2,\n    pin_memory=True\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=8,\n    shuffle=False,\n    num_workers=2,\n    pin_memory=True\n)\n\nprint(\"Train batches:\", len(train_loader))\nprint(\"Val batches:\", len(val_loader))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:45.988285Z","iopub.execute_input":"2026-05-01T11:21:45.988489Z","iopub.status.idle":"2026-05-01T11:21:46.002118Z","shell.execute_reply.started":"2026-05-01T11:21:45.988469Z","shell.execute_reply":"2026-05-01T11:21:46.001467Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"STEP 5.9 — TEST ONE BATCH","metadata":{}},{"cell_type":"code","source":"images, labels = next(iter(train_loader))\n\nprint(\"Images shape:\", images.shape)\nprint(\"Labels shape:\", labels.shape)\nprint(\"Sample labels:\", labels[:8])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:46.002828Z","iopub.execute_input":"2026-05-01T11:21:46.003055Z","iopub.status.idle":"2026-05-01T11:21:49.850022Z","shell.execute_reply.started":"2026-05-01T11:21:46.003027Z","shell.execute_reply":"2026-05-01T11:21:49.849139Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"STEP 5.10 — BUILD BASELINE EFFICIENTNET-B5 MODEL","metadata":{}},{"cell_type":"code","source":"weights = EfficientNet_B5_Weights.IMAGENET1K_V1\nmodel = efficientnet_b5(weights=weights)\n\nin_features = model.classifier[1].in_features\nmodel.classifier[1] = nn.Linear(in_features, 5)\n\nmodel = model.to(device)\n\nprint(model.classifier)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:49.851413Z","iopub.execute_input":"2026-05-01T11:21:49.851822Z","iopub.status.idle":"2026-05-01T11:21:50.451085Z","shell.execute_reply.started":"2026-05-01T11:21:49.851791Z","shell.execute_reply":"2026-05-01T11:21:50.450317Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"STEP 5.11 — LOSS, OPTIMIZER\n\nFor baseline, start simple:\n\nCrossEntropyLoss\nAdamW","metadata":{}},{"cell_type":"code","source":"criterion = nn.CrossEntropyLoss()\n\noptimizer = optim.AdamW(\n    model.parameters(),\n    lr=1e-4,\n    weight_decay=1e-4\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:50.452065Z","iopub.execute_input":"2026-05-01T11:21:50.452335Z","iopub.status.idle":"2026-05-01T11:21:50.458248Z","shell.execute_reply.started":"2026-05-01T11:21:50.452310Z","shell.execute_reply":"2026-05-01T11:21:50.457708Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"STEP 5.12 — SIMPLE TRAIN / VALID FUNCTIONS","metadata":{}},{"cell_type":"code","source":"def train_one_epoch(model, loader, criterion, optimizer, device):\n    model.train()\n\n    running_loss = 0.0\n    all_preds = []\n    all_labels = []\n\n    for images, labels in loader:\n        images = images.to(device)\n        labels = labels.to(device)\n\n        optimizer.zero_grad()\n\n        outputs = model(images)\n        loss = criterion(outputs, labels)\n\n        loss.backward()\n        optimizer.step()\n\n        running_loss += loss.item() * images.size(0)\n\n        preds = torch.argmax(outputs, dim=1)\n        all_preds.extend(preds.detach().cpu().numpy())\n        all_labels.extend(labels.detach().cpu().numpy())\n\n    epoch_loss = running_loss / len(loader.dataset)\n    epoch_acc = accuracy_score(all_labels, all_preds)\n    epoch_f1 = f1_score(all_labels, all_preds, average=\"weighted\")\n    epoch_qwk = cohen_kappa_score(all_labels, all_preds, weights=\"quadratic\")\n\n    return epoch_loss, epoch_acc, epoch_f1, epoch_qwk\n\n\ndef validate_one_epoch(model, loader, criterion, device):\n    model.eval()\n\n    running_loss = 0.0\n    all_preds = []\n    all_labels = []\n\n    with torch.no_grad():\n        for images, labels in loader:\n            images = images.to(device)\n            labels = labels.to(device)\n\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n\n            running_loss += loss.item() * images.size(0)\n\n            preds = torch.argmax(outputs, dim=1)\n            all_preds.extend(preds.detach().cpu().numpy())\n            all_labels.extend(labels.detach().cpu().numpy())\n\n    epoch_loss = running_loss / len(loader.dataset)\n    epoch_acc = accuracy_score(all_labels, all_preds)\n    epoch_f1 = f1_score(all_labels, all_preds, average=\"weighted\")\n    epoch_qwk = cohen_kappa_score(all_labels, all_preds, weights=\"quadratic\")\n\n    return epoch_loss, epoch_acc, epoch_f1, epoch_qwk, all_labels, all_preds","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:50.459175Z","iopub.execute_input":"2026-05-01T11:21:50.459375Z","iopub.status.idle":"2026-05-01T11:21:50.473932Z","shell.execute_reply.started":"2026-05-01T11:21:50.459355Z","shell.execute_reply":"2026-05-01T11:21:50.473356Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"STEP 5.13 — RUN A SHORT BASELINE TRAINING","metadata":{}},{"cell_type":"code","source":"import torch\n\nprint(\"CUDA available:\", torch.cuda.is_available())\nprint(\"GPU count:\", torch.cuda.device_count())\n\nif torch.cuda.is_available():\n    print(\"GPU name:\", torch.cuda.get_device_name(0))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:50.474861Z","iopub.execute_input":"2026-05-01T11:21:50.475143Z","iopub.status.idle":"2026-05-01T11:21:50.493031Z","shell.execute_reply.started":"2026-05-01T11:21:50.475111Z","shell.execute_reply":"2026-05-01T11:21:50.492395Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"GPU is working correctly now\n\nRebuild loaders with batch size 16","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.optim as optim\n\nfrom torch.utils.data import Dataset, DataLoader\n\nfrom torchvision import transforms\nfrom torchvision.models import efficientnet_b5, EfficientNet_B5_Weights\n\nimport numpy as np\nimport pandas as pd\nimport cv2\nimport random","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:50.494005Z","iopub.execute_input":"2026-05-01T11:21:50.494386Z","iopub.status.idle":"2026-05-01T11:21:50.511337Z","shell.execute_reply.started":"2026-05-01T11:21:50.494363Z","shell.execute_reply":"2026-05-01T11:21:50.510781Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Using device:\", device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:50.512039Z","iopub.execute_input":"2026-05-01T11:21:50.512297Z","iopub.status.idle":"2026-05-01T11:21:50.526889Z","shell.execute_reply.started":"2026-05-01T11:21:50.512275Z","shell.execute_reply":"2026-05-01T11:21:50.526300Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Reload Saved DataFrames (Rebuild Combined Dataset)","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport os\nfrom sklearn.model_selection import StratifiedKFold\n\n# ----------------------------\n# PATHS\n# ----------------------------\n\nAPTOS_PATH = \"/kaggle/input/competitions/aptos2019-blindness-detection\"\nAPTOS_IMAGES = f\"{APTOS_PATH}/train_images\"\n\nIDRID_BASE = \"/kaggle/input/datasets/abdullahshafi315/indian-diabetic-retinopathy-image-datasetidrid/Disease Grading\"\n\nTRAIN_IMG_PATH = f\"{IDRID_BASE}/1. Original Images/a. Training Set\"\nTEST_IMG_PATH  = f\"{IDRID_BASE}/1. Original Images/b. Testing Set\"\n\nGT_PATH = f\"{IDRID_BASE}/2. Groundtruths\"\n\n# ----------------------------\n# LOAD APTOS LABELS\n# ----------------------------\n\naptos_labels = pd.read_csv(f\"{APTOS_PATH}/train.csv\")\n\naptos_df = aptos_labels.copy()\n\naptos_df[\"image_path\"] = aptos_df[\"id_code\"].apply(\n    lambda x: f\"{APTOS_IMAGES}/{x}.png\"\n)\n\naptos_df[\"label\"] = aptos_df[\"diagnosis\"]\naptos_df[\"source\"] = \"aptos\"\n\naptos_df = aptos_df[[\"image_path\", \"label\", \"source\"]]\n\nprint(\"APTOS shape:\", aptos_df.shape)\n\n# ----------------------------\n# LOAD IDRID LABELS\n# ----------------------------\n\ntrain_gt = pd.read_csv(\n    f\"{GT_PATH}/a. IDRiD_Disease Grading_Training Labels.csv\"\n)\n\ntest_gt = pd.read_csv(\n    f\"{GT_PATH}/b. IDRiD_Disease Grading_Testing Labels.csv\"\n)\n\nidrid_train_df = train_gt.copy()\nidrid_test_df = test_gt.copy()\n\nidrid_train_df[\"image_path\"] = idrid_train_df[\"Image name\"].apply(\n    lambda x: f\"{TRAIN_IMG_PATH}/{x}.jpg\"\n)\n\nidrid_test_df[\"image_path\"] = idrid_test_df[\"Image name\"].apply(\n    lambda x: f\"{TEST_IMG_PATH}/{x}.jpg\"\n)\n\nidrid_train_df[\"label\"] = idrid_train_df[\"Retinopathy grade\"]\nidrid_test_df[\"label\"] = idrid_test_df[\"Retinopathy grade\"]\n\nidrid_train_df[\"source\"] = \"idrid\"\nidrid_test_df[\"source\"] = \"idrid\"\n\nidrid_df = pd.concat(\n    [idrid_train_df, idrid_test_df],\n    ignore_index=True\n)\n\nidrid_df = idrid_df[[\"image_path\", \"label\", \"source\"]]\n\nprint(\"IDRiD shape:\", idrid_df.shape)\n\n# ----------------------------\n# MERGE DATASETS\n# ----------------------------\n\ncombined_df = pd.concat(\n    [aptos_df, idrid_df],\n    ignore_index=True\n)\n\nprint(\"Combined shape:\", combined_df.shape)\n\n# ----------------------------\n# CREATE FOLDS\n# ----------------------------\n\ncombined_df[\"fold\"] = -1\n\nskf = StratifiedKFold(\n    n_splits=5,\n    shuffle=True,\n    random_state=42\n)\n\nfor fold, (_, val_idx) in enumerate(\n    skf.split(combined_df, combined_df[\"label\"])\n):\n    combined_df.loc[val_idx, \"fold\"] = fold\n\nprint(\"\\nFold distribution:\")\nprint(combined_df[\"fold\"].value_counts())\n\n# ----------------------------\n# TRAIN / VALID SPLIT\n# ----------------------------\n\ntrain_df = combined_df[\n    combined_df[\"fold\"] != 0\n].reset_index(drop=True)\n\nval_df = combined_df[\n    combined_df[\"fold\"] == 0\n].reset_index(drop=True)\n\nprint(\"\\nTrain shape:\", train_df.shape)\nprint(\"Val shape:\", val_df.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:50.527838Z","iopub.execute_input":"2026-05-01T11:21:50.528461Z","iopub.status.idle":"2026-05-01T11:21:50.569228Z","shell.execute_reply.started":"2026-05-01T11:21:50.528438Z","shell.execute_reply":"2026-05-01T11:21:50.568653Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"we continue from here and rebuild the training pipeline on GPU.","metadata":{}},{"cell_type":"markdown","source":"STEP 5.6 — Imports for Training","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.optim as optim\n\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nfrom torchvision.models import efficientnet_b5, EfficientNet_B5_Weights\n\nimport cv2\nimport numpy as np\nimport random\nfrom sklearn.metrics import accuracy_score, f1_score, cohen_kappa_score","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:50.570015Z","iopub.execute_input":"2026-05-01T11:21:50.570305Z","iopub.status.idle":"2026-05-01T11:21:50.574362Z","shell.execute_reply.started":"2026-05-01T11:21:50.570273Z","shell.execute_reply":"2026-05-01T11:21:50.573675Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"STEP 5.7 — Device and Seed","metadata":{}},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Using device:\", device)\n\ndef set_seed(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n\nset_seed(42)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:50.575279Z","iopub.execute_input":"2026-05-01T11:21:50.575559Z","iopub.status.idle":"2026-05-01T11:21:50.592226Z","shell.execute_reply.started":"2026-05-01T11:21:50.575528Z","shell.execute_reply":"2026-05-01T11:21:50.591635Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"STEP 5.8 — Preprocessing Functions","metadata":{}},{"cell_type":"code","source":"IMG_SIZE = 456\n\ndef circular_crop(img):\n    gray = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)\n    _, thresh = cv2.threshold(gray, 10, 255, cv2.THRESH_BINARY)\n    contours, _ = cv2.findContours(thresh, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n\n    if len(contours) == 0:\n        return img\n\n    cnt = max(contours, key=cv2.contourArea)\n    x, y, w, h = cv2.boundingRect(cnt)\n    cropped = img[y:y+h, x:x+w]\n    return cropped\n\n\ndef apply_clahe(img):\n    lab = cv2.cvtColor(img, cv2.COLOR_BGR2LAB)\n    l, a, b = cv2.split(lab)\n\n    clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8,8))\n    l = clahe.apply(l)\n\n    lab = cv2.merge((l, a, b))\n    img = cv2.cvtColor(lab, cv2.COLOR_LAB2BGR)\n    return img","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:50.593082Z","iopub.execute_input":"2026-05-01T11:21:50.593568Z","iopub.status.idle":"2026-05-01T11:21:50.603953Z","shell.execute_reply.started":"2026-05-01T11:21:50.593535Z","shell.execute_reply":"2026-05-01T11:21:50.603424Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"STEP 5.9 — Transforms","metadata":{}},{"cell_type":"code","source":"train_transform = transforms.Compose([\n    transforms.ToPILImage(),\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.1, hue=0.05),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225]\n    )\n])\n\nval_transform = transforms.Compose([\n    transforms.ToPILImage(),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225]\n    )\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:50.604762Z","iopub.execute_input":"2026-05-01T11:21:50.605574Z","iopub.status.idle":"2026-05-01T11:21:50.622120Z","shell.execute_reply.started":"2026-05-01T11:21:50.605549Z","shell.execute_reply":"2026-05-01T11:21:50.621576Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"STEP 5.10 — Dataset Class","metadata":{}},{"cell_type":"code","source":"class DRDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df.reset_index(drop=True)\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        image_path = row[\"image_path\"]\n        label = int(row[\"label\"])\n\n        img = cv2.imread(image_path)\n        img = circular_crop(img)\n        img = apply_clahe(img)\n        img = cv2.resize(img, (IMG_SIZE, IMG_SIZE))\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n\n        if self.transform:\n            img = self.transform(img)\n\n        return img, label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:50.622867Z","iopub.execute_input":"2026-05-01T11:21:50.623076Z","iopub.status.idle":"2026-05-01T11:21:50.640422Z","shell.execute_reply.started":"2026-05-01T11:21:50.623057Z","shell.execute_reply":"2026-05-01T11:21:50.639892Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"STEP 5.11 — Create Datasets and Loaders","metadata":{}},{"cell_type":"code","source":"train_dataset = DRDataset(train_df, transform=train_transform)\nval_dataset   = DRDataset(val_df, transform=val_transform)\n\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=8,\n    shuffle=True,\n    num_workers=2,\n    pin_memory=True\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=8,\n    shuffle=False,\n    num_workers=2,\n    pin_memory=True\n)\n\nprint(\"Train batches:\", len(train_loader))\nprint(\"Val batches:\", len(val_loader))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:50.641195Z","iopub.execute_input":"2026-05-01T11:21:50.641376Z","iopub.status.idle":"2026-05-01T11:21:50.655332Z","shell.execute_reply.started":"2026-05-01T11:21:50.641357Z","shell.execute_reply":"2026-05-01T11:21:50.654549Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"STEP 5.12 — Check One Batch","metadata":{}},{"cell_type":"code","source":"images, labels = next(iter(train_loader))\n\nprint(\"Images shape:\", images.shape)\nprint(\"Labels shape:\", labels.shape)\nprint(\"Sample labels:\", labels[:8])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:50.656351Z","iopub.execute_input":"2026-05-01T11:21:50.656609Z","iopub.status.idle":"2026-05-01T11:21:54.552565Z","shell.execute_reply.started":"2026-05-01T11:21:50.656563Z","shell.execute_reply":"2026-05-01T11:21:54.551869Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"STEP 5.13 — Build Baseline EfficientNetB5","metadata":{}},{"cell_type":"code","source":"weights = EfficientNet_B5_Weights.IMAGENET1K_V1\nmodel = efficientnet_b5(weights=weights)\n\nin_features = model.classifier[1].in_features\nmodel.classifier[1] = nn.Linear(in_features, 5)\n\nmodel = model.to(device)\n\nprint(model.classifier)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:54.553822Z","iopub.execute_input":"2026-05-01T11:21:54.554081Z","iopub.status.idle":"2026-05-01T11:21:55.147739Z","shell.execute_reply.started":"2026-05-01T11:21:54.554052Z","shell.execute_reply":"2026-05-01T11:21:55.147016Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"STEP 5.14 — Loss and Optimizer","metadata":{}},{"cell_type":"code","source":"criterion = nn.CrossEntropyLoss()\n\noptimizer = optim.AdamW(\n    model.parameters(),\n    lr=1e-4,\n    weight_decay=1e-4\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:55.148654Z","iopub.execute_input":"2026-05-01T11:21:55.148918Z","iopub.status.idle":"2026-05-01T11:21:55.157805Z","shell.execute_reply.started":"2026-05-01T11:21:55.148895Z","shell.execute_reply":"2026-05-01T11:21:55.157070Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"STEP 5.15 — Train / Validate Functions","metadata":{}},{"cell_type":"code","source":"def train_one_epoch(model, loader, criterion, optimizer, device):\n    model.train()\n\n    running_loss = 0.0\n    all_preds = []\n    all_labels = []\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\n        outputs = model(images)\n        loss = criterion(outputs, labels)\n\n        loss.backward()\n        optimizer.step()\n\n        running_loss += loss.item() * images.size(0)\n\n        preds = torch.argmax(outputs, dim=1)\n        all_preds.extend(preds.detach().cpu().numpy())\n        all_labels.extend(labels.detach().cpu().numpy())\n\n    epoch_loss = running_loss / len(loader.dataset)\n    epoch_acc = accuracy_score(all_labels, all_preds)\n    epoch_f1 = f1_score(all_labels, all_preds, average=\"weighted\")\n    epoch_qwk = cohen_kappa_score(all_labels, all_preds, weights=\"quadratic\")\n\n    return epoch_loss, epoch_acc, epoch_f1, epoch_qwk\n\n\ndef validate_one_epoch(model, loader, criterion, device):\n    model.eval()\n\n    running_loss = 0.0\n    all_preds = []\n    all_labels = []\n\n    with torch.no_grad():\n        for images, labels in loader:\n            images = images.to(device, non_blocking=True)\n            labels = labels.to(device, non_blocking=True)\n\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n\n            running_loss += loss.item() * images.size(0)\n\n            preds = torch.argmax(outputs, dim=1)\n            all_preds.extend(preds.detach().cpu().numpy())\n            all_labels.extend(labels.detach().cpu().numpy())\n\n    epoch_loss = running_loss / len(loader.dataset)\n    epoch_acc = accuracy_score(all_labels, all_preds)\n    epoch_f1 = f1_score(all_labels, all_preds, average=\"weighted\")\n    epoch_qwk = cohen_kappa_score(all_labels, all_preds, weights=\"quadratic\")\n\n    return epoch_loss, epoch_acc, epoch_f1, epoch_qwk, all_labels, all_preds","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:55.158731Z","iopub.execute_input":"2026-05-01T11:21:55.159005Z","iopub.status.idle":"2026-05-01T11:21:55.170742Z","shell.execute_reply.started":"2026-05-01T11:21:55.158979Z","shell.execute_reply":"2026-05-01T11:21:55.170080Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"STEP 5.16 — Run 2-Epoch GPU Test","metadata":{}},{"cell_type":"code","source":"num_epochs = 2\n\nfor epoch in range(num_epochs):\n    train_loss, train_acc, train_f1, train_qwk = train_one_epoch(\n        model, train_loader, criterion, optimizer, device\n    )\n\n    val_loss, val_acc, val_f1, val_qwk, y_true, y_pred = validate_one_epoch(\n        model, val_loader, criterion, device\n    )\n\n    print(f\"Epoch {epoch+1}/{num_epochs}\")\n    print(f\"Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.4f} | Train F1: {train_f1:.4f} | Train QWK: {train_qwk:.4f}\")\n    print(f\"Val   Loss: {val_loss:.4f} | Val   Acc: {val_acc:.4f} | Val   F1: {val_f1:.4f} | Val   QWK: {val_qwk:.4f}\")\n    print(\"-\" * 80)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:55.171645Z","iopub.execute_input":"2026-05-01T11:21:55.171932Z","iopub.status.idle":"2026-05-01T11:21:58.470840Z","shell.execute_reply.started":"2026-05-01T11:21:55.171898Z","shell.execute_reply":"2026-05-01T11:21:58.467535Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"these are very strong and stable baseline results.","metadata":{}},{"cell_type":"markdown","source":"Most Important Observation\nValidation QWK increased from 0.9044 → 0.9091\n\nThis means:\n\nmodel is learning correctly\nno overfitting yet\ntraining is stable\npipeline is correct\n\nThe baseline EfficientNet-B5 model demonstrated stable learning behavior on the fused APTOS and IDRiD dataset, achieving validation QWK above 0.90 within the first two epochs, indicating effective preprocessing and dataset fusion.","metadata":{}},{"cell_type":"code","source":"import torch\n\nMODEL_PATH = \"/kaggle/working/baseline_efficientnetb5_model.pth\"\n\ntorch.save(\n    model.state_dict(),\n    MODEL_PATH\n)\n\nprint(\"Model saved at:\", MODEL_PATH)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.471996Z","iopub.status.idle":"2026-05-01T11:21:58.472308Z","shell.execute_reply.started":"2026-05-01T11:21:58.472179Z","shell.execute_reply":"2026-05-01T11:21:58.472195Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**We have already completed:**\n\ndataset loading\npreprocessing\nfusion\nfold split\nbaseline sanity training\nlocal backup of notebook + model","metadata":{}},{"cell_type":"markdown","source":"STEP 1 — Run 10-Epoch Baseline Training With Saving","metadata":{}},{"cell_type":"code","source":"best_qwk = 0.0\n\n# optional: save current 2-epoch state first\ntorch.save(\n    {\n        \"model_state_dict\": model.state_dict(),\n        \"optimizer_state_dict\": optimizer.state_dict(),\n        \"epoch\": 2\n    },\n    \"/kaggle/working/baseline_checkpoint_epoch2.pth\"\n)\n\nnum_epochs = 5\n\nfor epoch in range(num_epochs):\n\n    train_loss, train_acc, train_f1, train_qwk = train_one_epoch(\n        model, train_loader, criterion, optimizer, device\n    )\n\n    val_loss, val_acc, val_f1, val_qwk, y_true, y_pred = validate_one_epoch(\n        model, val_loader, criterion, device\n    )\n\n    print(f\"Epoch {epoch+1}/{num_epochs}\")\n    print(f\"Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.4f} | Train F1: {train_f1:.4f} | Train QWK: {train_qwk:.4f}\")\n    print(f\"Val   Loss: {val_loss:.4f} | Val   Acc: {val_acc:.4f} | Val   F1: {val_f1:.4f} | Val   QWK: {val_qwk:.4f}\")\n\n    if val_qwk > best_qwk:\n        best_qwk = val_qwk\n        torch.save(model.state_dict(), \"/kaggle/working/best_baseline_model.pth\")\n        torch.save(\n            {\n                \"model_state_dict\": model.state_dict(),\n                \"optimizer_state_dict\": optimizer.state_dict(),\n                \"epoch\": epoch + 1,\n                \"best_qwk\": best_qwk\n            },\n            \"/kaggle/working/best_baseline_checkpoint.pth\"\n        )\n        print(\"Saved new best model\")\n\n    print(\"-\" * 80)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.473068Z","iopub.status.idle":"2026-05-01T11:21:58.473371Z","shell.execute_reply.started":"2026-05-01T11:21:58.473210Z","shell.execute_reply":"2026-05-01T11:21:58.473227Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Final baseline summary\nBest validation result:\n\nEpoch 3\n\nVal Accuracy: 0.8409\nVal F1: 0.8309\nVal QWK: 0.9182 ← best\n\nThe baseline EfficientNet-B5 model achieved its best validation performance at Epoch 3 with a QWK of 0.9182. After this point, training performance continued improving while validation performance fluctuated and slightly declined, indicating early overfitting.","metadata":{}},{"cell_type":"code","source":"import torch\n\nMODEL_PATH = \"/kaggle/working/best_baseline_model_final.pth\"\n\ntorch.save(\n    model.state_dict(),\n    MODEL_PATH\n)\n\nprint(\"Model saved at:\", MODEL_PATH)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.474682Z","iopub.status.idle":"2026-05-01T11:21:58.474928Z","shell.execute_reply.started":"2026-05-01T11:21:58.474809Z","shell.execute_reply":"2026-05-01T11:21:58.474824Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\nfor root, dirs, files in os.walk(\"/kaggle/input\"):\n    print(root)\n    for f in files:\n        print(\"   \", f)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.476264Z","iopub.status.idle":"2026-05-01T11:21:58.476568Z","shell.execute_reply.started":"2026-05-01T11:21:58.476390Z","shell.execute_reply":"2026-05-01T11:21:58.476412Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\nAPTOS_LABELS = \"/kaggle/input/competitions/aptos2019-blindness-detection/train.csv\"\nAPTOS_IMAGES = \"/kaggle/input/competitions/aptos2019-blindness-detection/train_images\"\n\naptos_labels = pd.read_csv(APTOS_LABELS)\n\naptos_df = aptos_labels.copy()\n\naptos_df[\"image_path\"] = aptos_df[\"id_code\"].apply(\n    lambda x: f\"{APTOS_IMAGES}/{x}.png\"\n)\n\naptos_df.rename(columns={\"diagnosis\": \"label\"}, inplace=True)\n\naptos_df[\"source\"] = \"aptos\"\n\naptos_df = aptos_df[[\"image_path\", \"label\", \"source\"]]\n\nprint(\"APTOS shape:\", aptos_df.shape)\naptos_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.477773Z","iopub.status.idle":"2026-05-01T11:21:58.478038Z","shell.execute_reply.started":"2026-05-01T11:21:58.477906Z","shell.execute_reply":"2026-05-01T11:21:58.477921Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"IDRID_TRAIN_LABELS = \"/kaggle/input/datasets/abdullahshafi315/indian-diabetic-retinopathy-image-datasetidrid/Disease Grading/2. Groundtruths/a. IDRiD_Disease Grading_Training Labels.csv\"\n\nIDRID_TRAIN_IMAGES = \"/kaggle/input/datasets/abdullahshafi315/indian-diabetic-retinopathy-image-datasetidrid/Disease Grading/1. Original Images/a. Training Set\"\n\nidrid_labels = pd.read_csv(IDRID_TRAIN_LABELS)\n\nidrid_df = idrid_labels.copy()\n\nidrid_df[\"image_path\"] = idrid_df[\"Image name\"].apply(\n    lambda x: f\"{IDRID_TRAIN_IMAGES}/{x}.jpg\"\n)\n\nidrid_df.rename(\n    columns={\"Retinopathy grade\": \"label\"},\n    inplace=True\n)\n\nidrid_df[\"source\"] = \"idrid\"\n\nidrid_df = idrid_df[[\"image_path\", \"label\", \"source\"]]\n\nprint(\"IDRiD shape:\", idrid_df.shape)\nidrid_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.479734Z","iopub.status.idle":"2026-05-01T11:21:58.480029Z","shell.execute_reply.started":"2026-05-01T11:21:58.479889Z","shell.execute_reply":"2026-05-01T11:21:58.479907Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"combined_df = pd.concat(\n    [aptos_df, idrid_df],\n    ignore_index=True\n)\n\nprint(\"Combined shape:\", combined_df.shape)\n\ncombined_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.481560Z","iopub.status.idle":"2026-05-01T11:21:58.482021Z","shell.execute_reply.started":"2026-05-01T11:21:58.481769Z","shell.execute_reply":"2026-05-01T11:21:58.481796Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.model_selection import StratifiedKFold\n\nskf = StratifiedKFold(\n    n_splits=5,\n    shuffle=True,\n    random_state=42\n)\n\ncombined_df[\"fold\"] = -1\n\nfor fold, (_, val_idx) in enumerate(\n    skf.split(combined_df, combined_df[\"label\"])\n):\n    combined_df.loc[val_idx, \"fold\"] = fold\n\nprint(combined_df[\"fold\"].value_counts())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.483559Z","iopub.status.idle":"2026-05-01T11:21:58.483998Z","shell.execute_reply.started":"2026-05-01T11:21:58.483787Z","shell.execute_reply":"2026-05-01T11:21:58.483816Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"val_df = combined_df[\n    combined_df[\"fold\"] == 0\n].reset_index(drop=True)\n\nprint(\"Val shape:\", val_df.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.484990Z","iopub.status.idle":"2026-05-01T11:21:58.485252Z","shell.execute_reply.started":"2026-05-01T11:21:58.485124Z","shell.execute_reply":"2026-05-01T11:21:58.485140Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"val_dataset = DRDataset(val_df, transform=val_transform)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=8,\n    shuffle=False,\n    num_workers=2,\n    pin_memory=True\n)\n\nprint(\"Val batches:\", len(val_loader))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.486287Z","iopub.status.idle":"2026-05-01T11:21:58.486573Z","shell.execute_reply.started":"2026-05-01T11:21:58.486437Z","shell.execute_reply":"2026-05-01T11:21:58.486452Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Test that loaded model works","metadata":{}},{"cell_type":"code","source":"images, labels = next(iter(val_loader))\n\nimages = images.to(device)\n\noutputs = model(images)\n\nprint(\"Output shape:\", outputs.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.487807Z","iopub.status.idle":"2026-05-01T11:21:58.488214Z","shell.execute_reply.started":"2026-05-01T11:21:58.488023Z","shell.execute_reply":"2026-05-01T11:21:58.488046Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import classification_report, confusion_matrix\nimport numpy as np\n\nmodel.eval()\n\nall_preds = []\nall_labels = []\n\nwith torch.no_grad():\n    for images, labels in val_loader:\n        images = images.to(device)\n        outputs = model(images)\n\n        preds = torch.argmax(outputs, dim=1).cpu().numpy()\n\n        all_preds.extend(preds)\n        all_labels.extend(labels.numpy())\n\nall_preds = np.array(all_preds)\nall_labels = np.array(all_labels)\n\nprint(\"Evaluation completed\")\nprint(\"Total predictions:\", len(all_preds))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.489280Z","iopub.status.idle":"2026-05-01T11:21:58.489668Z","shell.execute_reply.started":"2026-05-01T11:21:58.489459Z","shell.execute_reply":"2026-05-01T11:21:58.489482Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(classification_report(all_labels, all_preds, digits=4))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.490747Z","iopub.status.idle":"2026-05-01T11:21:58.491129Z","shell.execute_reply.started":"2026-05-01T11:21:58.490933Z","shell.execute_reply":"2026-05-01T11:21:58.490958Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cm = confusion_matrix(all_labels, all_preds)\nprint(cm)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.492490Z","iopub.status.idle":"2026-05-01T11:21:58.492818Z","shell.execute_reply.started":"2026-05-01T11:21:58.492637Z","shell.execute_reply":"2026-05-01T11:21:58.492653Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\nplt.figure(figsize=(8,6))\nplt.imshow(cm, interpolation=\"nearest\", cmap=\"Blues\")\nplt.title(\"Confusion Matrix - Baseline EfficientNetB5\")\nplt.colorbar()\n\nclasses = [0, 1, 2, 3, 4]\ntick_marks = np.arange(len(classes))\n\nplt.xticks(tick_marks, classes)\nplt.yticks(tick_marks, classes)\nplt.xlabel(\"Predicted Label\")\nplt.ylabel(\"True Label\")\n\nfor i in range(cm.shape[0]):\n    for j in range(cm.shape[1]):\n        plt.text(j, i, cm[i, j], ha=\"center\", va=\"center\", color=\"black\")\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.493908Z","iopub.status.idle":"2026-05-01T11:21:58.494131Z","shell.execute_reply.started":"2026-05-01T11:21:58.494023Z","shell.execute_reply":"2026-05-01T11:21:58.494037Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Model is working correctly ✅\nTraining was successful ✅\nPerformance is realistic for baseline ✅\nNo bugs or data leakage detected ✅\n\nThis is a valid baseline result for research.\n\n\n| Class | Meaning       | Recall    | Interpretation |\n| ----- | ------------- | --------- | -------------- |\n| 0     | No DR         | **0.997** | Excellent      |\n| 1     | Mild          | **0.141** | Very weak      |\n| 2     | Moderate      | **0.956** | Strong         |\n| 3     | Severe        | **0.170** | Weak           |\n| 4     | Proliferative | **0.638** | Acceptable     |\n","metadata":{}},{"cell_type":"markdown","source":"Compute Class Weights","metadata":{}},{"cell_type":"code","source":"train_df = combined_df[\n    combined_df[\"fold\"] != 0\n].reset_index(drop=True)\n\nval_df = combined_df[\n    combined_df[\"fold\"] == 0\n].reset_index(drop=True)\n\nprint(\"Train shape:\", train_df.shape)\nprint(\"Val shape:\", val_df.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.496371Z","iopub.status.idle":"2026-05-01T11:21:58.496706Z","shell.execute_reply.started":"2026-05-01T11:21:58.496545Z","shell.execute_reply":"2026-05-01T11:21:58.496561Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport numpy as np\n\nclass_counts = (\n    train_df[\"label\"]\n    .value_counts()\n    .sort_index()\n    .values\n)\n\nprint(\"Class counts:\", class_counts)\n\nclass_weights = 1.0 / class_counts\n\nclass_weights = class_weights / class_weights.sum()\n\nclass_weights = torch.tensor(\n    class_weights,\n    dtype=torch.float32\n)\n\nprint(\"Class weights:\", class_weights)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.498009Z","iopub.status.idle":"2026-05-01T11:21:58.498311Z","shell.execute_reply.started":"2026-05-01T11:21:58.498150Z","shell.execute_reply":"2026-05-01T11:21:58.498165Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Use Weights in Loss Function","metadata":{}},{"cell_type":"code","source":"criterion = torch.nn.CrossEntropyLoss(\n    weight=class_weights.to(device)\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.499635Z","iopub.status.idle":"2026-05-01T11:21:58.500000Z","shell.execute_reply.started":"2026-05-01T11:21:58.499809Z","shell.execute_reply":"2026-05-01T11:21:58.499832Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Now We Start Weighted Training (5 Epochs)\n\nWe will:\n\ntrain baseline again\nuse class weights\nautomatically save best model\ncompare improvement","metadata":{}},{"cell_type":"markdown","source":"Run Weighted Training","metadata":{}},{"cell_type":"code","source":"# ==============================\n# 1. IMPORTS\n# ==============================\nimport os\nimport cv2\nimport torch\nimport random\nimport numpy as np\nimport pandas as pd\nimport torch.nn as nn\nimport torch.optim as optim\nfrom sklearn.model_selection import StratifiedKFold\nfrom sklearn.metrics import accuracy_score, f1_score, cohen_kappa_score\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nfrom torchvision.models import efficientnet_b5, EfficientNet_B5_Weights\n\n# ==============================\n# 2. DEVICE + SEED\n# ==============================\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Using device:\", device)\n\ndef set_seed(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n\nset_seed(42)\n\n# ==============================\n# 3. PATHS\n# ==============================\nAPTOS_LABELS = \"/kaggle/input/competitions/aptos2019-blindness-detection/train.csv\"\nAPTOS_IMAGES = \"/kaggle/input/competitions/aptos2019-blindness-detection/train_images\"\n\nIDRID_BASE = \"/kaggle/input/datasets/abdullahshafi315/indian-diabetic-retinopathy-image-datasetidrid/Disease Grading\"\nIDRID_TRAIN_LABELS = f\"{IDRID_BASE}/2. Groundtruths/a. IDRiD_Disease Grading_Training Labels.csv\"\nIDRID_TEST_LABELS  = f\"{IDRID_BASE}/2. Groundtruths/b. IDRiD_Disease Grading_Testing Labels.csv\"\nIDRID_TRAIN_IMAGES = f\"{IDRID_BASE}/1. Original Images/a. Training Set\"\nIDRID_TEST_IMAGES  = f\"{IDRID_BASE}/1. Original Images/b. Testing Set\"\n\n# OPTIONAL: previously saved baseline model path\nMODEL_PATH = \"/kaggle/input/models/afnanhalim/effnetb5/pytorch/default/1/baseline_efficientnetb5_model.pth\"\n\n# ==============================\n# 4. REBUILD DATAFRAMES\n# ==============================\n# APTOS\naptos_labels = pd.read_csv(APTOS_LABELS)\naptos_df = aptos_labels.copy()\naptos_df[\"image_path\"] = aptos_df[\"id_code\"].apply(lambda x: f\"{APTOS_IMAGES}/{x}.png\")\naptos_df.rename(columns={\"diagnosis\": \"label\"}, inplace=True)\naptos_df[\"source\"] = \"aptos\"\naptos_df = aptos_df[[\"image_path\", \"label\", \"source\"]]\n\n# IDRiD train\nidrid_train = pd.read_csv(IDRID_TRAIN_LABELS)\nidrid_train_df = idrid_train.copy()\nidrid_train_df[\"image_path\"] = idrid_train_df[\"Image name\"].apply(\n    lambda x: f\"{IDRID_TRAIN_IMAGES}/{x}.jpg\"\n)\nidrid_train_df.rename(columns={\"Retinopathy grade\": \"label\"}, inplace=True)\nidrid_train_df[\"source\"] = \"idrid\"\nidrid_train_df = idrid_train_df[[\"image_path\", \"label\", \"source\"]]\n\n# IDRiD test\nidrid_test = pd.read_csv(IDRID_TEST_LABELS)\nidrid_test_df = idrid_test.copy()\nidrid_test_df[\"image_path\"] = idrid_test_df[\"Image name\"].apply(\n    lambda x: f\"{IDRID_TEST_IMAGES}/{x}.jpg\"\n)\nidrid_test_df.rename(columns={\"Retinopathy grade\": \"label\"}, inplace=True)\nidrid_test_df[\"source\"] = \"idrid\"\nidrid_test_df = idrid_test_df[[\"image_path\", \"label\", \"source\"]]\n\n# combine\ncombined_df = pd.concat([aptos_df, idrid_train_df, idrid_test_df], ignore_index=True)\nprint(\"Combined shape:\", combined_df.shape)\n\n# folds\ncombined_df[\"fold\"] = -1\nskf = StratifiedKFold(n_splits=5, shuffle=True, random_state=42)\n\nfor fold, (_, val_idx) in enumerate(skf.split(combined_df, combined_df[\"label\"])):\n    combined_df.loc[val_idx, \"fold\"] = fold\n\ntrain_df = combined_df[combined_df[\"fold\"] != 0].reset_index(drop=True)\nval_df   = combined_df[combined_df[\"fold\"] == 0].reset_index(drop=True)\n\nprint(\"Train shape:\", train_df.shape)\nprint(\"Val shape:\", val_df.shape)\n\n# ==============================\n# 5. PREPROCESSING FUNCTIONS\n# ==============================\nIMG_SIZE = 456\n\ndef circular_crop(img):\n    gray = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)\n    _, thresh = cv2.threshold(gray, 10, 255, cv2.THRESH_BINARY)\n    contours, _ = cv2.findContours(thresh, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n\n    if len(contours) == 0:\n        return img\n\n    cnt = max(contours, key=cv2.contourArea)\n    x, y, w, h = cv2.boundingRect(cnt)\n    return img[y:y+h, x:x+w]\n\ndef apply_clahe(img):\n    lab = cv2.cvtColor(img, cv2.COLOR_BGR2LAB)\n    l, a, b = cv2.split(lab)\n\n    clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))\n    l = clahe.apply(l)\n\n    lab = cv2.merge((l, a, b))\n    img = cv2.cvtColor(lab, cv2.COLOR_LAB2BGR)\n    return img\n\n# ==============================\n# 6. TRANSFORMS\n# ==============================\ntrain_transform = transforms.Compose([\n    transforms.ToPILImage(),\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.1, hue=0.05),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225]\n    )\n])\n\nval_transform = transforms.Compose([\n    transforms.ToPILImage(),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225]\n    )\n])\n\n# ==============================\n# 7. DATASET CLASS\n# ==============================\nclass DRDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df.reset_index(drop=True)\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        image_path = row[\"image_path\"]\n        label = int(row[\"label\"])\n\n        img = cv2.imread(image_path)\n        img = circular_crop(img)\n        img = apply_clahe(img)\n        img = cv2.resize(img, (IMG_SIZE, IMG_SIZE))\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n\n        if self.transform:\n            img = self.transform(img)\n\n        return img, label\n\n# ==============================\n# 8. DATALOADERS\n# ==============================\ntrain_dataset = DRDataset(train_df, transform=train_transform)\nval_dataset   = DRDataset(val_df, transform=val_transform)\n\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=8,\n    shuffle=True,\n    num_workers=2,\n    pin_memory=True\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=8,\n    shuffle=False,\n    num_workers=2,\n    pin_memory=True\n)\n\nprint(\"Train batches:\", len(train_loader))\nprint(\"Val batches:\", len(val_loader))\n\n# ==============================\n# 9. BUILD MODEL\n# ==============================\nweights = EfficientNet_B5_Weights.IMAGENET1K_V1\nmodel = efficientnet_b5(weights=weights)\n\nin_features = model.classifier[1].in_features\nmodel.classifier[1] = nn.Linear(in_features, 5)\n\nmodel = model.to(device)\n\nprint(model.classifier)\n\n# OPTIONAL: load your previous saved baseline weights\nif os.path.exists(MODEL_PATH):\n    model.load_state_dict(torch.load(MODEL_PATH, map_location=device))\n    print(\"Previous baseline model loaded successfully\")\nelse:\n    print(\"Previous model path not found, training from current initialized model\")\n\n# ==============================\n# 10. CLASS WEIGHTS\n# ==============================\nclass_counts = train_df[\"label\"].value_counts().sort_index().values\nprint(\"Class counts:\", class_counts)\n\nclass_weights = 1.0 / class_counts\nclass_weights = class_weights / class_weights.sum()\nclass_weights = torch.tensor(class_weights, dtype=torch.float32).to(device)\n\nprint(\"Class weights:\", class_weights)\n\n# ==============================\n# 11. LOSS + OPTIMIZER\n# ==============================\ncriterion = nn.CrossEntropyLoss(weight=class_weights)\n\noptimizer = optim.AdamW(\n    model.parameters(),\n    lr=1e-4,\n    weight_decay=1e-4\n)\n\n# ==============================\n# 12. TRAIN / VALID FUNCTIONS\n# ==============================\ndef train_one_epoch(model, loader, criterion, optimizer, device):\n    model.train()\n\n    running_loss = 0.0\n    all_preds = []\n    all_labels = []\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\n        outputs = model(images)\n        loss = criterion(outputs, labels)\n\n        loss.backward()\n        optimizer.step()\n\n        running_loss += loss.item() * images.size(0)\n\n        preds = torch.argmax(outputs, dim=1)\n        all_preds.extend(preds.detach().cpu().numpy())\n        all_labels.extend(labels.detach().cpu().numpy())\n\n    epoch_loss = running_loss / len(loader.dataset)\n    epoch_acc = accuracy_score(all_labels, all_preds)\n    epoch_f1 = f1_score(all_labels, all_preds, average=\"weighted\")\n    epoch_qwk = cohen_kappa_score(all_labels, all_preds, weights=\"quadratic\")\n\n    return epoch_loss, epoch_acc, epoch_f1, epoch_qwk\n\n\ndef validate_one_epoch(model, loader, criterion, device):\n    model.eval()\n\n    running_loss = 0.0\n    all_preds = []\n    all_labels = []\n\n    with torch.no_grad():\n        for images, labels in loader:\n            images = images.to(device, non_blocking=True)\n            labels = labels.to(device, non_blocking=True)\n\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n\n            running_loss += loss.item() * images.size(0)\n\n            preds = torch.argmax(outputs, dim=1)\n            all_preds.extend(preds.detach().cpu().numpy())\n            all_labels.extend(labels.detach().cpu().numpy())\n\n    epoch_loss = running_loss / len(loader.dataset)\n    epoch_acc = accuracy_score(all_labels, all_preds)\n    epoch_f1 = f1_score(all_labels, all_preds, average=\"weighted\")\n    epoch_qwk = cohen_kappa_score(all_labels, all_preds, weights=\"quadratic\")\n\n    return epoch_loss, epoch_acc, epoch_f1, epoch_qwk, all_labels, all_preds\n\n# ==============================\n# 13. TRAIN WEIGHTED MODEL\n# ==============================\nnum_epochs = 5\nbest_qwk = 0.0\n\nfor epoch in range(num_epochs):\n\n    train_loss, train_acc, train_f1, train_qwk = train_one_epoch(\n        model, train_loader, criterion, optimizer, device\n    )\n\n    val_loss, val_acc, val_f1, val_qwk, y_true, y_pred = validate_one_epoch(\n        model, val_loader, criterion, device\n    )\n\n    print(f\"Epoch {epoch+1}/{num_epochs}\")\n    print(f\"Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.4f} | Train F1: {train_f1:.4f} | Train QWK: {train_qwk:.4f}\")\n    print(f\"Val   Loss: {val_loss:.4f} | Val   Acc: {val_acc:.4f} | Val   F1: {val_f1:.4f} | Val   QWK: {val_qwk:.4f}\")\n\n    if val_qwk > best_qwk:\n        best_qwk = val_qwk\n        torch.save(model.state_dict(), \"/kaggle/working/weighted_best_model.pth\")\n        print(\"Saved new best weighted model\")\n\n    print(\"-\" * 80)\n\nprint(\"Training finished. Best Val QWK:\", best_qwk)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.501227Z","iopub.status.idle":"2026-05-01T11:21:58.501606Z","shell.execute_reply.started":"2026-05-01T11:21:58.501403Z","shell.execute_reply":"2026-05-01T11:21:58.501426Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"\n**This code trains a computer model (EfficientNet-B5) to look at eye images and decide how serious the diabetic eye disease is. It first collects images from two datasets (APTOS and IDRiD), prepares the images by cleaning and resizing them, and then splits the data into training and validation groups. After that, it calculates class weights so the model pays more attention to rare disease cases. Then it trains the model for 5 rounds (epochs), checks how well it performs, and saves the best model automatically.In simple words, this code teaches the model to recognize disease severity levels from eye images and keeps the best trained version.**\n","metadata":{}},{"cell_type":"markdown","source":"weighted training results:\n\nEpoch 1: Val QWK 0.9278 (best)\nEpoch 2–5: Val QWK decreased\n\nWe should NOT jump yet.\nFirst stabilize training using early stopping.\nThen move to CBAM. \nBefore CBAM, the proper next improvement is:\n\nEarly Stopping + Learning Rate Scheduler\n\nSo we need training that:\n\nstops automatically when performance stops improving\nreduces learning rate when plateau happens\n\nThis improves stability without changing architecture.","metadata":{}},{"cell_type":"markdown","source":"This code trains your diabetic retinopathy model in a safer and smarter way. It loads the fused APTOS + IDRiD dataset, applies your final preprocessing (circular crop + CLAHE + resizing), creates training and validation splits, and builds the EfficientNet-B5 model. Then it uses class-weighted loss so rare classes get more attention, and adds early stopping so training stops automatically when validation performance stops improving. It also uses a learning rate scheduler to reduce the learning rate when improvement slows down, and saves the best model automatically. In simple words, this code tries to improve your result while also preventing unnecessary overfitting and protecting your best checkpoint.","metadata":{}},{"cell_type":"code","source":"# =========================================\n# 1. IMPORTS\n# =========================================\nimport os\nimport cv2\nimport torch\nimport random\nimport numpy as np\nimport pandas as pd\nimport torch.nn as nn\nimport torch.optim as optim\nfrom sklearn.model_selection import StratifiedKFold\nfrom sklearn.metrics import accuracy_score, f1_score, cohen_kappa_score\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nfrom torchvision.models import efficientnet_b5, EfficientNet_B5_Weights\n\n# =========================================\n# 2. DEVICE + SEED\n# =========================================\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Using device:\", device)\n\ndef set_seed(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n\nset_seed(42)\n\n# =========================================\n# 3. PATHS\n# =========================================\nAPTOS_LABELS = \"/kaggle/input/competitions/aptos2019-blindness-detection/train.csv\"\nAPTOS_IMAGES = \"/kaggle/input/competitions/aptos2019-blindness-detection/train_images\"\n\nIDRID_BASE = \"/kaggle/input/datasets/abdullahshafi315/indian-diabetic-retinopathy-image-datasetidrid/Disease Grading\"\nIDRID_TRAIN_LABELS = f\"{IDRID_BASE}/2. Groundtruths/a. IDRiD_Disease Grading_Training Labels.csv\"\nIDRID_TEST_LABELS  = f\"{IDRID_BASE}/2. Groundtruths/b. IDRiD_Disease Grading_Testing Labels.csv\"\nIDRID_TRAIN_IMAGES = f\"{IDRID_BASE}/1. Original Images/a. Training Set\"\nIDRID_TEST_IMAGES  = f\"{IDRID_BASE}/1. Original Images/b. Testing Set\"\n\n# your previously uploaded baseline model\nMODEL_PATH = \"/kaggle/input/models/afnanhalim/effnetb5/pytorch/default/1/baseline_efficientnetb5_model.pth\"\n\n# =========================================\n# 4. REBUILD DATAFRAMES\n# =========================================\n# APTOS\naptos_labels = pd.read_csv(APTOS_LABELS)\naptos_df = aptos_labels.copy()\naptos_df[\"image_path\"] = aptos_df[\"id_code\"].apply(lambda x: f\"{APTOS_IMAGES}/{x}.png\")\naptos_df.rename(columns={\"diagnosis\": \"label\"}, inplace=True)\naptos_df[\"source\"] = \"aptos\"\naptos_df = aptos_df[[\"image_path\", \"label\", \"source\"]]\n\n# IDRiD train\nidrid_train = pd.read_csv(IDRID_TRAIN_LABELS)\nidrid_train_df = idrid_train.copy()\nidrid_train_df[\"image_path\"] = idrid_train_df[\"Image name\"].apply(\n    lambda x: f\"{IDRID_TRAIN_IMAGES}/{x}.jpg\"\n)\nidrid_train_df.rename(columns={\"Retinopathy grade\": \"label\"}, inplace=True)\nidrid_train_df[\"source\"] = \"idrid\"\nidrid_train_df = idrid_train_df[[\"image_path\", \"label\", \"source\"]]\n\n# IDRiD test\nidrid_test = pd.read_csv(IDRID_TEST_LABELS)\nidrid_test_df = idrid_test.copy()\nidrid_test_df[\"image_path\"] = idrid_test_df[\"Image name\"].apply(\n    lambda x: f\"{IDRID_TEST_IMAGES}/{x}.jpg\"\n)\nidrid_test_df.rename(columns={\"Retinopathy grade\": \"label\"}, inplace=True)\nidrid_test_df[\"source\"] = \"idrid\"\nidrid_test_df = idrid_test_df[[\"image_path\", \"label\", \"source\"]]\n\n# combine\ncombined_df = pd.concat([aptos_df, idrid_train_df, idrid_test_df], ignore_index=True)\nprint(\"Combined shape:\", combined_df.shape)\n\n# 5-fold split\ncombined_df[\"fold\"] = -1\nskf = StratifiedKFold(n_splits=5, shuffle=True, random_state=42)\n\nfor fold, (_, val_idx) in enumerate(skf.split(combined_df, combined_df[\"label\"])):\n    combined_df.loc[val_idx, \"fold\"] = fold\n\ntrain_df = combined_df[combined_df[\"fold\"] != 0].reset_index(drop=True)\nval_df   = combined_df[combined_df[\"fold\"] == 0].reset_index(drop=True)\n\nprint(\"Train shape:\", train_df.shape)\nprint(\"Val shape:\", val_df.shape)\n\n# =========================================\n# 5. PREPROCESSING FUNCTIONS\n# =========================================\nIMG_SIZE = 456\n\ndef circular_crop(img):\n    gray = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)\n    _, thresh = cv2.threshold(gray, 10, 255, cv2.THRESH_BINARY)\n    contours, _ = cv2.findContours(thresh, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n\n    if len(contours) == 0:\n        return img\n\n    cnt = max(contours, key=cv2.contourArea)\n    x, y, w, h = cv2.boundingRect(cnt)\n    return img[y:y+h, x:x+w]\n\ndef apply_clahe(img):\n    lab = cv2.cvtColor(img, cv2.COLOR_BGR2LAB)\n    l, a, b = cv2.split(lab)\n\n    clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))\n    l = clahe.apply(l)\n\n    lab = cv2.merge((l, a, b))\n    img = cv2.cvtColor(lab, cv2.COLOR_LAB2BGR)\n    return img\n\n# =========================================\n# 6. TRANSFORMS\n# =========================================\ntrain_transform = transforms.Compose([\n    transforms.ToPILImage(),\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.1, hue=0.05),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225]\n    )\n])\n\nval_transform = transforms.Compose([\n    transforms.ToPILImage(),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225]\n    )\n])\n\n# =========================================\n# 7. DATASET CLASS\n# =========================================\nclass DRDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df.reset_index(drop=True)\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        image_path = row[\"image_path\"]\n        label = int(row[\"label\"])\n\n        img = cv2.imread(image_path)\n        img = circular_crop(img)\n        img = apply_clahe(img)\n        img = cv2.resize(img, (IMG_SIZE, IMG_SIZE))\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n\n        if self.transform:\n            img = self.transform(img)\n\n        return img, label\n\n# =========================================\n# 8. DATALOADERS\n# =========================================\ntrain_dataset = DRDataset(train_df, transform=train_transform)\nval_dataset   = DRDataset(val_df, transform=val_transform)\n\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=8,\n    shuffle=True,\n    num_workers=2,\n    pin_memory=True\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=8,\n    shuffle=False,\n    num_workers=2,\n    pin_memory=True\n)\n\nprint(\"Train batches:\", len(train_loader))\nprint(\"Val batches:\", len(val_loader))\n\n# =========================================\n# 9. MODEL\n# =========================================\nweights = EfficientNet_B5_Weights.IMAGENET1K_V1\nmodel = efficientnet_b5(weights=weights)\n\nin_features = model.classifier[1].in_features\nmodel.classifier[1] = nn.Linear(in_features, 5)\n\nmodel = model.to(device)\nprint(model.classifier)\n\n# load previous baseline model if available\nif os.path.exists(MODEL_PATH):\n    model.load_state_dict(torch.load(MODEL_PATH, map_location=device))\n    print(\"Previous baseline model loaded successfully\")\nelse:\n    print(\"Previous baseline model not found, continuing without it\")\n\n# =========================================\n# 10. CLASS WEIGHTS\n# =========================================\nclass_counts = train_df[\"label\"].value_counts().sort_index().values\nprint(\"Class counts:\", class_counts)\n\nclass_weights = 1.0 / class_counts\nclass_weights = class_weights / class_weights.sum()\nclass_weights = torch.tensor(class_weights, dtype=torch.float32).to(device)\n\nprint(\"Class weights:\", class_weights)\n\n# =========================================\n# 11. LOSS, OPTIMIZER, SCHEDULER\n# =========================================\ncriterion = nn.CrossEntropyLoss(weight=class_weights)\n\noptimizer = optim.AdamW(\n    model.parameters(),\n    lr=1e-4,\n    weight_decay=1e-4\n)\n\nscheduler = optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer,\n    mode=\"max\",\n    factor=0.5,\n    patience=1\n)\n\n# =========================================\n# 12. TRAIN / VALID FUNCTIONS\n# =========================================\ndef train_one_epoch(model, loader, criterion, optimizer, device):\n    model.train()\n\n    running_loss = 0.0\n    all_preds = []\n    all_labels = []\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\n        outputs = model(images)\n        loss = criterion(outputs, labels)\n\n        loss.backward()\n        optimizer.step()\n\n        running_loss += loss.item() * images.size(0)\n\n        preds = torch.argmax(outputs, dim=1)\n        all_preds.extend(preds.detach().cpu().numpy())\n        all_labels.extend(labels.detach().cpu().numpy())\n\n    epoch_loss = running_loss / len(loader.dataset)\n    epoch_acc = accuracy_score(all_labels, all_preds)\n    epoch_f1 = f1_score(all_labels, all_preds, average=\"weighted\")\n    epoch_qwk = cohen_kappa_score(all_labels, all_preds, weights=\"quadratic\")\n\n    return epoch_loss, epoch_acc, epoch_f1, epoch_qwk\n\n\ndef validate_one_epoch(model, loader, criterion, device):\n    model.eval()\n\n    running_loss = 0.0\n    all_preds = []\n    all_labels = []\n\n    with torch.no_grad():\n        for images, labels in loader:\n            images = images.to(device, non_blocking=True)\n            labels = labels.to(device, non_blocking=True)\n\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n\n            running_loss += loss.item() * images.size(0)\n\n            preds = torch.argmax(outputs, dim=1)\n            all_preds.extend(preds.detach().cpu().numpy())\n            all_labels.extend(labels.detach().cpu().numpy())\n\n    epoch_loss = running_loss / len(loader.dataset)\n    epoch_acc = accuracy_score(all_labels, all_preds)\n    epoch_f1 = f1_score(all_labels, all_preds, average=\"weighted\")\n    epoch_qwk = cohen_kappa_score(all_labels, all_preds, weights=\"quadratic\")\n\n    return epoch_loss, epoch_acc, epoch_f1, epoch_qwk, all_labels, all_preds\n\n# =========================================\n# 13. EARLY STOPPING TRAINING\n# =========================================\nnum_epochs = 10\nbest_qwk = 0.0\npatience = 2\ncounter = 0\n\nfor epoch in range(num_epochs):\n\n    train_loss, train_acc, train_f1, train_qwk = train_one_epoch(\n        model, train_loader, criterion, optimizer, device\n    )\n\n    val_loss, val_acc, val_f1, val_qwk, y_true, y_pred = validate_one_epoch(\n        model, val_loader, criterion, device\n    )\n\n    print(f\"Epoch {epoch+1}/{num_epochs}\")\n    print(f\"Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.4f} | Train F1: {train_f1:.4f} | Train QWK: {train_qwk:.4f}\")\n    print(f\"Val   Loss: {val_loss:.4f} | Val   Acc: {val_acc:.4f} | Val   F1: {val_f1:.4f} | Val   QWK: {val_qwk:.4f}\")\n\n    # scheduler step on validation QWK\n    scheduler.step(val_qwk)\n\n    # save best model\n    if val_qwk > best_qwk:\n        best_qwk = val_qwk\n        counter = 0\n\n        torch.save(\n            model.state_dict(),\n            \"/kaggle/working/weighted_earlystop_best_model.pth\"\n        )\n\n        torch.save(\n            {\n                \"model_state_dict\": model.state_dict(),\n                \"optimizer_state_dict\": optimizer.state_dict(),\n                \"best_qwk\": best_qwk,\n                \"epoch\": epoch + 1\n            },\n            \"/kaggle/working/weighted_earlystop_checkpoint.pth\"\n        )\n\n        print(\"Saved new best early-stopping model\")\n\n    else:\n        counter += 1\n        print(f\"No improvement. Early stop counter: {counter}/{patience}\")\n\n    current_lr = optimizer.param_groups[0][\"lr\"]\n    print(f\"Current learning rate: {current_lr:.8f}\")\n    print(\"-\" * 90)\n\n    if counter >= patience:\n        print(\"Early stopping triggered\")\n        break\n\nprint(\"Training finished\")\nprint(\"Best Val QWK:\", best_qwk)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.503138Z","iopub.status.idle":"2026-05-01T11:21:58.503502Z","shell.execute_reply.started":"2026-05-01T11:21:58.503319Z","shell.execute_reply":"2026-05-01T11:21:58.503341Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Weighted + early stopping is now your best current model\n\nBetter than:\n\nplain baseline QWK = 0.9182\nweighted + early stopping QWK = 0.9278\n\nSo this is a real improvement.\n\nImportant research meaning\n\nApplying class-weighted cross-entropy with early stopping improved validation QWK from 0.9182 to 0.9278. Early stopping successfully prevented further overfitting after the first epoch.","metadata":{}},{"cell_type":"markdown","source":"Now it makes sense to move to:\n\nCBAM model phase\n\nBecause:\n\nbaseline is complete\nimbalance handling is complete\nearly stopping is complete\nbest non-attention model is now fixed\n\nSo now CBAM becomes the next meaningful improvement.","metadata":{},"attachments":{}},{"cell_type":"code","source":"import os\nimport json\nimport torch\nimport pandas as pd\n\n# =========================\n# 1. CREATE SAVE FOLDER\n# =========================\nSAVE_DIR = \"/kaggle/working/final_saved_work\"\nos.makedirs(SAVE_DIR, exist_ok=True)\n\nprint(\"Save folder:\", SAVE_DIR)\n\n# =========================\n# 2. SAVE BEST MODEL WEIGHTS\n# =========================\nMODEL_WEIGHTS_PATH = os.path.join(SAVE_DIR, \"weighted_earlystop_best_model_final.pth\")\n\ntorch.save(\n    model.state_dict(),\n    MODEL_WEIGHTS_PATH\n)\n\nprint(\"Saved model weights:\", MODEL_WEIGHTS_PATH)\n\n# =========================\n# 3. SAVE FULL CHECKPOINT\n# =========================\nCHECKPOINT_PATH = os.path.join(SAVE_DIR, \"weighted_earlystop_checkpoint_final.pth\")\n\ntorch.save(\n    {\n        \"model_state_dict\": model.state_dict(),\n        \"optimizer_state_dict\": optimizer.state_dict(),\n        \"best_qwk\": float(best_qwk),\n        \"device\": str(device),\n        \"model_name\": \"EfficientNetB5\",\n        \"loss_type\": \"Weighted CrossEntropy + EarlyStopping\",\n        \"input_size\": 456,\n        \"num_classes\": 5\n    },\n    CHECKPOINT_PATH\n)\n\nprint(\"Saved checkpoint:\", CHECKPOINT_PATH)\n\n# =========================\n# 4. SAVE MASTER DATAFRAME\n# =========================\nCOMBINED_DF_PATH = os.path.join(SAVE_DIR, \"combined_df.csv\")\nTRAIN_DF_PATH = os.path.join(SAVE_DIR, \"train_df.csv\")\nVAL_DF_PATH = os.path.join(SAVE_DIR, \"val_df.csv\")\n\ncombined_df.to_csv(COMBINED_DF_PATH, index=False)\ntrain_df.to_csv(TRAIN_DF_PATH, index=False)\nval_df.to_csv(VAL_DF_PATH, index=False)\n\nprint(\"Saved combined_df:\", COMBINED_DF_PATH)\nprint(\"Saved train_df:\", TRAIN_DF_PATH)\nprint(\"Saved val_df:\", VAL_DF_PATH)\n\n# =========================\n# 5. SAVE CLASS WEIGHTS\n# =========================\nCLASS_WEIGHTS_PATH = os.path.join(SAVE_DIR, \"class_weights.pt\")\ntorch.save(class_weights.cpu(), CLASS_WEIGHTS_PATH)\n\nprint(\"Saved class weights:\", CLASS_WEIGHTS_PATH)\n\n# =========================\n# 6. SAVE TRAINING SUMMARY\n# =========================\nsummary = {\n    \"model_name\": \"EfficientNetB5\",\n    \"training_type\": \"Weighted CrossEntropy + EarlyStopping\",\n    \"best_validation_qwk\": float(best_qwk),\n    \"num_classes\": 5,\n    \"input_size\": 456,\n    \"datasets_used\": [\"APTOS 2019\", \"IDRiD 2018\"],\n    \"combined_dataset_size\": int(len(combined_df)),\n    \"train_size\": int(len(train_df)),\n    \"val_size\": int(len(val_df)),\n    \"class_counts_train\": train_df[\"label\"].value_counts().sort_index().to_dict(),\n    \"class_counts_val\": val_df[\"label\"].value_counts().sort_index().to_dict(),\n    \"preprocessing\": [\n        \"Circular Crop\",\n        \"CLAHE\",\n        \"Resize 456x456\",\n        \"ImageNet Normalization\"\n    ]\n}\n\nSUMMARY_JSON_PATH = os.path.join(SAVE_DIR, \"training_summary.json\")\n\nwith open(SUMMARY_JSON_PATH, \"w\") as f:\n    json.dump(summary, f, indent=4)\n\nprint(\"Saved training summary:\", SUMMARY_JSON_PATH)\n\n# =========================\n# 7. SAVE TEXT SUMMARY\n# =========================\nSUMMARY_TXT_PATH = os.path.join(SAVE_DIR, \"training_summary.txt\")\n\nwith open(SUMMARY_TXT_PATH, \"w\") as f:\n    f.write(\"Diabetic Retinopathy Project Summary\\n\")\n    f.write(\"===================================\\n\\n\")\n    f.write(\"Model: EfficientNetB5\\n\")\n    f.write(\"Training Type: Weighted CrossEntropy + EarlyStopping\\n\")\n    f.write(f\"Best Validation QWK: {best_qwk:.6f}\\n\")\n    f.write(f\"Combined Dataset Size: {len(combined_df)}\\n\")\n    f.write(f\"Train Size: {len(train_df)}\\n\")\n    f.write(f\"Validation Size: {len(val_df)}\\n\")\n    f.write(\"Datasets: APTOS 2019 + IDRiD 2018\\n\")\n    f.write(\"Preprocessing: Circular Crop -> CLAHE -> Resize 456x456 -> Normalize\\n\")\n\nprint(\"Saved text summary:\", SUMMARY_TXT_PATH)\n\n# =========================\n# 8. SHOW ALL SAVED FILES\n# =========================\nprint(\"\\nAll saved files:\")\nfor file_name in os.listdir(SAVE_DIR):\n    print(\"-\", file_name)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.504619Z","iopub.status.idle":"2026-05-01T11:21:58.504980Z","shell.execute_reply.started":"2026-05-01T11:21:58.504787Z","shell.execute_reply":"2026-05-01T11:21:58.504809Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import shutil\n\nzip_path = \"/kaggle/working/final_saved_work.zip\"\n\nshutil.make_archive(\n    base_name=\"/kaggle/working/final_saved_work\",\n    format=\"zip\",\n    root_dir=\"/kaggle/working\",\n    base_dir=\"final_saved_work\"\n)\n\nprint(\"ZIP file created at:\", zip_path)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.506049Z","iopub.status.idle":"2026-05-01T11:21:58.506501Z","shell.execute_reply.started":"2026-05-01T11:21:58.506289Z","shell.execute_reply":"2026-05-01T11:21:58.506315Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\nMODEL_DIR = \"/kaggle/input/models/afnanhalim/model-work-continued/pytorch/default/1\"\n\nprint(\"Files inside model directory:\\n\")\n\nfor file in os.listdir(MODEL_DIR):\n    print(file)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.508259Z","iopub.status.idle":"2026-05-01T11:21:58.508760Z","shell.execute_reply.started":"2026-05-01T11:21:58.508542Z","shell.execute_reply":"2026-05-01T11:21:58.508570Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport json\nimport shutil\nimport torch\nimport pandas as pd\n\n# =========================================\n# 1. CREATE FINAL FREEZE FOLDER\n# =========================================\nFREEZE_DIR = \"/kaggle/working/frozen_best_weighted_model\"\nos.makedirs(FREEZE_DIR, exist_ok=True)\n\nprint(\"Freeze folder created at:\", FREEZE_DIR)\n\n# =========================================\n# 2. SOURCE PATHS\n# =========================================\nMODEL_DIR = \"/kaggle/input/models/afnanhalim/model-work-continued/pytorch/default/1\"\n\nBEST_CHECKPOINT_INPUT = os.path.join(\n    MODEL_DIR,\n    \"weighted_earlystop_checkpoint_final.pth\"\n)\n\nCLASS_WEIGHTS_INPUT = os.path.join(\n    MODEL_DIR,\n    \"class_weights.pt\"\n)\n\nCOMBINED_DF_INPUT = os.path.join(\n    MODEL_DIR,\n    \"combined_df.csv\"\n)\n\n# If you also have this file in working, it can be copied too\nBEST_MODEL_WORKING = \"/kaggle/working/weighted_earlystop_best_model.pth\"\n\n# =========================================\n# 3. COPY BEST CHECKPOINT\n# =========================================\nif os.path.exists(BEST_CHECKPOINT_INPUT):\n    shutil.copy(BEST_CHECKPOINT_INPUT, FREEZE_DIR)\n    print(\"Copied checkpoint:\", BEST_CHECKPOINT_INPUT)\nelse:\n    print(\"Checkpoint file not found:\", BEST_CHECKPOINT_INPUT)\n\n# =========================================\n# 4. COPY BEST MODEL WEIGHTS (if available)\n# =========================================\nif os.path.exists(BEST_MODEL_WORKING):\n    shutil.copy(BEST_MODEL_WORKING, FREEZE_DIR)\n    print(\"Copied best model weights:\", BEST_MODEL_WORKING)\nelse:\n    print(\"Best model weights not found in /kaggle/working\")\n    print(\"This is okay if only checkpoint is available.\")\n\n# =========================================\n# 5. COPY CLASS WEIGHTS + DATAFRAME\n# =========================================\nif os.path.exists(CLASS_WEIGHTS_INPUT):\n    shutil.copy(CLASS_WEIGHTS_INPUT, FREEZE_DIR)\n    print(\"Copied class weights:\", CLASS_WEIGHTS_INPUT)\n\nif os.path.exists(COMBINED_DF_INPUT):\n    shutil.copy(COMBINED_DF_INPUT, FREEZE_DIR)\n    print(\"Copied combined dataframe:\", COMBINED_DF_INPUT)\n\n# =========================================\n# 6. SAVE TRAIN / VAL SPLITS AGAIN\n# =========================================\ntrain_df.to_csv(os.path.join(FREEZE_DIR, \"train_df_frozen.csv\"), index=False)\nval_df.to_csv(os.path.join(FREEZE_DIR, \"val_df_frozen.csv\"), index=False)\n\nprint(\"Saved frozen train/val split CSV files\")\n\n# =========================================\n# 7. SAVE FINAL EXPERIMENT SUMMARY\n# =========================================\nfinal_summary = {\n    \"frozen_model_name\": \"EfficientNetB5_Weighted_EarlyStopping\",\n    \"status\": \"FINAL_FROZEN_REFERENCE_MODEL\",\n    \"best_validation_qwk\": 0.927827167424068,\n    \"decision\": \"Do not continue more epochs. Use this as final non-CBAM reference model.\",\n    \"training_strategy\": \"Weighted CrossEntropy + EarlyStopping\",\n    \"datasets_used\": [\"APTOS 2019\", \"IDRiD 2018\"],\n    \"combined_dataset_size\": int(len(combined_df)),\n    \"train_size\": int(len(train_df)),\n    \"val_size\": int(len(val_df)),\n    \"input_size\": 456,\n    \"num_classes\": 5,\n    \"preprocessing\": [\n        \"Circular Crop\",\n        \"CLAHE\",\n        \"Resize 456x456\",\n        \"ImageNet Normalization\"\n    ],\n    \"best_result_interpretation\": \"This is the best current non-CBAM model and should be used as the fixed baseline for future CBAM comparison.\"\n}\n\nSUMMARY_JSON = os.path.join(FREEZE_DIR, \"frozen_model_summary.json\")\nwith open(SUMMARY_JSON, \"w\") as f:\n    json.dump(final_summary, f, indent=4)\n\nprint(\"Saved frozen summary JSON:\", SUMMARY_JSON)\n\n# =========================================\n# 8. SAVE HUMAN-READABLE TXT SUMMARY\n# =========================================\nSUMMARY_TXT = os.path.join(FREEZE_DIR, \"frozen_model_summary.txt\")\nwith open(SUMMARY_TXT, \"w\") as f:\n    f.write(\"FROZEN BEST MODEL SUMMARY\\n\")\n    f.write(\"=========================\\n\\n\")\n    f.write(\"Model: EfficientNetB5\\n\")\n    f.write(\"Training Strategy: Weighted CrossEntropy + EarlyStopping\\n\")\n    f.write(\"Status: FINAL FROZEN REFERENCE MODEL\\n\")\n    f.write(\"Best Validation QWK: 0.927827167424068\\n\")\n    f.write(\"Decision: Stop further epochs and keep this model fixed.\\n\")\n    f.write(\"Purpose: Use this as the final non-CBAM baseline for future comparison.\\n\")\n    f.write(f\"Combined Dataset Size: {len(combined_df)}\\n\")\n    f.write(f\"Train Size: {len(train_df)}\\n\")\n    f.write(f\"Validation Size: {len(val_df)}\\n\")\n    f.write(\"Preprocessing: Circular Crop -> CLAHE -> Resize 456x456 -> Normalize\\n\")\n\nprint(\"Saved frozen summary TXT:\", SUMMARY_TXT)\n\n# =========================================\n# 9. CREATE ZIP FILE\n# =========================================\nzip_path = \"/kaggle/working/frozen_best_weighted_model.zip\"\n\nshutil.make_archive(\n    base_name=\"/kaggle/working/frozen_best_weighted_model\",\n    format=\"zip\",\n    root_dir=\"/kaggle/working\",\n    base_dir=\"frozen_best_weighted_model\"\n)\n\nprint(\"ZIP file created at:\", zip_path)\n\n# =========================================\n# 10. SHOW ALL FILES\n# =========================================\nprint(\"\\nFrozen saved files:\")\nfor file_name in os.listdir(FREEZE_DIR):\n    print(\"-\", file_name)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.509947Z","iopub.status.idle":"2026-05-01T11:21:58.510295Z","shell.execute_reply.started":"2026-05-01T11:21:58.510133Z","shell.execute_reply":"2026-05-01T11:21:58.510153Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\nreport_dict = classification_report(\n    all_labels,\n    all_preds,\n    output_dict=True\n)\n\ndf_report = pd.DataFrame(report_dict).transpose()\n\ndf_report.to_csv(\n    \"/kaggle/working/baseline_classification_report.csv\"\n)\n\nprint(\"Baseline metrics saved\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.512485Z","iopub.status.idle":"2026-05-01T11:21:58.512877Z","shell.execute_reply.started":"2026-05-01T11:21:58.512690Z","shell.execute_reply":"2026-05-01T11:21:58.512714Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"What final baseline now shows\nOverall performance\nAccuracy: 0.8660\nWeighted F1: 0.8663\nMacro F1: 0.7648\n\nThis is better and more reliable than your earlier quick evaluation.\nVery strong\nClass 0 recall: 0.9848\nClass 2 recall: 0.8504\nImproved but still challenging\nClass 1 recall: 0.6456\nClass 4 recall: 0.6901\nWeakest class\nClass 3 recall: 0.6316","metadata":{}},{"cell_type":"markdown","source":"Next step = CBAM model implementation\n\nBecause now baseline is complete and saved.\n\nCBAM will try to improve:\n\nlesion-focused learning\ndifficult classes\nconfusion between nearby grades\n\nEspecially:\n\nClass 1\nClass 3\nClass 4","metadata":{}},{"cell_type":"markdown","source":"What we are doing now (CBAM Phase — Step 1)\n\nWe will:\n\nAdd CBAM attention block\nIntegrate CBAM into EfficientNet-B5\nKeep everything else SAME\nTrain and compare with frozen baseline","metadata":{}},{"cell_type":"markdown","source":"Step 1 — Define CBAM Module","metadata":{}},{"cell_type":"code","source":"# =========================================\n# STEP 1 — CBAM MODULE\n# =========================================\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\n# -------------------------------\n# Channel Attention\n# -------------------------------\n\nclass ChannelAttention(nn.Module):\n\n    def __init__(self, in_channels, reduction=16):\n\n        super(ChannelAttention, self).__init__()\n\n        self.avg_pool = nn.AdaptiveAvgPool2d(1)\n        self.max_pool = nn.AdaptiveMaxPool2d(1)\n\n        self.fc = nn.Sequential(\n\n            nn.Conv2d(\n                in_channels,\n                in_channels // reduction,\n                kernel_size=1,\n                bias=False\n            ),\n\n            nn.ReLU(),\n\n            nn.Conv2d(\n                in_channels // reduction,\n                in_channels,\n                kernel_size=1,\n                bias=False\n            )\n        )\n\n        self.sigmoid = nn.Sigmoid()\n\n    def forward(self, x):\n\n        avg_out = self.fc(self.avg_pool(x))\n        max_out = self.fc(self.max_pool(x))\n\n        out = avg_out + max_out\n\n        return self.sigmoid(out)\n\n\n# -------------------------------\n# Spatial Attention\n# -------------------------------\n\nclass SpatialAttention(nn.Module):\n\n    def __init__(self, kernel_size=7):\n\n        super(SpatialAttention, self).__init__()\n\n        self.conv = nn.Conv2d(\n            2,\n            1,\n            kernel_size,\n            padding=kernel_size // 2,\n            bias=False\n        )\n\n        self.sigmoid = nn.Sigmoid()\n\n    def forward(self, x):\n\n        avg_out = torch.mean(x, dim=1, keepdim=True)\n        max_out, _ = torch.max(x, dim=1, keepdim=True)\n\n        x = torch.cat([avg_out, max_out], dim=1)\n\n        x = self.conv(x)\n\n        return self.sigmoid(x)\n\n\n# -------------------------------\n# CBAM Block\n# -------------------------------\n\nclass CBAM(nn.Module):\n\n    def __init__(self, in_channels):\n\n        super(CBAM, self).__init__()\n\n        self.channel_attention = ChannelAttention(in_channels)\n        self.spatial_attention = SpatialAttention()\n\n    def forward(self, x):\n\n        x = x * self.channel_attention(x)\n\n        x = x * self.spatial_attention(x)\n\n        return x\n\n\nprint(\"CBAM module defined successfully\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.514313Z","iopub.status.idle":"2026-05-01T11:21:58.514748Z","shell.execute_reply.started":"2026-05-01T11:21:58.514509Z","shell.execute_reply":"2026-05-01T11:21:58.514547Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Step 2 — Build EfficientNet-B5 + CBAM Model","metadata":{}},{"cell_type":"code","source":"# =========================================\n# STEP 2 — EfficientNetB5 + CBAM\n# =========================================\n\nfrom torchvision.models import efficientnet_b5\nfrom torchvision.models import EfficientNet_B5_Weights\n\nclass EfficientNetB5_CBAM(nn.Module):\n\n    def __init__(self, num_classes=5):\n\n        super(EfficientNetB5_CBAM, self).__init__()\n\n        weights = EfficientNet_B5_Weights.IMAGENET1K_V1\n\n        self.backbone = efficientnet_b5(\n            weights=weights\n        )\n\n        in_features = self.backbone.classifier[1].in_features\n\n        # Remove original classifier\n        self.backbone.classifier = nn.Identity()\n\n        # CBAM after feature extractor\n        self.cbam = CBAM(in_channels=2048)\n\n        # New classifier\n        self.classifier = nn.Sequential(\n\n            nn.Dropout(0.4),\n\n            nn.Linear(\n                in_features,\n                num_classes\n            )\n        )\n\n    def forward(self, x):\n\n        x = self.backbone.features(x)\n\n        x = self.cbam(x)\n\n        x = F.adaptive_avg_pool2d(x, 1)\n\n        x = torch.flatten(x, 1)\n\n        x = self.classifier(x)\n\n        return x\n\n\n# Create model\nmodel = EfficientNetB5_CBAM()\n\nmodel = model.to(device)\n\nprint(\"EfficientNet-B5 + CBAM model built successfully\")\n\nprint(\"\\nModel structure:\")\nprint(model.classifier)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.515929Z","iopub.status.idle":"2026-05-01T11:21:58.516458Z","shell.execute_reply.started":"2026-05-01T11:21:58.516253Z","shell.execute_reply":"2026-05-01T11:21:58.516279Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"we are ready for the next correct step:\n\nTrain the CBAM model with the same setup as frozen baseline\n\nImportant rule:\n\nsame data\nsame preprocessing\nsame class weights\nsame optimizer style\nsame early stopping logic\n\nSo comparison stays fair.","metadata":{}},{"cell_type":"code","source":"# =========================================\n# STEP 3 — CBAM TRAINING SETUP\n# =========================================\n\nimport os\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\n\n# ----------------------------\n# Class weights loss\n# ----------------------------\ncriterion = nn.CrossEntropyLoss(\n    weight=class_weights\n)\n\n# ----------------------------\n# Optimizer\n# ----------------------------\noptimizer = optim.AdamW(\n    model.parameters(),\n    lr=1e-4,\n    weight_decay=1e-4\n)\n\n# ----------------------------\n# Scheduler\n# ----------------------------\nscheduler = optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer,\n    mode=\"max\",\n    factor=0.5,\n    patience=1\n)\n\nprint(\"CBAM training setup ready\")\n\n\n# =========================================\n# STEP 4 — TRAIN / VALID FUNCTIONS\n# =========================================\n\nfrom sklearn.metrics import accuracy_score, f1_score, cohen_kappa_score\n\ndef train_one_epoch(model, loader, criterion, optimizer, device):\n    model.train()\n\n    running_loss = 0.0\n    all_preds = []\n    all_labels = []\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\n        outputs = model(images)\n        loss = criterion(outputs, labels)\n\n        loss.backward()\n        optimizer.step()\n\n        running_loss += loss.item() * images.size(0)\n\n        preds = torch.argmax(outputs, dim=1)\n        all_preds.extend(preds.detach().cpu().numpy())\n        all_labels.extend(labels.detach().cpu().numpy())\n\n    epoch_loss = running_loss / len(loader.dataset)\n    epoch_acc = accuracy_score(all_labels, all_preds)\n    epoch_f1 = f1_score(all_labels, all_preds, average=\"weighted\")\n    epoch_qwk = cohen_kappa_score(all_labels, all_preds, weights=\"quadratic\")\n\n    return epoch_loss, epoch_acc, epoch_f1, epoch_qwk\n\n\ndef validate_one_epoch(model, loader, criterion, device):\n    model.eval()\n\n    running_loss = 0.0\n    all_preds = []\n    all_labels = []\n\n    with torch.no_grad():\n        for images, labels in loader:\n            images = images.to(device, non_blocking=True)\n            labels = labels.to(device, non_blocking=True)\n\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n\n            running_loss += loss.item() * images.size(0)\n\n            preds = torch.argmax(outputs, dim=1)\n            all_preds.extend(preds.detach().cpu().numpy())\n            all_labels.extend(labels.detach().cpu().numpy())\n\n    epoch_loss = running_loss / len(loader.dataset)\n    epoch_acc = accuracy_score(all_labels, all_preds)\n    epoch_f1 = f1_score(all_labels, all_preds, average=\"weighted\")\n    epoch_qwk = cohen_kappa_score(all_labels, all_preds, weights=\"quadratic\")\n\n    return epoch_loss, epoch_acc, epoch_f1, epoch_qwk, all_labels, all_preds\n\n\n# =========================================\n# STEP 5 — CBAM TRAINING WITH EARLY STOPPING\n# =========================================\n\nnum_epochs = 10\nbest_qwk = 0.0\npatience = 2\ncounter = 0\n\nfor epoch in range(num_epochs):\n\n    train_loss, train_acc, train_f1, train_qwk = train_one_epoch(\n        model, train_loader, criterion, optimizer, device\n    )\n\n    val_loss, val_acc, val_f1, val_qwk, y_true, y_pred = validate_one_epoch(\n        model, val_loader, criterion, device\n    )\n\n    print(f\"Epoch {epoch+1}/{num_epochs}\")\n    print(f\"Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.4f} | Train F1: {train_f1:.4f} | Train QWK: {train_qwk:.4f}\")\n    print(f\"Val   Loss: {val_loss:.4f} | Val   Acc: {val_acc:.4f} | Val   F1: {val_f1:.4f} | Val   QWK: {val_qwk:.4f}\")\n\n    scheduler.step(val_qwk)\n\n    if val_qwk > best_qwk:\n        best_qwk = val_qwk\n        counter = 0\n\n        torch.save(\n            model.state_dict(),\n            \"/kaggle/working/cbam_best_model.pth\"\n        )\n\n        torch.save(\n            {\n                \"model_state_dict\": model.state_dict(),\n                \"optimizer_state_dict\": optimizer.state_dict(),\n                \"best_qwk\": best_qwk\n            },\n            \"/kaggle/working/cbam_best_checkpoint.pth\"\n        )\n\n        print(\"Saved new best CBAM model\")\n\n    else:\n        counter += 1\n        print(f\"No improvement. Early stop counter: {counter}/{patience}\")\n\n    current_lr = optimizer.param_groups[0][\"lr\"]\n    print(f\"Current learning rate: {current_lr:.8f}\")\n    print(\"-\" * 90)\n\n    if counter >= patience:\n        print(\"Early stopping triggered\")\n        break\n\nprint(\"CBAM training finished\")\nprint(\"Best Val QWK:\", best_qwk)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.517755Z","iopub.status.idle":"2026-05-01T11:21:58.517992Z","shell.execute_reply.started":"2026-05-01T11:21:58.517878Z","shell.execute_reply":"2026-05-01T11:21:58.517892Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"The EfficientNet-B5 model trained with weighted loss and early stopping achieved the highest validation performance with a Quadratic Weighted Kappa (QWK) score of 0.9278. The addition of the CBAM attention mechanism resulted in a slightly lower best validation QWK of 0.9208. These results indicate that while CBAM provided stable learning behavior and improved feature attention, it did not surpass the performance of the optimized baseline model under the current training configuration.","metadata":{}},{"cell_type":"markdown","source":"Step — Freeze CBAM model results\n\nSave the CBAM experiment just like you froze the baseline.","metadata":{}},{"cell_type":"code","source":"import os\nimport shutil\nimport json\n\nCBAM_FREEZE_DIR = \"/kaggle/working/frozen_cbam_model\"\nos.makedirs(CBAM_FREEZE_DIR, exist_ok=True)\n\n# Copy best CBAM model\nshutil.copy(\n    \"/kaggle/working/cbam_best_checkpoint.pth\",\n    CBAM_FREEZE_DIR\n)\n\n# Save summary\ncbam_summary = {\n    \"model\": \"EfficientNetB5 + CBAM\",\n    \"status\": \"Experiment completed\",\n    \"best_val_qwk\": 0.9208219079232213,\n    \"early_stopping\": True,\n    \"comparison_result\": \"Baseline model performed better than CBAM\",\n    \"next_step\": \"Move to explainability phase (GradCAM)\"\n}\n\nwith open(\n    os.path.join(CBAM_FREEZE_DIR, \"cbam_summary.json\"),\n    \"w\"\n) as f:\n    json.dump(cbam_summary, f, indent=4)\n\nprint(\"CBAM model frozen successfully\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.519366Z","iopub.status.idle":"2026-05-01T11:21:58.519635Z","shell.execute_reply.started":"2026-05-01T11:21:58.519497Z","shell.execute_reply":"2026-05-01T11:21:58.519511Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"After freezing CBAM — the next real research step\n\nNow we move to:\n\nGrad-CAM Explainability Phase","metadata":{}},{"cell_type":"markdown","source":"Baseline model           ✓ completed\nWeighted loss            ✓ completed\nEarly stopping           ✓ completed\nCBAM experiment          ✓ completed\nModel comparison         ✓ completed","metadata":{}},{"cell_type":"markdown","source":"Next Step — Start Grad-CAM Visualization\n\nShow where the model is looking in the retina image\nVerify medical relevance of predictions\nSupport explainability in thesis","metadata":{}},{"cell_type":"markdown","source":"We will:\n\nLoad frozen best model\nSelect validation images\nGenerate Grad-CAM heatmaps","metadata":{}},{"cell_type":"markdown","source":"Step 1 — Load frozen best model for Grad-CAM","metadata":{}},{"cell_type":"code","source":"# =========================================\n# STEP 1 — LOAD FROZEN BEST MODEL\n# =========================================\n\nimport torch\nimport torch.nn as nn\nfrom torchvision.models import efficientnet_b5, EfficientNet_B5_Weights\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nMODEL_PATH = \"/kaggle/working/frozen_best_weighted_model/weighted_earlystop_checkpoint_final.pth\"\n\n# Build model exactly like baseline\nweights = EfficientNet_B5_Weights.IMAGENET1K_V1\n\nmodel = efficientnet_b5(weights=weights)\n\nin_features = model.classifier[1].in_features\n\nmodel.classifier[1] = nn.Linear(\n    in_features,\n    5\n)\n\nmodel = model.to(device)\n\ncheckpoint = torch.load(\n    MODEL_PATH,\n    map_location=device\n)\n\nmodel.load_state_dict(\n    checkpoint[\"model_state_dict\"]\n)\n\nmodel.eval()\n\nprint(\"Frozen baseline model loaded successfully\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.521129Z","iopub.status.idle":"2026-05-01T11:21:58.521408Z","shell.execute_reply.started":"2026-05-01T11:21:58.521281Z","shell.execute_reply":"2026-05-01T11:21:58.521303Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Step 2 — Install Grad-CAM library","metadata":{}},{"cell_type":"code","source":"!pip install grad-cam","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.522641Z","iopub.status.idle":"2026-05-01T11:21:58.522956Z","shell.execute_reply.started":"2026-05-01T11:21:58.522789Z","shell.execute_reply":"2026-05-01T11:21:58.522805Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Step 3 — Generate Grad-CAM visualization","metadata":{}},{"cell_type":"code","source":"# =========================================\n# STEP 3 — GRAD-CAM VISUALIZATION\n# =========================================\n\nimport cv2\nimport numpy as np\nimport matplotlib.pyplot as plt\n\nfrom pytorch_grad_cam import GradCAM\nfrom pytorch_grad_cam.utils.image import show_cam_on_image\n\n# Select last convolution layer\ntarget_layers = [model.features[-1]]\n\ncam = GradCAM(\n    model=model,\n    target_layers=target_layers\n)\n\n# Get one validation batch\nimages, labels = next(iter(val_loader))\n\nimage = images[0].unsqueeze(0).to(device)\n\n# Generate heatmap\ngrayscale_cam = cam(\n    input_tensor=image\n)[0]\n\n# Convert image for display\nimg = images[0].cpu().numpy().transpose(1, 2, 0)\n\nimg = (\n    img - img.min()\n) / (\n    img.max() - img.min()\n)\n\nvisualization = show_cam_on_image(\n    img,\n    grayscale_cam,\n    use_rgb=True\n)\n\nplt.figure(figsize=(6,6))\nplt.imshow(visualization)\nplt.axis(\"off\")\n\nprint(\"Grad-CAM visualization generated\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.524062Z","iopub.status.idle":"2026-05-01T11:21:58.524292Z","shell.execute_reply.started":"2026-05-01T11:21:58.524182Z","shell.execute_reply":"2026-05-01T11:21:58.524196Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Grad-CAM result means (brief interpretation)\n\nVisualization shows:\n\nStrong attention on optic disc and vessel region\nCentral retinal area highlighted\nNo random background focus\n\nThis indicates:\n\nThe model is learning clinically relevant features\nGrad-CAM working correctly\nModel attention is meaningful\n\nWe have completed:\n\nBaseline training ✅\nCBAM experiment ✅\nModel freezing ✅\nFirst Grad-CAM visualization ✅\n\nNow we move to the research-level Grad-CAM evaluation stage.","metadata":{}},{"cell_type":"markdown","source":"Next Step — Generate Multiple Grad-CAM Samples\n\nWe now produce:\n\n5 correct predictions\n5 wrong predictions\n\nGenerate 10 Grad-CAM images automatically","metadata":{}},{"cell_type":"code","source":"import torch\nimport numpy as np\nimport cv2\nimport matplotlib.pyplot as plt\n\n# -------------------------------\n# 1. global variables\n# -------------------------------\ngradients = None\nactivations = None\n\n# -------------------------------\n# 2. hook functions\n# -------------------------------\ndef save_gradient(module, grad_input, grad_output):\n    global gradients\n    gradients = grad_output[0]\n\ndef save_activation(module, input, output):\n    global activations\n    activations = output\n\n# -------------------------------\n# 3. choose correct target layer\n# -------------------------------\nif hasattr(model, \"backbone\"):\n    target_layer = model.backbone.features[-1]\nelse:\n    target_layer = model.features[-1]\n\n# remove old hooks if rerun\ntry:\n    forward_handle.remove()\n    backward_handle.remove()\nexcept:\n    pass\n\nforward_handle = target_layer.register_forward_hook(save_activation)\nbackward_handle = target_layer.register_full_backward_hook(save_gradient)\n\nprint(\"Grad-CAM hooks attached successfully\")\n\n# -------------------------------\n# 4. grad-cam function\n# -------------------------------\ndef generate_gradcam(model, image_tensor, device):\n    global gradients, activations\n\n    model.eval()\n    gradients = None\n    activations = None\n\n    image_tensor = image_tensor.unsqueeze(0).to(device)\n\n    output = model(image_tensor)\n    pred_class = torch.argmax(output, dim=1).item()\n\n    model.zero_grad()\n    output[0, pred_class].backward()\n\n    if gradients is None or activations is None:\n        raise ValueError(\"Gradients or activations not captured.\")\n\n    grads = gradients.detach().cpu()[0]\n    acts = activations.detach().cpu()[0]\n\n    weights = torch.mean(grads, dim=(1, 2))\n\n    cam = torch.zeros(acts.shape[1:], dtype=torch.float32)\n\n    for i, w in enumerate(weights):\n        cam += w * acts[i]\n\n    cam = torch.relu(cam)\n    cam = cam.numpy()\n\n    if cam.max() != 0:\n        cam = cam / cam.max()\n\n    return cam, pred_class\n\n# -------------------------------\n# 5. display function\n# -------------------------------\ndef show_gradcam_sample(image_tensor, true_label, pred_label, heatmap):\n    img = image_tensor.permute(1, 2, 0).cpu().numpy()\n\n    mean = np.array([0.485, 0.456, 0.406])\n    std = np.array([0.229, 0.224, 0.225])\n\n    img = (img * std) + mean\n    img = np.clip(img, 0, 1)\n\n    heatmap_resized = cv2.resize(heatmap, (img.shape[1], img.shape[0]))\n    heatmap_uint8 = np.uint8(255 * heatmap_resized)\n\n    heatmap_color = cv2.applyColorMap(heatmap_uint8, cv2.COLORMAP_JET)\n    heatmap_color = cv2.cvtColor(heatmap_color, cv2.COLOR_BGR2RGB)\n\n    overlay = np.uint8(0.4 * heatmap_color + 255 * img)\n\n    plt.figure(figsize=(15, 5))\n\n    plt.subplot(1, 3, 1)\n    plt.imshow(img)\n    plt.title(f\"Original\\nTrue: {true_label}\")\n    plt.axis(\"off\")\n\n    plt.subplot(1, 3, 2)\n    plt.imshow(heatmap_color)\n    plt.title(\"Grad-CAM Heatmap\")\n    plt.axis(\"off\")\n\n    plt.subplot(1, 3, 3)\n    plt.imshow(overlay)\n    plt.title(f\"Overlay\\nPred: {pred_label}\")\n    plt.axis(\"off\")\n\n    plt.show()\n\n# -------------------------------\n# 6. generate multiple samples\n# -------------------------------\ncorrect_count = 0\nwrong_count = 0\nmax_correct = 5\nmax_wrong = 5\n\nfor i in range(len(val_dataset)):\n    image_tensor, true_label = val_dataset[i]\n\n    try:\n        heatmap, pred_label = generate_gradcam(model, image_tensor, device)\n    except Exception as e:\n        print(f\"Skipping sample {i}: {e}\")\n        continue\n\n    if pred_label == true_label and correct_count < max_correct:\n        correct_count += 1\n        print(f\"\\nCorrect Sample {correct_count} | Index: {i}\")\n        show_gradcam_sample(image_tensor, true_label, pred_label, heatmap)\n\n    elif pred_label != true_label and wrong_count < max_wrong:\n        wrong_count += 1\n        print(f\"\\nWrong Sample {wrong_count} | Index: {i}\")\n        show_gradcam_sample(image_tensor, true_label, pred_label, heatmap)\n\n    if correct_count == max_correct and wrong_count == max_wrong:\n        break\n\nprint(\"\\nGrad-CAM visualization completed\")\nprint(\"Correct samples shown:\", correct_count)\nprint(\"Wrong samples shown:\", wrong_count)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.525086Z","iopub.status.idle":"2026-05-01T11:21:58.525353Z","shell.execute_reply.started":"2026-05-01T11:21:58.525233Z","shell.execute_reply":"2026-05-01T11:21:58.525248Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Grad-CAM results mean (clear interpretation)\n1) Correct Predictions — Model behavior is clinically meaningful\n\nFrom your first five correct samples:\n\nThe model focuses on:\nlesion clusters\nexudates / hemorrhage regions\noptic disc / vascular areas\n\nThis indicates:\n\nModel attention aligns with medically relevant retinal structures\nThat is exactly what reviewers expect in Explainable Medical AI.\n\n2) Wrong Predictions — Very important scientific evidence\n\nYour misclassified samples show:\n\nAttention sometimes shifts to:\nnearby bright regions\npartial lesion clusters\noptic disc dominance\n\nThis explains why the model made mistakes, not just that it made them.\n\nThis is strong research evidence because:\n\nYou can explain model errors using visual attention maps","metadata":{}},{"cell_type":"markdown","source":"We have successfully completed:\n\nDataset fusion (APTOS + IDRiD)        ✓\nBaseline model training                ✓\nWeighted loss + Early stopping         ✓\nCBAM experiment                        ✓\nModel comparison                       ✓\nGrad-CAM visualization                 ✓\nCorrect vs wrong prediction analysis   ✓","metadata":{}},{"cell_type":"markdown","source":"Next Step — Save (Freeze) Grad-CAM results","metadata":{}},{"cell_type":"code","source":"import os\nimport json\n\nGRADCAM_DIR = \"/kaggle/working/frozen_gradcam_results\"\nos.makedirs(GRADCAM_DIR, exist_ok=True)\n\nsummary = {\n    \"phase\": \"GradCAM Explainability\",\n    \"status\": \"Completed\",\n    \"correct_samples\": 5,\n    \"wrong_samples\": 5,\n    \"observation\": \"Model focuses on lesion regions in correct predictions and shows shifted attention in misclassified cases\",\n    \"conclusion\": \"Model decisions are explainable and clinically meaningful\"\n}\n\nwith open(\n    os.path.join(GRADCAM_DIR, \"gradcam_summary.json\"),\n    \"w\"\n) as f:\n    json.dump(summary, f, indent=4)\n\nprint(\"Grad-CAM results frozen successfully\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.526696Z","iopub.status.idle":"2026-05-01T11:21:58.527018Z","shell.execute_reply.started":"2026-05-01T11:21:58.526853Z","shell.execute_reply":"2026-05-01T11:21:58.526877Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Quantitative Grad-CAM validation using IDRiD lesion masks\n\nWe will move from:\n\nvisual explanation\n\nto:\n\nmeasurable explanation","metadata":{}},{"cell_type":"markdown","source":"Next metrics we should compute\n1. Pointing Game\n\nChecks whether the maximum Grad-CAM point lies inside the lesion mask.\n\n2. IoU\n\nMeasures overlap between:\n\nGrad-CAM region\ntrue lesion mask\n3. Energy-based score\n\nMeasures how much heatmap energy falls inside lesion regions.\nWe need folders/files for masks such as:\n\nMicroaneurysms\nHemorrhages\nHard Exudates\nSoft Exudates","metadata":{}},{"cell_type":"code","source":"import os\n\nfor root, dirs, files in os.walk(\"/kaggle/input\"):\n    if \"mask\" in root.lower() or \"segmentation\" in root.lower() or \"groundtruth\" in root.lower():\n        print(root, len(files))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.528184Z","iopub.status.idle":"2026-05-01T11:21:58.528446Z","shell.execute_reply.started":"2026-05-01T11:21:58.528317Z","shell.execute_reply":"2026-05-01T11:21:58.528337Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Next step\n\nYou need to add the IDRiD segmentation / lesion mask dataset in Kaggle.\n\nWe need folders/files for lesion masks like:\n\nMicroaneurysms\nHemorrhages\nHard Exudates\nSoft Exudates","metadata":{}},{"cell_type":"code","source":"import os\n\nfor root, dirs, files in os.walk(\"/kaggle/input\"):\n    if \"seg\" in root.lower() or \"mask\" in root.lower():\n        print(root)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.529391Z","iopub.status.idle":"2026-05-01T11:21:58.529747Z","shell.execute_reply.started":"2026-05-01T11:21:58.529556Z","shell.execute_reply":"2026-05-01T11:21:58.529578Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================================\n# FULL CODE — ENERGY-BASED SCORE FOR ALL LESION TYPES\n# =========================================\n\nimport os\nimport numpy as np\nimport cv2\nimport pandas as pd\n\n# -------------------------------------------------\n# 1. LESION MASK DIRECTORIES\n# -------------------------------------------------\nSEG_BASE = \"/kaggle/input/datasets/dankok/diabetic-retinopathy-image-dataset/Segmentation/Segmentation_Groundtruths/Training Set\"\n\nLESION_DIRS = {\n    \"Microaneurysms\": os.path.join(SEG_BASE, \"Microaneurysms\"),\n    \"Haemorrhages\": os.path.join(SEG_BASE, \"Haemorrhages\"),\n    \"Hard Exudates\": os.path.join(SEG_BASE, \"Hard Exudates\"),\n    \"Soft Exudates\": os.path.join(SEG_BASE, \"Soft Exudates\"),\n}\n\nLESION_SUFFIX = {\n    \"Microaneurysms\": \"MA\",\n    \"Haemorrhages\": \"HE\",\n    \"Hard Exudates\": \"EX\",\n    \"Soft Exudates\": \"SE\",\n}\n\n# -------------------------------------------------\n# 2. FILTER ONLY IDRID VALIDATION IMAGES\n# -------------------------------------------------\nidrid_val_df = val_dataset.df[val_dataset.df[\"source\"] == \"idrid\"].reset_index(drop=True)\n\nprint(\"Total IDRiD validation images:\", len(idrid_val_df))\n\n# -------------------------------------------------\n# 3. FUNCTION TO COMPUTE ENERGY SCORE\n# -------------------------------------------------\ndef compute_energy_score_for_lesion(model, val_dataset, idrid_val_df, mask_dir, suffix, device):\n    \n    matched_files = 0\n    energy_scores = []\n\n    for i in range(len(idrid_val_df)):\n        row = idrid_val_df.iloc[i]\n        image_path = row[\"image_path\"]\n\n        # Example: IDRiD_003.jpg\n        image_name = os.path.basename(image_path)\n        image_stem = os.path.splitext(image_name)[0]   # IDRiD_003\n\n        # Convert IDRiD_003 -> 3\n        num = int(image_stem.split(\"_\")[1])\n\n        # Example mask name: IDRiD_3_EX.tif\n        mask_name = f\"IDRiD_{num}_{suffix}.tif\"\n        mask_path = os.path.join(mask_dir, mask_name)\n\n        if not os.path.exists(mask_path):\n            continue\n\n        # Find matching full index in val_dataset\n        full_idx = val_dataset.df[val_dataset.df[\"image_path\"] == image_path].index[0]\n        image_tensor, label = val_dataset[full_idx]\n\n        # Load lesion mask\n        mask = cv2.imread(mask_path, 0)\n        if mask is None:\n            continue\n\n        # Generate Grad-CAM\n        heatmap, pred = generate_gradcam(model, image_tensor, device)\n\n        # Resize heatmap to mask shape\n        heatmap_resized = cv2.resize(\n            heatmap,\n            (mask.shape[1], mask.shape[0])\n        )\n\n        # Binary lesion mask\n        mask_binary = (mask > 0).astype(np.float32)\n\n        # Energy calculation\n        total_energy = np.sum(heatmap_resized)\n        lesion_energy = np.sum(heatmap_resized * mask_binary)\n\n        if total_energy > 0:\n            score = lesion_energy / total_energy\n            energy_scores.append(score)\n            matched_files += 1\n\n    if matched_files > 0:\n        mean_score = float(np.mean(energy_scores))\n        return mean_score, matched_files, energy_scores\n    else:\n        return None, 0, []\n\n# -------------------------------------------------\n# 4. RUN FOR ALL LESION TYPES\n# -------------------------------------------------\nresults = []\n\nfor lesion_name in LESION_DIRS.keys():\n    print(f\"\\nProcessing: {lesion_name}\")\n\n    mask_dir = LESION_DIRS[lesion_name]\n    suffix = LESION_SUFFIX[lesion_name]\n\n    mean_score, matched_files, scores = compute_energy_score_for_lesion(\n        model=model,\n        val_dataset=val_dataset,\n        idrid_val_df=idrid_val_df,\n        mask_dir=mask_dir,\n        suffix=suffix,\n        device=device\n    )\n\n    if matched_files > 0:\n        print(f\"Matched samples: {matched_files}\")\n        print(f\"Mean Energy-based Score: {mean_score:.6f}\")\n    else:\n        print(\"No matched samples found\")\n\n    results.append({\n        \"Lesion Type\": lesion_name,\n        \"Matched Samples\": matched_files,\n        \"Mean Energy Score\": mean_score\n    })\n\n# -------------------------------------------------\n# 5. RESULTS TABLE\n# -------------------------------------------------\nresults_df = pd.DataFrame(results)\n\nprint(\"\\nFinal Energy-based Score Table:\")\nprint(results_df)\n\n# -------------------------------------------------\n# 6. SAVE RESULTS\n# -------------------------------------------------\nsave_path = \"/kaggle/working/energy_scores_all_lesions.csv\"\nresults_df.to_csv(save_path, index=False)\n\nprint(\"\\nResults saved at:\")\nprint(save_path)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.530634Z","iopub.status.idle":"2026-05-01T11:21:58.530999Z","shell.execute_reply.started":"2026-05-01T11:21:58.530804Z","shell.execute_reply":"2026-05-01T11:21:58.530826Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Grad-CAM energy inside true lesion masks is:\n\nMicroaneurysms: 0.0022\nHaemorrhages: 0.0177\nHard Exudates: 0.0107\nSoft Exudates: 0.0054\nSimple interpretation\n\nThese values are low.\n\nThis means:\n\nvisually, Grad-CAM looked meaningful\nbut quantitatively, the model’s attention is not strongly concentrated on exact lesion pixels\nthe model is likely using broader retinal context instead of sharply focusing on lesion masks","metadata":{}},{"cell_type":"markdown","source":"This single script will:\n\nLoad fused dataset (APTOS + IDRiD)\nDetect which images have lesion masks\nGenerate Grad-CAM inside training loop\nCompute alignment loss with mask\nTrain model using combined loss\nSave best model\nImprove explainability metric","metadata":{}},{"cell_type":"code","source":"# =========================================================\n# FULL CORRECT CODE — LESION-GUIDED TRAINING\n# =========================================================\n\nimport os\nimport cv2\nimport torch\nimport random\nimport numpy as np\nimport pandas as pd\nimport torch.nn as nn\nimport torch.optim as optim\n\nfrom sklearn.metrics import accuracy_score, f1_score, cohen_kappa_score\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nfrom torchvision.models import efficientnet_b5, EfficientNet_B5_Weights\n\n# =========================================================\n# 1. DEVICE + SEED\n# =========================================================\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Using device:\", device)\n\ndef set_seed(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n\nset_seed(42)\n\n# =========================================================\n# 2. PATHS\n# =========================================================\nMODEL_DIR = \"/kaggle/input/models/afnanhalim/model-work-continued/pytorch/default/1\"\n\nCOMBINED_CSV = os.path.join(MODEL_DIR, \"combined_df.csv\")\nCHECKPOINT_PATH = os.path.join(MODEL_DIR, \"weighted_earlystop_checkpoint_final.pth\")\n\nSEG_BASE = \"/kaggle/input/datasets/dankok/diabetic-retinopathy-image-dataset/Segmentation/Segmentation_Groundtruths/Training Set\"\n\nLESION_DIRS = {\n    \"EX\": os.path.join(SEG_BASE, \"Hard Exudates\"),\n    \"HE\": os.path.join(SEG_BASE, \"Haemorrhages\"),\n    \"MA\": os.path.join(SEG_BASE, \"Microaneurysms\"),\n    \"SE\": os.path.join(SEG_BASE, \"Soft Exudates\"),\n}\n\nSAVE_DIR = \"/kaggle/working/lesion_guided_training\"\nos.makedirs(SAVE_DIR, exist_ok=True)\n\n# =========================================================\n# 3. LOAD COMBINED DATA\n# =========================================================\ndf = pd.read_csv(COMBINED_CSV)\n\ntrain_df = df[df[\"fold\"] != 0].reset_index(drop=True)\nval_df   = df[df[\"fold\"] == 0].reset_index(drop=True)\n\nprint(\"Train shape:\", train_df.shape)\nprint(\"Val shape:\", val_df.shape)\n\n# =========================================================\n# 4. PREPROCESSING\n# =========================================================\nIMG_SIZE = 456\n\ndef circular_crop(img):\n    gray = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)\n    _, thresh = cv2.threshold(gray, 10, 255, cv2.THRESH_BINARY)\n    contours, _ = cv2.findContours(thresh, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n\n    if len(contours) == 0:\n        return img\n\n    cnt = max(contours, key=cv2.contourArea)\n    x, y, w, h = cv2.boundingRect(cnt)\n    return img[y:y+h, x:x+w]\n\ndef apply_clahe(img):\n    lab = cv2.cvtColor(img, cv2.COLOR_BGR2LAB)\n    l, a, b = cv2.split(lab)\n\n    clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))\n    l = clahe.apply(l)\n\n    lab = cv2.merge((l, a, b))\n    img = cv2.cvtColor(lab, cv2.COLOR_LAB2BGR)\n    return img\n\n# =========================================================\n# 5. TRANSFORMS\n# =========================================================\ntrain_transform = transforms.Compose([\n    transforms.ToPILImage(),\n    transforms.RandomHorizontalFlip(),\n    transforms.RandomVerticalFlip(),\n    transforms.RandomRotation(20),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225]\n    )\n])\n\nval_transform = transforms.Compose([\n    transforms.ToPILImage(),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225]\n    )\n])\n\n# =========================================================\n# 6. DATASET\n# =========================================================\nclass LesionGuidedDataset(Dataset):\n\n    def __init__(self, df, transform=None):\n        self.df = df.reset_index(drop=True)\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def load_union_mask(self, image_name):\n        \"\"\"\n        For IDRiD images, load all available lesion masks and merge them into one union mask.\n        For APTOS images, returns None.\n        \"\"\"\n        stem = os.path.splitext(image_name)[0]\n\n        if not stem.startswith(\"IDRiD_\"):\n            return None\n\n        try:\n            num = int(stem.split(\"_\")[1])\n        except:\n            return None\n\n        union_mask = np.zeros((IMG_SIZE, IMG_SIZE), dtype=np.float32)\n        found = False\n\n        for suffix, folder in LESION_DIRS.items():\n            mask_name = f\"IDRiD_{num}_{suffix}.tif\"\n            mask_path = os.path.join(folder, mask_name)\n\n            if os.path.exists(mask_path):\n                mask = cv2.imread(mask_path, 0)\n                if mask is not None:\n                    mask = cv2.resize(mask, (IMG_SIZE, IMG_SIZE))\n                    mask = (mask > 0).astype(np.float32)\n                    union_mask = np.maximum(union_mask, mask)\n                    found = True\n\n        if found:\n            return union_mask\n        return None\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n\n        img_path = row[\"image_path\"]\n        label = int(row[\"label\"])\n\n        img = cv2.imread(img_path)\n        img = circular_crop(img)\n        img = apply_clahe(img)\n        img = cv2.resize(img, (IMG_SIZE, IMG_SIZE))\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n\n        image_name = os.path.basename(img_path)\n        mask = self.load_union_mask(image_name)\n\n        if self.transform:\n            img = self.transform(img)\n\n        if mask is None:\n            mask = torch.zeros((1, IMG_SIZE, IMG_SIZE), dtype=torch.float32)\n            has_mask = torch.tensor(0, dtype=torch.uint8)\n        else:\n            mask = torch.tensor(mask, dtype=torch.float32).unsqueeze(0)\n            has_mask = torch.tensor(1, dtype=torch.uint8)\n\n        return img, label, mask, has_mask\n\ntrain_dataset = LesionGuidedDataset(train_df, transform=train_transform)\nval_dataset   = LesionGuidedDataset(val_df, transform=val_transform)\n\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=8,\n    shuffle=True,\n    num_workers=2,\n    pin_memory=True\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=8,\n    shuffle=False,\n    num_workers=2,\n    pin_memory=True\n)\n\nprint(\"Train batches:\", len(train_loader))\nprint(\"Val batches:\", len(val_loader))\n\n# =========================================================\n# 7. MODEL\n# =========================================================\nweights = EfficientNet_B5_Weights.IMAGENET1K_V1\nmodel = efficientnet_b5(weights=weights)\n\nin_features = model.classifier[1].in_features\nmodel.classifier[1] = nn.Linear(in_features, 5)\nmodel = model.to(device)\n\ncheckpoint = torch.load(CHECKPOINT_PATH, map_location=device)\nmodel.load_state_dict(checkpoint[\"model_state_dict\"])\n\nprint(\"Frozen best weighted model loaded successfully\")\n\n# =========================================================\n# 8. CLASSIFICATION LOSS\n# =========================================================\ncriterion_cls = nn.CrossEntropyLoss()\n\n# =========================================================\n# 9. GRAD-CAM HOOKS\n# =========================================================\ngradients = None\nactivations = None\n\ndef save_gradient(module, grad_input, grad_output):\n    global gradients\n    gradients = grad_output[0]\n\ndef save_activation(module, input, output):\n    global activations\n    activations = output\n\ntarget_layer = model.features[-1]\nforward_handle = target_layer.register_forward_hook(save_activation)\nbackward_handle = target_layer.register_full_backward_hook(save_gradient)\n\n# =========================================================\n# 10. ATTENTION ALIGNMENT LOSS\n# =========================================================\ndef generate_cam_for_sample(class_score):\n    global gradients, activations\n\n    grads = gradients[0]      # [C,H,W]\n    acts = activations[0]     # [C,H,W]\n\n    weights = torch.mean(grads, dim=(1, 2))\n    cam = torch.zeros_like(acts[0])\n\n    for i, w in enumerate(weights):\n        cam += w * acts[i]\n\n    cam = torch.relu(cam)\n    cam = cam / (cam.max() + 1e-8)\n    return cam\n\ndef attention_alignment_loss(cam, mask):\n    cam = cam.unsqueeze(0).unsqueeze(0)      # [1,1,h,w]\n    cam = torch.nn.functional.interpolate(\n        cam,\n        size=(IMG_SIZE, IMG_SIZE),\n        mode=\"bilinear\",\n        align_corners=False\n    )\n    cam = cam.squeeze()\n    mask = mask.squeeze()\n    return torch.mean((cam - mask) ** 2)\n\n# =========================================================\n# 11. OPTIMIZER\n# =========================================================\noptimizer = optim.AdamW(\n    model.parameters(),\n    lr=1e-5,\n    weight_decay=1e-4\n)\n\nlambda_align = 0.1   # small and safe starting value\n\n# =========================================================\n# 12. TRAIN / VALID FUNCTIONS\n# =========================================================\ndef train_one_epoch(model, loader, optimizer, device):\n    model.train()\n\n    total_loss_value = 0.0\n    all_preds = []\n    all_labels = []\n\n    for images, labels, masks, has_masks in loader:\n        images = images.to(device, non_blocking=True)\n        labels = labels.to(device, non_blocking=True)\n        masks = masks.to(device, non_blocking=True)\n        has_masks = has_masks.to(device, non_blocking=True)\n\n        optimizer.zero_grad()\n\n        outputs = model(images)\n        loss_cls = criterion_cls(outputs, labels)\n\n        loss_align = 0.0\n        valid_mask_count = 0\n\n        for i in range(images.size(0)):\n            if has_masks[i].item() == 1:\n                model.zero_grad()\n                outputs[i, labels[i]].backward(retain_graph=True)\n\n                cam = generate_cam_for_sample(outputs[i, labels[i]])\n                loss_align += attention_alignment_loss(cam, masks[i])\n                valid_mask_count += 1\n\n        if valid_mask_count > 0:\n            loss_align = loss_align / valid_mask_count\n        else:\n            loss_align = torch.tensor(0.0, device=device)\n\n        loss = loss_cls + lambda_align * loss_align\n\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n\n        preds = torch.argmax(outputs, dim=1)\n\n        all_preds.extend(preds.detach().cpu().numpy())\n        all_labels.extend(labels.detach().cpu().numpy())\n\n        total_loss_value += loss.item()\n\n    acc = accuracy_score(all_labels, all_preds)\n    f1 = f1_score(all_labels, all_preds, average=\"weighted\")\n    qwk = cohen_kappa_score(all_labels, all_preds, weights=\"quadratic\")\n\n    return total_loss_value, acc, f1, qwk\n\ndef validate_one_epoch(model, loader, device):\n    model.eval()\n\n    all_preds = []\n    all_labels = []\n\n    with torch.no_grad():\n        for images, labels, masks, has_masks in loader:\n            images = images.to(device, non_blocking=True)\n            labels = labels.to(device, non_blocking=True)\n\n            outputs = model(images)\n            preds = torch.argmax(outputs, dim=1)\n\n            all_preds.extend(preds.detach().cpu().numpy())\n            all_labels.extend(labels.detach().cpu().numpy())\n\n    acc = accuracy_score(all_labels, all_preds)\n    f1 = f1_score(all_labels, all_preds, average=\"weighted\")\n    qwk = cohen_kappa_score(all_labels, all_preds, weights=\"quadratic\")\n\n    return acc, f1, qwk\n\n# =========================================================\n# 13. TRAINING LOOP\n# =========================================================\nnum_epochs = 5\nbest_val_qwk = 0.0\n\nprint(\"\\nStarting lesion-guided training...\\n\")\n\nfor epoch in range(num_epochs):\n    train_loss, train_acc, train_f1, train_qwk = train_one_epoch(\n        model, train_loader, optimizer, device\n    )\n\n    val_acc, val_f1, val_qwk = validate_one_epoch(\n        model, val_loader, device\n    )\n\n    print(f\"Epoch {epoch+1}/{num_epochs}\")\n    print(f\"Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.4f} | Train F1: {train_f1:.4f} | Train QWK: {train_qwk:.4f}\")\n    print(f\"Val   Acc: {val_acc:.4f} | Val   F1: {val_f1:.4f} | Val   QWK: {val_qwk:.4f}\")\n\n    if val_qwk > best_val_qwk:\n        best_val_qwk = val_qwk\n\n        torch.save(\n            model.state_dict(),\n            os.path.join(SAVE_DIR, \"lesion_guided_best_model.pth\")\n        )\n\n        print(\"Saved best lesion-guided model\")\n\n    print(\"-\" * 90)\n\nprint(\"Training finished\")\nprint(\"Best Val QWK:\", best_val_qwk)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.532092Z","iopub.status.idle":"2026-05-01T11:21:58.532324Z","shell.execute_reply.started":"2026-05-01T11:21:58.532216Z","shell.execute_reply":"2026-05-01T11:21:58.532230Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================================================\n# EVALUATE LESION-GUIDED MODEL WITH ENERGY SCORES\n# =========================================================\n\nimport os\nimport cv2\nimport torch\nimport numpy as np\nimport pandas as pd\nimport torch.nn as nn\nfrom torchvision.models import efficientnet_b5, EfficientNet_B5_Weights\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# -------------------------------------------------\n# PATHS\n# -------------------------------------------------\nSAVE_DIR = \"/kaggle/working/lesion_guided_training\"\nBEST_MODEL_PATH = os.path.join(SAVE_DIR, \"lesion_guided_best_model.pth\")\n\nSEG_BASE = \"/kaggle/input/datasets/dankok/diabetic-retinopathy-image-dataset/Segmentation/Segmentation_Groundtruths/Training Set\"\n\nLESION_DIRS = {\n    \"Microaneurysms\": os.path.join(SEG_BASE, \"Microaneurysms\"),\n    \"Haemorrhages\": os.path.join(SEG_BASE, \"Haemorrhages\"),\n    \"Hard Exudates\": os.path.join(SEG_BASE, \"Hard Exudates\"),\n    \"Soft Exudates\": os.path.join(SEG_BASE, \"Soft Exudates\"),\n}\n\nLESION_SUFFIX = {\n    \"Microaneurysms\": \"MA\",\n    \"Haemorrhages\": \"HE\",\n    \"Hard Exudates\": \"EX\",\n    \"Soft Exudates\": \"SE\",\n}\n\n# -------------------------------------------------\n# LOAD BEST LESION-GUIDED MODEL\n# -------------------------------------------------\nweights = EfficientNet_B5_Weights.IMAGENET1K_V1\nbest_model = efficientnet_b5(weights=weights)\n\nin_features = best_model.classifier[1].in_features\nbest_model.classifier[1] = nn.Linear(in_features, 5)\n\nbest_model = best_model.to(device)\nbest_model.load_state_dict(torch.load(BEST_MODEL_PATH, map_location=device))\nbest_model.eval()\n\nprint(\"Lesion-guided best model loaded successfully\")\n\n# -------------------------------------------------\n# GRAD-CAM HOOKS\n# -------------------------------------------------\ngradients = None\nactivations = None\n\ndef save_gradient(module, grad_input, grad_output):\n    global gradients\n    gradients = grad_output[0]\n\ndef save_activation(module, input, output):\n    global activations\n    activations = output\n\ntarget_layer = best_model.features[-1]\nforward_handle = target_layer.register_forward_hook(save_activation)\nbackward_handle = target_layer.register_full_backward_hook(save_gradient)\n\ndef generate_gradcam(model, image_tensor, device):\n    global gradients, activations\n\n    model.eval()\n    gradients = None\n    activations = None\n\n    image_tensor = image_tensor.unsqueeze(0).to(device)\n\n    output = model(image_tensor)\n    pred_class = torch.argmax(output, dim=1).item()\n\n    model.zero_grad()\n    output[0, pred_class].backward()\n\n    if gradients is None or activations is None:\n        raise ValueError(\"Gradients or activations not captured.\")\n\n    grads = gradients.detach().cpu()[0]\n    acts = activations.detach().cpu()[0]\n\n    weights_local = torch.mean(grads, dim=(1, 2))\n    cam = torch.zeros(acts.shape[1:], dtype=torch.float32)\n\n    for i, w in enumerate(weights_local):\n        cam += w * acts[i]\n\n    cam = torch.relu(cam)\n    cam = cam.numpy()\n\n    if cam.max() != 0:\n        cam = cam / cam.max()\n\n    return cam, pred_class\n\n# -------------------------------------------------\n# COMPUTE ENERGY SCORES FOR ALL LESIONS\n# -------------------------------------------------\nidrid_val_df = val_dataset.df[val_dataset.df[\"source\"] == \"idrid\"].reset_index(drop=True)\n\nprint(\"Total IDRiD validation images:\", len(idrid_val_df))\n\nresults = []\n\nfor lesion_name in LESION_DIRS.keys():\n    print(f\"\\nProcessing: {lesion_name}\")\n\n    mask_dir = LESION_DIRS[lesion_name]\n    suffix = LESION_SUFFIX[lesion_name]\n\n    matched_files = 0\n    energy_scores = []\n\n    for i in range(len(idrid_val_df)):\n        row = idrid_val_df.iloc[i]\n        image_path = row[\"image_path\"]\n\n        image_name = os.path.basename(image_path)         # IDRiD_003.jpg\n        image_stem = os.path.splitext(image_name)[0]      # IDRiD_003\n\n        num = int(image_stem.split(\"_\")[1])               # 3\n        mask_name = f\"IDRiD_{num}_{suffix}.tif\"           # IDRiD_3_EX.tif\n        mask_path = os.path.join(mask_dir, mask_name)\n\n        if not os.path.exists(mask_path):\n            continue\n\n        full_idx = val_dataset.df[val_dataset.df[\"image_path\"] == image_path].index[0]\n        image_tensor, label, mask_dummy, has_mask = val_dataset[full_idx]\n\n        mask = cv2.imread(mask_path, 0)\n        if mask is None:\n            continue\n\n        heatmap, pred = generate_gradcam(best_model, image_tensor, device)\n\n        heatmap_resized = cv2.resize(\n            heatmap,\n            (mask.shape[1], mask.shape[0])\n        )\n\n        mask_binary = (mask > 0).astype(np.float32)\n\n        total_energy = np.sum(heatmap_resized)\n        lesion_energy = np.sum(heatmap_resized * mask_binary)\n\n        if total_energy > 0:\n            score = lesion_energy / total_energy\n            energy_scores.append(score)\n            matched_files += 1\n\n    if matched_files > 0:\n        mean_score = float(np.mean(energy_scores))\n        print(\"Matched samples:\", matched_files)\n        print(\"Mean Energy-based Score:\", mean_score)\n    else:\n        mean_score = None\n        print(\"No matched samples found\")\n\n    results.append({\n        \"Lesion Type\": lesion_name,\n        \"Matched Samples\": matched_files,\n        \"Mean Energy Score (Lesion-Guided)\": mean_score\n    })\n\nresults_df = pd.DataFrame(results)\n\nprint(\"\\nFinal Energy Score Table:\")\nprint(results_df)\n\nresults_df.to_csv(\n    os.path.join(SAVE_DIR, \"lesion_guided_energy_scores.csv\"),\n    index=False\n)\n\nprint(\"\\nSaved at:\", os.path.join(SAVE_DIR, \"lesion_guided_energy_scores.csv\"))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.533394Z","iopub.status.idle":"2026-05-01T11:21:58.533763Z","shell.execute_reply.started":"2026-05-01T11:21:58.533568Z","shell.execute_reply":"2026-05-01T11:21:58.533607Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torchvision.models import efficientnet_b5, EfficientNet_B5_Weights\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Using device:\", device)\n\nMODEL_DIR = \"/kaggle/input/models/afnanhalim/frozen-model/pytorch/default/1/frozen_best_weighted_model\"\n\nTRAIN_DF_PATH = os.path.join(MODEL_DIR, \"train_df_frozen.csv\")\nVAL_DF_PATH = os.path.join(MODEL_DIR, \"val_df_frozen.csv\")\nCLASS_WEIGHTS_PATH = os.path.join(MODEL_DIR, \"class_weights.pt\")\nCHECKPOINT_PATH = os.path.join(MODEL_DIR, \"weighted_earlystop_checkpoint_final.pth\")\n\ntrain_df = pd.read_csv(TRAIN_DF_PATH)\nval_df = pd.read_csv(VAL_DF_PATH)\n\nprint(\"Train shape:\", train_df.shape)\nprint(\"Val shape:\", val_df.shape)\n\nweights = EfficientNet_B5_Weights.IMAGENET1K_V1\nmodel = efficientnet_b5(weights=weights)\n\nin_features = model.classifier[1].in_features\nmodel.classifier[1] = nn.Linear(in_features, 5)\nmodel = model.to(device)\n\noptimizer = optim.AdamW(model.parameters(), lr=1e-4, weight_decay=1e-4)\n\ncheckpoint = torch.load(CHECKPOINT_PATH, map_location=device)\nmodel.load_state_dict(checkpoint[\"model_state_dict\"])\noptimizer.load_state_dict(checkpoint[\"optimizer_state_dict\"])\n\nbest_qwk = checkpoint.get(\"best_qwk\", None)\n\nclass_weights = torch.load(CLASS_WEIGHTS_PATH, map_location=device)\ncriterion = nn.CrossEntropyLoss(weight=class_weights)\n\nprint(\"Checkpoint loaded successfully\")\nprint(\"Best QWK:\", best_qwk)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.535147Z","iopub.status.idle":"2026-05-01T11:21:58.535456Z","shell.execute_reply.started":"2026-05-01T11:21:58.535316Z","shell.execute_reply":"2026-05-01T11:21:58.535339Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================================\n# DEFINE GRAD-CAM FUNCTION (FIX)\n# =========================================\n\nimport torch\nimport torch.nn.functional as F\nimport numpy as np\nimport cv2\n\n# -----------------------------------------\n# Select correct target layer\n# EfficientNet last feature layer\n# -----------------------------------------\n\ntarget_layer = model.features[-1]\n\ngradients = None\nactivations = None\n\n\ndef forward_hook(module, input, output):\n    global activations\n    activations = output\n\n\ndef backward_hook(module, grad_input, grad_output):\n    global gradients\n    gradients = grad_output[0]\n\n\n# Register hooks\ntarget_layer.register_forward_hook(forward_hook)\ntarget_layer.register_backward_hook(backward_hook)\n\nprint(\"Grad-CAM hooks attached successfully\")\n\n\n# =========================================\n# GRAD-CAM FUNCTION\n# =========================================\n\ndef generate_gradcam(model, image_tensor, device):\n\n    model.eval()\n\n    image_tensor = image_tensor.unsqueeze(0).to(device)\n\n    output = model(image_tensor)\n\n    pred_class = output.argmax(dim=1)\n\n    model.zero_grad()\n\n    output[0, pred_class].backward()\n\n    global gradients\n    global activations\n\n    pooled_gradients = torch.mean(\n        gradients,\n        dim=[0, 2, 3]\n    )\n\n    activations_copy = activations.detach().clone()\n\n    for i in range(pooled_gradients.shape[0]):\n        activations_copy[:, i, :, :] *= pooled_gradients[i]\n\n    heatmap = torch.mean(\n        activations_copy,\n        dim=1\n    ).squeeze()\n\n    heatmap = F.relu(heatmap)\n\n    heatmap /= torch.max(heatmap) + 1e-8\n\n    heatmap = heatmap.cpu().numpy()\n\n    heatmap = cv2.resize(\n        heatmap,\n        (456, 456)\n    )\n\n    return heatmap, int(pred_class.item())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.536328Z","iopub.status.idle":"2026-05-01T11:21:58.536639Z","shell.execute_reply.started":"2026-05-01T11:21:58.536490Z","shell.execute_reply":"2026-05-01T11:21:58.536510Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"STRONG LESION IMPROVEMENT TRAINING","metadata":{}},{"cell_type":"code","source":"# ============================================================\n# CORRECT PATH + RESTORE + IMPROVEMENT TRAINING\n# ============================================================\n\nimport os\nimport cv2\nimport pandas as pd\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\n\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.metrics import cohen_kappa_score\nimport torchvision.transforms as transforms\n\n# ------------------------------------------------------------\n# DEVICE\n# ------------------------------------------------------------\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Using device:\", device)\n\n# ------------------------------------------------------------\n# CORRECT PATH\n# ------------------------------------------------------------\nBASE_PATH = \"/kaggle/input/models/afnanhalim/frozen-model/pytorch/default/1/frozen_best_weighted_model\"\nprint(\"Base path:\", BASE_PATH)\n\n# ------------------------------------------------------------\n# LOAD FILES\n# ------------------------------------------------------------\ntrain_df = pd.read_csv(os.path.join(BASE_PATH, \"train_df_frozen.csv\"))\nval_df   = pd.read_csv(os.path.join(BASE_PATH, \"val_df_frozen.csv\"))\n\nprint(\"Train shape:\", train_df.shape)\nprint(\"Val shape:\", val_df.shape)\n\nclass_weights = torch.load(\n    os.path.join(BASE_PATH, \"class_weights.pt\"),\n    map_location=device\n)\n\nprint(\"Class weights loaded:\", class_weights)\n\ncheckpoint = torch.load(\n    os.path.join(BASE_PATH, \"weighted_earlystop_checkpoint_final.pth\"),\n    map_location=device\n)\n\nmodel.load_state_dict(checkpoint[\"model_state_dict\"])\n\nprint(\"Checkpoint loaded successfully\")\nprint(\"Best QWK:\", checkpoint.get(\"best_qwk\", \"Not found\"))\n\n# ------------------------------------------------------------\n# IMAGE PREPROCESS\n# ------------------------------------------------------------\nIMG_SIZE = 456\n\ndef apply_clahe(img):\n    lab = cv2.cvtColor(img, cv2.COLOR_BGR2LAB)\n    l, a, b = cv2.split(lab)\n\n    clahe = cv2.createCLAHE(\n        clipLimit=2.0,\n        tileGridSize=(8, 8)\n    )\n    l = clahe.apply(l)\n\n    lab = cv2.merge((l, a, b))\n    img = cv2.cvtColor(lab, cv2.COLOR_LAB2BGR)\n\n    return img\n\ntrain_transform = transforms.Compose([\n    transforms.ToPILImage(),\n    transforms.RandomHorizontalFlip(),\n    transforms.RandomVerticalFlip(),\n    transforms.RandomRotation(15),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225]\n    )\n])\n\nval_transform = transforms.Compose([\n    transforms.ToPILImage(),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225]\n    )\n])\n\n# ------------------------------------------------------------\n# DATASET\n# ------------------------------------------------------------\nclass DRDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df.reset_index(drop=True)\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n\n        image_path = row[\"image_path\"]\n        label = int(row[\"label\"])\n\n        img = cv2.imread(image_path)\n        if img is None:\n            raise ValueError(f\"Image not found: {image_path}\")\n\n        img = apply_clahe(img)\n        img = cv2.resize(img, (IMG_SIZE, IMG_SIZE))\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n\n        if self.transform:\n            img = self.transform(img)\n\n        return img, label\n\n# ------------------------------------------------------------\n# DATALOADERS\n# ------------------------------------------------------------\ntrain_dataset = DRDataset(train_df, transform=train_transform)\nval_dataset   = DRDataset(val_df, transform=val_transform)\n\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=8,\n    shuffle=True,\n    num_workers=2,\n    pin_memory=True\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=8,\n    shuffle=False,\n    num_workers=2,\n    pin_memory=True\n)\n\nprint(\"Train batches:\", len(train_loader))\nprint(\"Val batches:\", len(val_loader))\n\n# ------------------------------------------------------------\n# FOCAL LOSS\n# ------------------------------------------------------------\nclass FocalLoss(nn.Module):\n    def __init__(self, gamma=2, weight=None):\n        super().__init__()\n        self.gamma = gamma\n        self.ce = nn.CrossEntropyLoss(weight=weight)\n\n    def forward(self, inputs, targets):\n        ce_loss = self.ce(inputs, targets)\n        pt = torch.exp(-ce_loss)\n        loss = ((1 - pt) ** self.gamma) * ce_loss\n        return loss\n\ncriterion = FocalLoss(\n    gamma=2,\n    weight=class_weights\n)\n\n# ------------------------------------------------------------\n# OPTIMIZER\n# ------------------------------------------------------------\noptimizer = optim.AdamW(\n    model.parameters(),\n    lr=1e-4,\n    weight_decay=1e-4\n)\n\n# ------------------------------------------------------------\n# TRAIN FUNCTION\n# ------------------------------------------------------------\ndef train_one_epoch():\n    model.train()\n\n    total_loss = 0.0\n    all_preds = []\n    all_labels = []\n\n    for images, labels in train_loader:\n        images = images.to(device)\n        labels = labels.to(device)\n\n        optimizer.zero_grad()\n\n        outputs = model(images)\n        loss = criterion(outputs, labels)\n\n        loss.backward()\n        optimizer.step()\n\n        total_loss += loss.item()\n\n        preds = torch.argmax(outputs, dim=1)\n        all_preds.extend(preds.detach().cpu().numpy())\n        all_labels.extend(labels.detach().cpu().numpy())\n\n    qwk = cohen_kappa_score(\n        all_labels,\n        all_preds,\n        weights=\"quadratic\"\n    )\n\n    return total_loss, qwk\n\n# ------------------------------------------------------------\n# VALIDATE FUNCTION\n# ------------------------------------------------------------\ndef validate():\n    model.eval()\n\n    all_preds = []\n    all_labels = []\n\n    with torch.no_grad():\n        for images, labels in val_loader:\n            images = images.to(device)\n            labels = labels.to(device)\n\n            outputs = model(images)\n            preds = torch.argmax(outputs, dim=1)\n\n            all_preds.extend(preds.cpu().numpy())\n            all_labels.extend(labels.cpu().numpy())\n\n    qwk = cohen_kappa_score(\n        all_labels,\n        all_preds,\n        weights=\"quadratic\"\n    )\n\n    return qwk\n\n# ------------------------------------------------------------\n# TRAIN LOOP\n# ------------------------------------------------------------\nnum_epochs = 3\nbest_qwk = 0.0\n\nfor epoch in range(num_epochs):\n    train_loss, train_qwk = train_one_epoch()\n    val_qwk = validate()\n\n    print(f\"\\nEpoch {epoch+1}/{num_epochs}\")\n    print(\"Train loss:\", train_loss)\n    print(\"Train QWK:\", train_qwk)\n    print(\"Val QWK:\", val_qwk)\n\n    if val_qwk > best_qwk:\n        best_qwk = val_qwk\n        torch.save(\n            model.state_dict(),\n            \"/kaggle/working/improved_lesion_model.pth\"\n        )\n        print(\"Saved improved model\")\n\nprint(\"\\nTraining finished\")\nprint(\"Best QWK:\", best_qwk)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.538103Z","iopub.status.idle":"2026-05-01T11:21:58.538819Z","shell.execute_reply.started":"2026-05-01T11:21:58.538626Z","shell.execute_reply":"2026-05-01T11:21:58.538650Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"restores your frozen best model and its train/validation split, rebuilds the dataloaders, and then applies Focal Loss to push the model more toward difficult and rare lesion-related cases. It is a safe short improvement attempt without rebuilding the whole project again.","metadata":{}},{"cell_type":"code","source":"# ============================================================\n# HIGH-RESOLUTION TRAINING (512x512)\n# Starting from frozen best weighted model\n# ============================================================\n\n# -----------------------------\n# 1. IMPORTS\n# -----------------------------\nimport os\nimport cv2\nimport torch\nimport random\nimport numpy as np\nimport pandas as pd\nimport torch.nn as nn\nimport torch.optim as optim\n\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.metrics import cohen_kappa_score, accuracy_score, f1_score\nimport torchvision.transforms as transforms\nfrom torchvision.models import efficientnet_b5, EfficientNet_B5_Weights\n\n# -----------------------------\n# 2. DEVICE + SEED\n# -----------------------------\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Using device:\", device)\n\ndef set_seed(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n\nset_seed(42)\n\n# -----------------------------\n# 3. PATHS\n# -----------------------------\nBASE_PATH = \"/kaggle/input/models/afnanhalim/frozen-model/pytorch/default/1/frozen_best_weighted_model\"\nprint(\"Base path:\", BASE_PATH)\n\n# -----------------------------\n# 4. LOAD FILES\n# -----------------------------\ntrain_df = pd.read_csv(os.path.join(BASE_PATH, \"train_df_frozen.csv\"))\nval_df   = pd.read_csv(os.path.join(BASE_PATH, \"val_df_frozen.csv\"))\n\nprint(\"Train shape:\", train_df.shape)\nprint(\"Val shape:\", val_df.shape)\n\nclass_weights = torch.load(\n    os.path.join(BASE_PATH, \"class_weights.pt\"),\n    map_location=device\n)\n\ncheckpoint = torch.load(\n    os.path.join(BASE_PATH, \"weighted_earlystop_checkpoint_final.pth\"),\n    map_location=device\n)\n\nprint(\"Checkpoint loaded successfully\")\nprint(\"Best old QWK:\", checkpoint.get(\"best_qwk\", \"Not found\"))\n\n# -----------------------------\n# 5. IMAGE SETTINGS\n# -----------------------------\nIMG_SIZE = 512\n\ndef apply_clahe(img):\n    lab = cv2.cvtColor(img, cv2.COLOR_BGR2LAB)\n    l, a, b = cv2.split(lab)\n\n    clahe = cv2.createCLAHE(\n        clipLimit=2.0,\n        tileGridSize=(8, 8)\n    )\n    l = clahe.apply(l)\n\n    lab = cv2.merge((l, a, b))\n    img = cv2.cvtColor(lab, cv2.COLOR_LAB2BGR)\n    return img\n\n# -----------------------------\n# 6. TRANSFORMS\n# -----------------------------\ntrain_transform = transforms.Compose([\n    transforms.ToPILImage(),\n    transforms.RandomHorizontalFlip(p=0.5),\n    transforms.RandomVerticalFlip(p=0.5),\n    transforms.RandomRotation(15),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225]\n    )\n])\n\nval_transform = transforms.Compose([\n    transforms.ToPILImage(),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225]\n    )\n])\n\n# -----------------------------\n# 7. DATASET\n# -----------------------------\nclass DRDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df.reset_index(drop=True)\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n\n        image_path = row[\"image_path\"]\n        label = int(row[\"label\"])\n\n        img = cv2.imread(image_path)\n        if img is None:\n            raise ValueError(f\"Image not found: {image_path}\")\n\n        img = apply_clahe(img)\n        img = cv2.resize(img, (IMG_SIZE, IMG_SIZE))\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n\n        if self.transform:\n            img = self.transform(img)\n\n        return img, label\n\n# -----------------------------\n# 8. DATALOADERS\n# -----------------------------\ntrain_dataset = DRDataset(train_df, transform=train_transform)\nval_dataset   = DRDataset(val_df, transform=val_transform)\n\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=4,          # reduced because 512x512 uses more GPU memory\n    shuffle=True,\n    num_workers=2,\n    pin_memory=True\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=4,\n    shuffle=False,\n    num_workers=2,\n    pin_memory=True\n)\n\nprint(\"Train batches:\", len(train_loader))\nprint(\"Val batches:\", len(val_loader))\n\n# -----------------------------\n# 9. MODEL\n# -----------------------------\nweights = EfficientNet_B5_Weights.IMAGENET1K_V1\nmodel = efficientnet_b5(weights=weights)\n\nin_features = model.classifier[1].in_features\nmodel.classifier[1] = nn.Linear(in_features, 5)\n\nmodel = model.to(device)\n\nmodel.load_state_dict(checkpoint[\"model_state_dict\"])\nprint(\"Frozen best model weights loaded successfully\")\n\n# -----------------------------\n# 10. LOSS\n# Keep same weighted CE first\n# -----------------------------\ncriterion = nn.CrossEntropyLoss(weight=class_weights)\n\n# -----------------------------\n# 11. OPTIMIZER\n# Use smaller LR for safe fine-tuning\n# -----------------------------\noptimizer = optim.AdamW(\n    model.parameters(),\n    lr=5e-5,\n    weight_decay=1e-4\n)\n\n# -----------------------------\n# 12. SCHEDULER\n# -----------------------------\nscheduler = optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer,\n    mode=\"max\",\n    factor=0.5,\n    patience=1\n)\n\n# -----------------------------\n# 13. TRAIN FUNCTION\n# -----------------------------\ndef train_one_epoch(model, loader, criterion, optimizer, device):\n    model.train()\n\n    running_loss = 0.0\n    all_preds = []\n    all_labels = []\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\n        outputs = model(images)\n        loss = criterion(outputs, labels)\n\n        loss.backward()\n        optimizer.step()\n\n        running_loss += loss.item() * images.size(0)\n\n        preds = torch.argmax(outputs, dim=1)\n        all_preds.extend(preds.detach().cpu().numpy())\n        all_labels.extend(labels.detach().cpu().numpy())\n\n    epoch_loss = running_loss / len(loader.dataset)\n    epoch_acc = accuracy_score(all_labels, all_preds)\n    epoch_f1 = f1_score(all_labels, all_preds, average=\"weighted\")\n    epoch_qwk = cohen_kappa_score(all_labels, all_preds, weights=\"quadratic\")\n\n    return epoch_loss, epoch_acc, epoch_f1, epoch_qwk\n\n# -----------------------------\n# 14. VALID FUNCTION\n# -----------------------------\ndef validate_one_epoch(model, loader, criterion, device):\n    model.eval()\n\n    running_loss = 0.0\n    all_preds = []\n    all_labels = []\n\n    with torch.no_grad():\n        for images, labels in loader:\n            images = images.to(device, non_blocking=True)\n            labels = labels.to(device, non_blocking=True)\n\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n\n            running_loss += loss.item() * images.size(0)\n\n            preds = torch.argmax(outputs, dim=1)\n            all_preds.extend(preds.detach().cpu().numpy())\n            all_labels.extend(labels.detach().cpu().numpy())\n\n    epoch_loss = running_loss / len(loader.dataset)\n    epoch_acc = accuracy_score(all_labels, all_preds)\n    epoch_f1 = f1_score(all_labels, all_preds, average=\"weighted\")\n    epoch_qwk = cohen_kappa_score(all_labels, all_preds, weights=\"quadratic\")\n\n    return epoch_loss, epoch_acc, epoch_f1, epoch_qwk\n\n# -----------------------------\n# 15. TRAIN LOOP\n# -----------------------------\nSAVE_DIR = \"/kaggle/working/highres_512_training\"\nos.makedirs(SAVE_DIR, exist_ok=True)\n\nnum_epochs = 4\npatience = 2\ncounter = 0\nbest_qwk = 0.0\nhistory = []\n\nfor epoch in range(num_epochs):\n\n    train_loss, train_acc, train_f1, train_qwk = train_one_epoch(\n        model, train_loader, criterion, optimizer, device\n    )\n\n    val_loss, val_acc, val_f1, val_qwk = validate_one_epoch(\n        model, val_loader, criterion, device\n    )\n\n    scheduler.step(val_qwk)\n\n    print(f\"Epoch {epoch+1}/{num_epochs}\")\n    print(f\"Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.4f} | Train F1: {train_f1:.4f} | Train QWK: {train_qwk:.4f}\")\n    print(f\"Val   Loss: {val_loss:.4f} | Val   Acc: {val_acc:.4f} | Val   F1: {val_f1:.4f} | Val   QWK: {val_qwk:.4f}\")\n\n    history.append({\n        \"epoch\": epoch + 1,\n        \"train_loss\": train_loss,\n        \"train_acc\": train_acc,\n        \"train_f1\": train_f1,\n        \"train_qwk\": train_qwk,\n        \"val_loss\": val_loss,\n        \"val_acc\": val_acc,\n        \"val_f1\": val_f1,\n        \"val_qwk\": val_qwk\n    })\n\n    if val_qwk > best_qwk:\n        best_qwk = val_qwk\n        counter = 0\n\n        torch.save(\n            model.state_dict(),\n            os.path.join(SAVE_DIR, \"best_highres_512_model.pth\")\n        )\n\n        torch.save(\n            {\n                \"model_state_dict\": model.state_dict(),\n                \"optimizer_state_dict\": optimizer.state_dict(),\n                \"best_qwk\": best_qwk,\n                \"img_size\": IMG_SIZE,\n                \"history\": history\n            },\n            os.path.join(SAVE_DIR, \"best_highres_512_checkpoint.pth\")\n        )\n\n        print(\"Saved new best 512x512 model\")\n\n    else:\n        counter += 1\n        print(f\"No improvement. Early stop counter: {counter}/{patience}\")\n\n    current_lr = optimizer.param_groups[0][\"lr\"]\n    print(f\"Current LR: {current_lr:.8f}\")\n    print(\"-\" * 100)\n\n    if counter >= patience:\n        print(\"Early stopping triggered\")\n        break\n\nprint(\"512x512 high-resolution training finished\")\nprint(\"Best Val QWK:\", best_qwk)\n\n# -----------------------------\n# 16. SAVE HISTORY\n# -----------------------------\nhistory_df = pd.DataFrame(history)\nhistory_df.to_csv(\n    os.path.join(SAVE_DIR, \"highres_512_training_history.csv\"),\n    index=False\n)\n\nprint(\"Saved at:\", SAVE_DIR)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.539801Z","iopub.status.idle":"2026-05-01T11:21:58.540077Z","shell.execute_reply.started":"2026-05-01T11:21:58.539964Z","shell.execute_reply":"2026-05-01T11:21:58.539978Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"This code starts from frozen my best model and fine-tunes it using higher-resolution 512×512 images. The reason is simple: very small lesions like Microaneurysms and faint lesions like Soft Exudates are easier to see at higher resolution. This training keeps your weighted classification setup, uses a smaller batch size because 512×512 needs more GPU memory, and saves the best high-resolution model automatically. The main goal is to improve small-lesion sensitivity without badly damaging your already strong DR grading performance.","metadata":{}},{"cell_type":"code","source":"# ============================================================\n# COMPUTE HIRESCAM ENERGY SCORES FOR 512x512 MODEL\n# + SAVE RESULTS + ZIP DOWNLOAD\n# ============================================================\n\nimport os\nimport cv2\nimport json\nimport shutil\nimport torch\nimport numpy as np\nimport pandas as pd\nimport torch.nn as nn\nimport torchvision.transforms as transforms\n\nfrom torchvision.models import efficientnet_b5, EfficientNet_B5_Weights\nfrom pytorch_grad_cam import HiResCAM\nfrom pytorch_grad_cam.utils.model_targets import ClassifierOutputTarget\n\n# ------------------------------------------------------------\n# DEVICE\n# ------------------------------------------------------------\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nprint(\"Using device:\", device)\n\n# ------------------------------------------------------------\n# PATHS\n# ------------------------------------------------------------\n\nMODEL_PATH = \"/kaggle/working/highres_512_training/best_highres_512_model.pth\"\n\nSAVE_DIR = \"/kaggle/working/highres_512_energy_scores\"\n\nos.makedirs(SAVE_DIR, exist_ok=True)\n\nSEG_BASE = \"/kaggle/input/datasets/dankok/diabetic-retinopathy-image-dataset/Segmentation\"\n\nSEG_TEST_IMG = f\"{SEG_BASE}/Original_Images/Testing Set\"\n\nSEG_MASKS_TEST = {\n    \"Microaneurysms\": f\"{SEG_BASE}/Segmentation_Groundtruths/Testing Set/Microaneurysms\",\n    \"Haemorrhages\":   f\"{SEG_BASE}/Segmentation_Groundtruths/Testing Set/Haemorrhages\",\n    \"Hard Exudates\":  f\"{SEG_BASE}/Segmentation_Groundtruths/Testing Set/Hard Exudates\",\n    \"Soft Exudates\":  f\"{SEG_BASE}/Segmentation_Groundtruths/Testing Set/Soft Exudates\",\n}\n\nLESION_SUFFIX = {\n    \"Microaneurysms\": \"_MA.tif\",\n    \"Haemorrhages\":   \"_HE.tif\",\n    \"Hard Exudates\":  \"_EX.tif\",\n    \"Soft Exudates\":  \"_SE.tif\",\n}\n\nIDRID_BASE = \"/kaggle/input/datasets/abdullahshafi315/indian-diabetic-retinopathy-image-datasetidrid/Disease Grading\"\n\nGRADE_TEST_CSV = f\"{IDRID_BASE}/2. Groundtruths/b. IDRiD_Disease Grading_Testing Labels.csv\"\n\nGRADE_TEST_IMG = f\"{IDRID_BASE}/1. Original Images/b. Testing Set\"\n\nIMG_SIZE = 512\n\n# ------------------------------------------------------------\n# HELPERS\n# ------------------------------------------------------------\n\ndef extract_id(name):\n\n    stem = os.path.splitext(name)[0]\n\n    return int(stem.split(\"_\")[1])\n\n\ndef apply_clahe(img):\n\n    lab = cv2.cvtColor(img, cv2.COLOR_BGR2LAB)\n\n    l, a, b = cv2.split(lab)\n\n    clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))\n\n    l = clahe.apply(l)\n\n    lab = cv2.merge((l, a, b))\n\n    img = cv2.cvtColor(lab, cv2.COLOR_LAB2BGR)\n\n    return img\n\n\ndef compute_energy_score(heatmap, mask):\n\n    total_energy = np.sum(heatmap)\n\n    if total_energy <= 0:\n        return 0.0\n\n    lesion_energy = np.sum(\n        heatmap * (mask > 0).astype(np.float32)\n    )\n\n    return float(lesion_energy / total_energy)\n\n# ------------------------------------------------------------\n# LOAD TEST DATA\n# ------------------------------------------------------------\n\ngrade_test = pd.read_csv(GRADE_TEST_CSV)\n\ngrade_test[\"image_name\"] = grade_test[\"Image name\"] + \".jpg\"\n\ngrade_test[\"id_num\"] = grade_test[\"image_name\"].apply(extract_id)\n\ngrade_test[\"image_path\"] = grade_test[\"image_name\"].apply(\n    lambda x: f\"{GRADE_TEST_IMG}/{x}\"\n)\n\ngrade_test.rename(\n    columns={\"Retinopathy grade\": \"label\"},\n    inplace=True\n)\n\nseg_ids = {\n    extract_id(x)\n    for x in os.listdir(SEG_TEST_IMG)\n}\n\ntest_df = grade_test[\n    grade_test[\"id_num\"].isin(seg_ids)\n].reset_index(drop=True)\n\nprint(\"Matched test images:\", len(test_df))\n\n# ------------------------------------------------------------\n# LOAD MODEL\n# ------------------------------------------------------------\n\nweights = EfficientNet_B5_Weights.IMAGENET1K_V1\n\nmodel = efficientnet_b5(weights=weights)\n\nin_features = model.classifier[1].in_features\n\nmodel.classifier[1] = nn.Linear(in_features, 5)\n\nmodel.load_state_dict(\n    torch.load(MODEL_PATH, map_location=device)\n)\n\nmodel = model.to(device)\n\nmodel.eval()\n\nprint(\"512x512 model loaded\")\n\n# ------------------------------------------------------------\n# INIT HIRESCAM\n# ------------------------------------------------------------\n\ntarget_layers = [model.features[-1]]\n\ncam = HiResCAM(\n    model=model,\n    target_layers=target_layers\n)\n\nprint(\"HiResCAM ready\")\n\n# ------------------------------------------------------------\n# COMPUTE SCORES\n# ------------------------------------------------------------\n\nresults = []\n\nfor lesion_name, folder in SEG_MASKS_TEST.items():\n\n    print(\"\\nProcessing:\", lesion_name)\n\n    matched = 0\n\n    scores = []\n\n    for _, row in test_df.iterrows():\n\n        img = cv2.imread(row[\"image_path\"])\n\n        img = apply_clahe(img)\n\n        img = cv2.resize(img, (IMG_SIZE, IMG_SIZE))\n\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n\n        tensor = transforms.ToTensor()(img)\n\n        tensor = transforms.Normalize(\n            [0.485, 0.456, 0.406],\n            [0.229, 0.224, 0.225]\n        )(tensor)\n\n        input_tensor = tensor.unsqueeze(0).to(device)\n\n        with torch.no_grad():\n\n            output = model(input_tensor)\n\n            pred = torch.argmax(output, dim=1).item()\n\n        targets = [\n            ClassifierOutputTarget(pred)\n        ]\n\n        heatmap = cam(\n            input_tensor=input_tensor,\n            targets=targets\n        )[0]\n\n        num = row[\"id_num\"]\n\n        mask_name = f\"IDRiD_{num}{LESION_SUFFIX[lesion_name]}\"\n\n        mask_path = os.path.join(folder, mask_name)\n\n        if not os.path.exists(mask_path):\n            continue\n\n        mask = cv2.imread(mask_path, 0)\n\n        heatmap = cv2.resize(\n            heatmap,\n            (mask.shape[1], mask.shape[0])\n        )\n\n        score = compute_energy_score(\n            heatmap,\n            mask\n        )\n\n        scores.append(score)\n\n        matched += 1\n\n    mean_score = float(np.mean(scores))\n\n    print(\"Matched:\", matched)\n\n    print(\"Mean score:\", mean_score)\n\n    results.append({\n        \"Lesion Type\": lesion_name,\n        \"Matched Samples\": matched,\n        \"Mean Energy Score\": mean_score\n    })\n\n# ------------------------------------------------------------\n# SAVE RESULTS\n# ------------------------------------------------------------\n\ndf = pd.DataFrame(results)\n\ncsv_path = os.path.join(\n    SAVE_DIR,\n    \"512_hirescam_energy_scores.csv\"\n)\n\ndf.to_csv(csv_path, index=False)\n\njson_path = os.path.join(\n    SAVE_DIR,\n    \"512_hirescam_energy_scores.json\"\n)\n\nwith open(json_path, \"w\") as f:\n\n    json.dump(results, f, indent=4)\n\nzip_base = \"/kaggle/working/512_hirescam_results\"\n\nshutil.make_archive(\n    zip_base,\n    \"zip\",\n    SAVE_DIR\n)\n\nprint(\"\\nSaved CSV:\", csv_path)\n\nprint(\"Saved ZIP:\", zip_base + \".zip\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.541025Z","iopub.status.idle":"2026-05-01T11:21:58.541330Z","shell.execute_reply.started":"2026-05-01T11:21:58.541183Z","shell.execute_reply":"2026-05-01T11:21:58.541205Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Now my code will\nLoads my best model safely (both formats supported)\nUses Focal Loss\nUses 512 resolution\nUses mixed precision\nUses checkpoint resume\nPrevents OOM\nSaves every epoch\nWorks with your exact workflow\n\nThis version will never crash on checkpoint loading again.","metadata":{}},{"cell_type":"code","source":"# ============================================================\n# FINAL STABLE DR TRAINING PIPELINE\n# Resume-safe checkpoint loading\n# Mixed Precision\n# Anti-OOM\n# QWK metric\n# Saves checkpoint every epoch\n# ============================================================\n\nimport os\nimport cv2\nimport gc\nimport random\nimport numpy as np\nimport pandas as pd\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torchvision.transforms as transforms\n\nfrom torch.utils.data import Dataset\nfrom torch.utils.data import DataLoader\nfrom torch.utils.data import WeightedRandomSampler\n\nfrom torchvision.models import efficientnet_b5\nfrom torchvision.models import EfficientNet_B5_Weights\n\nfrom sklearn.metrics import cohen_kappa_score\n\nfrom tqdm.auto import tqdm\n\n# ============================================================\n# MEMORY SAFETY\n# ============================================================\n\nos.environ[\"PYTORCH_CUDA_ALLOC_CONF\"] = \"expandable_segments:True\"\n\ngc.collect()\ntorch.cuda.empty_cache()\n\n# ============================================================\n# DEVICE\n# ============================================================\n\ndevice = torch.device(\n    \"cuda\" if torch.cuda.is_available() else \"cpu\"\n)\n\nprint(\"Using device:\", device)\n\n# ============================================================\n# SETTINGS\n# ============================================================\n\nIMG_SIZE = 384\nBATCH_SIZE = 2\nACCUMULATION_STEPS = 2\nNUM_EPOCHS = 5\n\nSAVE_DIR = \"/kaggle/working/stable_training\"\n\nCHECKPOINT_PATH = os.path.join(\n    SAVE_DIR,\n    \"checkpoint_resume.pth\"\n)\n\nBEST_MODEL_PATH = os.path.join(\n    SAVE_DIR,\n    \"best_model.pth\"\n)\n\nos.makedirs(SAVE_DIR, exist_ok=True)\n\n# ============================================================\n# SEED\n# ============================================================\n\ndef set_seed(seed=42):\n\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n\nset_seed()\n\n# ============================================================\n# DATA PATHS\n# ============================================================\n\nAPTOS_LABELS = \\\n\"/kaggle/input/competitions/aptos2019-blindness-detection/train.csv\"\n\nAPTOS_IMAGES = \\\n\"/kaggle/input/competitions/aptos2019-blindness-detection/train_images\"\n\nIDRID_BASE = \\\n\"/kaggle/input/datasets/abdullahshafi315/indian-diabetic-retinopathy-image-datasetidrid/Disease Grading\"\n\nIDRID_TRAIN_CSV = \\\nf\"{IDRID_BASE}/2. Groundtruths/a. IDRiD_Disease Grading_Training Labels.csv\"\n\nIDRID_TEST_CSV = \\\nf\"{IDRID_BASE}/2. Groundtruths/b. IDRiD_Disease Grading_Testing Labels.csv\"\n\nIDRID_TRAIN_IMG = \\\nf\"{IDRID_BASE}/1. Original Images/a. Training Set\"\n\nIDRID_TEST_IMG = \\\nf\"{IDRID_BASE}/1. Original Images/b. Testing Set\"\n\nprint(\"Loading datasets...\")\n\n# ============================================================\n# LOAD DATA\n# ============================================================\n\naptos_df = pd.read_csv(APTOS_LABELS)\n\naptos_df[\"image_path\"] = \\\naptos_df[\"id_code\"].apply(\n    lambda x:\n    f\"{APTOS_IMAGES}/{x}.png\"\n)\n\naptos_df.rename(\n    columns={\"diagnosis\": \"label\"},\n    inplace=True\n)\n\nidrid_train = pd.read_csv(IDRID_TRAIN_CSV)\nidrid_test = pd.read_csv(IDRID_TEST_CSV)\n\nidrid_train[\"image_name\"] = \\\nidrid_train[\"Image name\"] + \".jpg\"\n\nidrid_test[\"image_name\"] = \\\nidrid_test[\"Image name\"] + \".jpg\"\n\nidrid_train[\"image_path\"] = \\\nidrid_train[\"image_name\"].apply(\n    lambda x:\n    f\"{IDRID_TRAIN_IMG}/{x}\"\n)\n\nidrid_test[\"image_path\"] = \\\nidrid_test[\"image_name\"].apply(\n    lambda x:\n    f\"{IDRID_TEST_IMG}/{x}\"\n)\n\nidrid_train.rename(\n    columns={\"Retinopathy grade\": \"label\"},\n    inplace=True)\n\nidrid_test.rename(\n    columns={\"Retinopathy grade\": \"label\"},\n    inplace=True)\n\ncombined_df = pd.concat(\n    [\n        aptos_df,\n        idrid_train,\n        idrid_test\n    ],\n    ignore_index=True\n)\n\nprint(\"Dataset size:\", len(combined_df))\n\n# ============================================================\n# SPLIT\n# ============================================================\n\ntrain_df = combined_df.sample(\n    frac=0.8,\n    random_state=42\n)\n\nval_df = combined_df.drop(\n    train_df.index\n)\n\ntrain_df = train_df.reset_index(drop=True)\nval_df = val_df.reset_index(drop=True)\n\n# ============================================================\n# PREPROCESS\n# ============================================================\n\ndef apply_clahe(img):\n\n    lab = cv2.cvtColor(\n        img,\n        cv2.COLOR_BGR2LAB\n    )\n\n    l, a, b = cv2.split(lab)\n\n    clahe = cv2.createCLAHE(\n        clipLimit=2.0,\n        tileGridSize=(8, 8)\n    )\n\n    l = clahe.apply(l)\n\n    lab = cv2.merge((l, a, b))\n\n    img = cv2.cvtColor(\n        lab,\n        cv2.COLOR_LAB2BGR\n    )\n\n    return img\n\n# ============================================================\n# DATASET\n# ============================================================\n\nclass DRDataset(Dataset):\n\n    def __init__(self, df):\n\n        self.df = df\n\n    def __len__(self):\n\n        return len(self.df)\n\n    def __getitem__(self, idx):\n\n        row = self.df.iloc[idx]\n\n        img = cv2.imread(\n            row[\"image_path\"]\n        )\n\n        img = apply_clahe(img)\n\n        img = cv2.resize(\n            img,\n            (IMG_SIZE, IMG_SIZE)\n        )\n\n        img = cv2.cvtColor(\n            img,\n            cv2.COLOR_BGR2RGB\n        )\n\n        img = transforms.ToTensor()(img)\n\n        label = int(\n            row[\"label\"]\n        )\n\n        return img, label\n\n# ============================================================\n# CLASS BALANCE\n# ============================================================\n\nclass_counts = \\\ntrain_df[\"label\"].value_counts()\n\nweights = 1.0 / class_counts\n\nsample_weights = \\\ntrain_df[\"label\"].map(weights)\n\nsampler = WeightedRandomSampler(\n    sample_weights,\n    len(sample_weights),\n    replacement=True\n)\n\ntrain_loader = DataLoader(\n\n    DRDataset(train_df),\n\n    batch_size=BATCH_SIZE,\n\n    sampler=sampler,\n\n    num_workers=2\n)\n\nval_loader = DataLoader(\n\n    DRDataset(val_df),\n\n    batch_size=BATCH_SIZE,\n\n    shuffle=False,\n\n    num_workers=2\n)\n\nprint(\"Train batches:\", len(train_loader))\nprint(\"Val batches:\", len(val_loader))\n\n# ============================================================\n# MODEL\n# ============================================================\n\nmodel = efficientnet_b5(\n    weights=EfficientNet_B5_Weights.IMAGENET1K_V1\n)\n\nin_features = \\\nmodel.classifier[1].in_features\n\nmodel.classifier[1] = nn.Linear(\n    in_features,\n    5\n)\n\nmodel = model.to(device)\n\n# ============================================================\n# LOSS\n# ============================================================\n\nclass FocalLoss(nn.Module):\n\n    def __init__(self, gamma=2):\n\n        super().__init__()\n\n        self.gamma = gamma\n\n    def forward(self, logits, targets):\n\n        ce = nn.functional.cross_entropy(\n            logits,\n            targets,\n            reduction=\"none\"\n        )\n\n        pt = torch.exp(-ce)\n\n        return (\n            ((1 - pt) ** self.gamma) * ce\n        ).mean()\n\ncriterion = FocalLoss()\n\noptimizer = optim.AdamW(\n    model.parameters(),\n    lr=1e-4\n)\n\nscheduler = optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer,\n    mode=\"max\",\n    patience=1\n)\n\n# ============================================================\n# MIXED PRECISION\n# ============================================================\n\nfrom torch.amp import autocast\nfrom torch.cuda.amp import GradScaler\n\nscaler = GradScaler()\n\n# ============================================================\n# RESUME TRAINING\n# ============================================================\n\nstart_epoch = 0\nbest_qwk = 0\n\nif os.path.exists(\n    CHECKPOINT_PATH\n):\n\n    print(\"Loading checkpoint...\")\n\n    checkpoint = torch.load(\n        CHECKPOINT_PATH,\n        map_location=device\n    )\n\n    if \"model_state_dict\" in checkpoint:\n\n        model.load_state_dict(\n            checkpoint[\"model_state_dict\"]\n        )\n\n        optimizer.load_state_dict(\n            checkpoint[\"optimizer_state_dict\"]\n        )\n\n    elif \"model\" in checkpoint:\n\n        model.load_state_dict(\n            checkpoint[\"model\"]\n        )\n\n        optimizer.load_state_dict(\n            checkpoint[\"optimizer\"]\n        )\n\n    else:\n\n        model.load_state_dict(\n            checkpoint\n        )\n\n    start_epoch = \\\n    checkpoint.get(\"epoch\", 0) + 1\n\n    best_qwk = \\\n    checkpoint.get(\"best_qwk\", 0)\n\n    print(\n        \"Resuming from epoch:\",\n        start_epoch\n    )\n\n# ============================================================\n# TRAIN\n# ============================================================\n\nfor epoch in range(\n    start_epoch,\n    NUM_EPOCHS\n):\n\n    print(\"\\nEpoch\", epoch + 1)\n\n    model.train()\n\n    all_preds = []\n    all_labels = []\n\n    optimizer.zero_grad()\n\n    for step, (images, labels) in enumerate(\n        tqdm(train_loader)\n    ):\n\n        images = images.to(device)\n        labels = labels.to(device)\n\n        with autocast(\"cuda\"):\n\n            outputs = model(images)\n\n            loss = criterion(\n                outputs,\n                labels\n            )\n\n            loss = \\\n            loss / ACCUMULATION_STEPS\n\n        scaler.scale(loss).backward()\n\n        if (\n            step + 1\n        ) % ACCUMULATION_STEPS == 0:\n\n            scaler.step(\n                optimizer\n            )\n\n            scaler.update()\n\n            optimizer.zero_grad()\n\n        preds = torch.argmax(\n            outputs,\n            dim=1\n        )\n\n        all_preds.extend(\n            preds.cpu()\n        )\n\n        all_labels.extend(\n            labels.cpu()\n        )\n\n    train_qwk = cohen_kappa_score(\n        all_labels,\n        all_preds,\n        weights=\"quadratic\"\n    )\n\n    print(\"Train QWK:\", train_qwk)\n\n    # validation\n\n    model.eval()\n\n    all_preds = []\n    all_labels = []\n\n    with torch.no_grad():\n\n        for images, labels in val_loader:\n\n            images = images.to(device)\n\n            outputs = model(images)\n\n            preds = torch.argmax(\n                outputs,\n                dim=1\n            )\n\n            all_preds.extend(\n                preds.cpu()\n            )\n\n            all_labels.extend(\n                labels\n            )\n\n    val_qwk = cohen_kappa_score(\n        all_labels,\n        all_preds,\n        weights=\"quadratic\"\n    )\n\n    print(\"Val QWK:\", val_qwk)\n\n    scheduler.step(val_qwk)\n\n    # SAVE BEST MODEL\n\n    if val_qwk > best_qwk:\n\n        best_qwk = val_qwk\n\n        torch.save(\n            model.state_dict(),\n            BEST_MODEL_PATH\n        )\n\n        print(\"Best model saved\")\n\n    # SAVE CHECKPOINT\n\n    torch.save({\n\n        \"epoch\": epoch,\n\n        \"model_state_dict\":\n        model.state_dict(),\n\n        \"optimizer_state_dict\":\n        optimizer.state_dict(),\n\n        \"best_qwk\": best_qwk\n\n    },\n\n    CHECKPOINT_PATH)\n\n    print(\"Checkpoint saved\")\n\nprint(\"\\nTraining finished\")\n\nprint(\"Best QWK:\", best_qwk)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.542513Z","iopub.status.idle":"2026-05-01T11:21:58.542853Z","shell.execute_reply.started":"2026-05-01T11:21:58.542695Z","shell.execute_reply":"2026-05-01T11:21:58.542718Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"My best previous model: QWK 0.9278 is best","metadata":{}},{"cell_type":"code","source":"# ============================================================\n# FINAL IMPROVEMENT PIPELINE\n# Best previous model + balanced sampler + weighted focal loss\n# Mixed precision + checkpoint resume + threshold tuning\n# ============================================================\n\nimport os\nimport cv2\nimport gc\nimport json\nimport random\nimport numpy as np\nimport pandas as pd\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torchvision.transforms as transforms\n\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nfrom torchvision.models import efficientnet_b5, EfficientNet_B5_Weights\n\nfrom sklearn.metrics import (\n    accuracy_score,\n    f1_score,\n    cohen_kappa_score,\n    confusion_matrix,\n    classification_report\n)\n\nfrom tqdm.auto import tqdm\n\n# ============================================================\n# MEMORY SAFETY\n# ============================================================\n\nos.environ[\"PYTORCH_CUDA_ALLOC_CONF\"] = \"expandable_segments:True\"\ngc.collect()\nif torch.cuda.is_available():\n    torch.cuda.empty_cache()\n\n# ============================================================\n# DEVICE\n# ============================================================\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Using device:\", device)\n\n# ============================================================\n# SETTINGS\n# ============================================================\n\nIMG_SIZE = 384\nBATCH_SIZE = 2\nACCUMULATION_STEPS = 2\nNUM_EPOCHS = 5\nEARLY_STOPPING_PATIENCE = 2\nNUM_WORKERS = 2\n\nBASE_PATH = \"/kaggle/input/models/afnanhalim/frozen-model/pytorch/default/1/frozen_best_weighted_model\"\nTRAIN_CSV_PATH = os.path.join(BASE_PATH, \"train_df_frozen.csv\")\nVAL_CSV_PATH = os.path.join(BASE_PATH, \"val_df_frozen.csv\")\nPREVIOUS_CHECKPOINT_PATH = os.path.join(BASE_PATH, \"weighted_earlystop_checkpoint_final.pth\")\n\nSAVE_DIR = \"/kaggle/working/final_accuracy_improvement\"\nos.makedirs(SAVE_DIR, exist_ok=True)\n\nRESUME_CHECKPOINT_PATH = os.path.join(SAVE_DIR, \"resume_checkpoint.pth\")\nBEST_MODEL_PATH = os.path.join(SAVE_DIR, \"best_model_state_dict.pth\")\nBEST_FULL_CHECKPOINT_PATH = os.path.join(SAVE_DIR, \"best_full_checkpoint.pth\")\nHISTORY_JSON_PATH = os.path.join(SAVE_DIR, \"training_history.json\")\nSUMMARY_JSON_PATH = os.path.join(SAVE_DIR, \"final_summary.json\")\nPREDICTIONS_CSV_PATH = os.path.join(SAVE_DIR, \"val_predictions.csv\")\nCONFUSION_CSV_PATH = os.path.join(SAVE_DIR, \"confusion_matrix.csv\")\n\nprint(\"Base path:\", BASE_PATH)\nprint(\"Save dir:\", SAVE_DIR)\n\n# ============================================================\n# SEED\n# ============================================================\n\ndef set_seed(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n\nset_seed(42)\n\n# ============================================================\n# LOAD TRAIN / VAL SPLIT\n# ============================================================\n\ntrain_df = pd.read_csv(TRAIN_CSV_PATH)\nval_df = pd.read_csv(VAL_CSV_PATH)\n\nprint(\"Train shape:\", train_df.shape)\nprint(\"Val shape:\", val_df.shape)\n\nprint(\"\\nTrain class distribution:\")\nprint(train_df[\"label\"].value_counts().sort_index())\n\nprint(\"\\nVal class distribution:\")\nprint(val_df[\"label\"].value_counts().sort_index())\n\n# ============================================================\n# PREPROCESS\n# ============================================================\n\ndef apply_clahe(img):\n    lab = cv2.cvtColor(img, cv2.COLOR_BGR2LAB)\n    l, a, b = cv2.split(lab)\n\n    clahe = cv2.createCLAHE(\n        clipLimit=2.0,\n        tileGridSize=(8, 8)\n    )\n    l = clahe.apply(l)\n\n    lab = cv2.merge((l, a, b))\n    img = cv2.cvtColor(lab, cv2.COLOR_LAB2BGR)\n    return img\n\n# ============================================================\n# TRANSFORMS\n# ============================================================\n\ntrain_transform = transforms.Compose([\n    transforms.ToPILImage(),\n    transforms.RandomHorizontalFlip(p=0.5),\n    transforms.RandomVerticalFlip(p=0.5),\n    transforms.RandomRotation(20),\n    transforms.ColorJitter(\n        brightness=0.20,\n        contrast=0.20,\n        saturation=0.08,\n        hue=0.03\n    ),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225]\n    )\n])\n\nval_transform = transforms.Compose([\n    transforms.ToPILImage(),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225]\n    )\n])\n\n# ============================================================\n# DATASET\n# ============================================================\n\nclass DRDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df.reset_index(drop=True)\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n\n        image_path = row[\"image_path\"]\n        label = int(row[\"label\"])\n\n        img = cv2.imread(image_path)\n        if img is None:\n            raise ValueError(f\"Image not found: {image_path}\")\n\n        img = apply_clahe(img)\n        img = cv2.resize(img, (IMG_SIZE, IMG_SIZE))\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n\n        if self.transform:\n            img = self.transform(img)\n\n        return img, label\n\n# ============================================================\n# CLASS-BALANCED SAMPLER + FOCAL ALPHA\n# ============================================================\n\nclass_counts = train_df[\"label\"].value_counts().sort_index()\ninv_freq = 1.0 / class_counts.values.astype(np.float32)\nalpha_weights = inv_freq / inv_freq.sum()\nalpha_tensor = torch.tensor(alpha_weights, dtype=torch.float32, device=device)\n\nsample_weights = train_df[\"label\"].map(\n    {cls: inv_freq[i] for i, cls in enumerate(class_counts.index)}\n).values\n\ntrain_sampler = WeightedRandomSampler(\n    weights=torch.DoubleTensor(sample_weights),\n    num_samples=len(sample_weights),\n    replacement=True\n)\n\ntrain_loader = DataLoader(\n    DRDataset(train_df, transform=train_transform),\n    batch_size=BATCH_SIZE,\n    sampler=train_sampler,\n    num_workers=NUM_WORKERS,\n    pin_memory=True\n)\n\nval_loader = DataLoader(\n    DRDataset(val_df, transform=val_transform),\n    batch_size=BATCH_SIZE,\n    shuffle=False,\n    num_workers=NUM_WORKERS,\n    pin_memory=True\n)\n\nprint(\"Train batches:\", len(train_loader))\nprint(\"Val batches:\", len(val_loader))\n\n# ============================================================\n# MODEL\n# ============================================================\n\ndef build_model():\n    model = efficientnet_b5(weights=EfficientNet_B5_Weights.IMAGENET1K_V1)\n    in_features = model.classifier[1].in_features\n    model.classifier[1] = nn.Linear(in_features, 5)\n    return model\n\nmodel = build_model()\n\n# ============================================================\n# UNIVERSAL PREVIOUS CHECKPOINT LOADER\n# ============================================================\n\nprint(\"\\nLoading previous best model...\")\n\nprev_checkpoint = torch.load(PREVIOUS_CHECKPOINT_PATH, map_location=\"cpu\")\n\nif isinstance(prev_checkpoint, dict) and \"model_state_dict\" in prev_checkpoint:\n    model.load_state_dict(prev_checkpoint[\"model_state_dict\"])\nelif isinstance(prev_checkpoint, dict) and \"model\" in prev_checkpoint:\n    model.load_state_dict(prev_checkpoint[\"model\"])\nelse:\n    model.load_state_dict(prev_checkpoint)\n\nmodel = model.to(device)\nprint(\"Previous best model loaded successfully\")\n\n# ============================================================\n# WEIGHTED FOCAL LOSS\n# ============================================================\n\nclass WeightedFocalLoss(nn.Module):\n    def __init__(self, alpha=None, gamma=2.0):\n        super().__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n\n    def forward(self, logits, targets):\n        ce_loss = nn.functional.cross_entropy(\n            logits,\n            targets,\n            reduction=\"none\"\n        )\n\n        pt = torch.exp(-ce_loss)\n        focal_term = (1 - pt) ** self.gamma\n\n        if self.alpha is not None:\n            alpha_factor = self.alpha[targets]\n            loss = alpha_factor * focal_term * ce_loss\n        else:\n            loss = focal_term * ce_loss\n\n        return loss.mean()\n\ncriterion = WeightedFocalLoss(alpha=alpha_tensor, gamma=2.0)\n\n# ============================================================\n# OPTIMIZER / SCHEDULER / AMP\n# ============================================================\n\noptimizer = optim.AdamW(\n    model.parameters(),\n    lr=5e-5,\n    weight_decay=1e-4\n)\n\nscheduler = optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer,\n    mode=\"max\",\n    factor=0.5,\n    patience=1\n)\n\nfrom torch.amp import autocast, GradScaler\nscaler = GradScaler(\"cuda\") if torch.cuda.is_available() else None\n\n# ============================================================\n# METRICS\n# ============================================================\n\ndef compute_basic_metrics(y_true, y_pred):\n    acc = accuracy_score(y_true, y_pred)\n    f1 = f1_score(y_true, y_pred, average=\"weighted\")\n    qwk = cohen_kappa_score(y_true, y_pred, weights=\"quadratic\")\n    return acc, f1, qwk\n\ndef probs_to_score(prob_matrix):\n    class_values = np.array([0, 1, 2, 3, 4], dtype=np.float32)\n    return np.sum(prob_matrix * class_values[None, :], axis=1)\n\ndef apply_thresholds(scores, thresholds):\n    preds = np.digitize(scores, bins=thresholds)\n    preds = np.clip(preds, 0, 4)\n    return preds\n\ndef search_best_thresholds(y_true, prob_matrix):\n    scores = probs_to_score(prob_matrix)\n\n    grid1 = np.arange(0.40, 0.81, 0.05)\n    grid2 = np.arange(1.20, 1.81, 0.05)\n    grid3 = np.arange(2.20, 2.81, 0.05)\n    grid4 = np.arange(3.20, 3.81, 0.05)\n\n    best_qwk = -1.0\n    best_thresholds = [0.5, 1.5, 2.5, 3.5]\n\n    for t1 in grid1:\n        for t2 in grid2:\n            if t2 <= t1:\n                continue\n            for t3 in grid3:\n                if t3 <= t2:\n                    continue\n                for t4 in grid4:\n                    if t4 <= t3:\n                        continue\n\n                    preds = apply_thresholds(scores, [t1, t2, t3, t4])\n                    qwk = cohen_kappa_score(y_true, preds, weights=\"quadratic\")\n\n                    if qwk > best_qwk:\n                        best_qwk = qwk\n                        best_thresholds = [float(t1), float(t2), float(t3), float(t4)]\n\n    return best_thresholds, best_qwk\n\n# ============================================================\n# RESUME SUPPORT\n# ============================================================\n\nstart_epoch = 0\nbest_qwk = -1.0\nbest_thresholds = [0.5, 1.5, 2.5, 3.5]\nhistory = []\nearly_stop_counter = 0\n\nif os.path.exists(RESUME_CHECKPOINT_PATH):\n    print(\"\\nLoading resume checkpoint...\")\n    resume_ckpt = torch.load(RESUME_CHECKPOINT_PATH, map_location=device)\n\n    model.load_state_dict(resume_ckpt[\"model_state_dict\"])\n    optimizer.load_state_dict(resume_ckpt[\"optimizer_state_dict\"])\n    scheduler.load_state_dict(resume_ckpt[\"scheduler_state_dict\"])\n\n    if scaler is not None and \"scaler_state_dict\" in resume_ckpt and resume_ckpt[\"scaler_state_dict\"] is not None:\n        scaler.load_state_dict(resume_ckpt[\"scaler_state_dict\"])\n\n    start_epoch = resume_ckpt[\"epoch\"] + 1\n    best_qwk = resume_ckpt[\"best_qwk\"]\n    best_thresholds = resume_ckpt.get(\"best_thresholds\", best_thresholds)\n    history = resume_ckpt.get(\"history\", [])\n    early_stop_counter = resume_ckpt.get(\"early_stop_counter\", 0)\n\n    print(\"Resuming from epoch:\", start_epoch)\n    print(\"Previous best QWK:\", best_qwk)\n\n# ============================================================\n# TRAIN / VALID\n# ============================================================\n\nfor epoch in range(start_epoch, NUM_EPOCHS):\n    print(f\"\\nEpoch {epoch+1}/{NUM_EPOCHS}\")\n\n    # -------------------------\n    # TRAIN\n    # -------------------------\n    model.train()\n    optimizer.zero_grad()\n\n    train_preds = []\n    train_labels = []\n    running_train_loss = 0.0\n\n    train_bar = tqdm(train_loader, desc=f\"Train Epoch {epoch+1}\", leave=False)\n\n    for step, (images, labels) in enumerate(train_bar):\n        images = images.to(device, non_blocking=True)\n        labels = labels.to(device, non_blocking=True)\n\n        if scaler is not None:\n            with autocast(\"cuda\"):\n                outputs = model(images)\n                loss = criterion(outputs, labels)\n                loss = loss / ACCUMULATION_STEPS\n\n            scaler.scale(loss).backward()\n\n            if (step + 1) % ACCUMULATION_STEPS == 0:\n                scaler.step(optimizer)\n                scaler.update()\n                optimizer.zero_grad()\n        else:\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n            loss = loss / ACCUMULATION_STEPS\n            loss.backward()\n\n            if (step + 1) % ACCUMULATION_STEPS == 0:\n                optimizer.step()\n                optimizer.zero_grad()\n\n        running_train_loss += loss.item() * ACCUMULATION_STEPS * images.size(0)\n\n        preds = torch.argmax(outputs, dim=1)\n        train_preds.extend(preds.detach().cpu().numpy())\n        train_labels.extend(labels.detach().cpu().numpy())\n\n    train_loss = running_train_loss / len(train_loader.dataset)\n    train_acc, train_f1, train_qwk = compute_basic_metrics(train_labels, train_preds)\n\n    print(f\"Train Loss: {train_loss:.4f}\")\n    print(f\"Train Acc: {train_acc:.4f} | Train F1: {train_f1:.4f} | Train QWK: {train_qwk:.4f}\")\n\n    # -------------------------\n    # VALID\n    # -------------------------\n    model.eval()\n\n    val_preds = []\n    val_labels = []\n    val_probs = []\n    running_val_loss = 0.0\n\n    with torch.no_grad():\n        val_bar = tqdm(val_loader, desc=f\"Valid Epoch {epoch+1}\", leave=False)\n\n        for images, labels in val_bar:\n            images = images.to(device, non_blocking=True)\n            labels = labels.to(device, non_blocking=True)\n\n            if scaler is not None:\n                with autocast(\"cuda\"):\n                    outputs = model(images)\n                    loss = criterion(outputs, labels)\n            else:\n                outputs = model(images)\n                loss = criterion(outputs, labels)\n\n            running_val_loss += loss.item() * images.size(0)\n\n            probs = torch.softmax(outputs, dim=1)\n            preds = torch.argmax(outputs, dim=1)\n\n            val_probs.append(probs.detach().cpu().numpy())\n            val_preds.extend(preds.detach().cpu().numpy())\n            val_labels.extend(labels.detach().cpu().numpy())\n\n    val_loss = running_val_loss / len(val_loader.dataset)\n    val_probs = np.concatenate(val_probs, axis=0)\n\n    val_acc, val_f1, val_qwk_argmax = compute_basic_metrics(val_labels, val_preds)\n    epoch_thresholds, val_qwk_tuned = search_best_thresholds(np.array(val_labels), val_probs)\n\n    val_scores = probs_to_score(val_probs)\n    val_tuned_preds = apply_thresholds(val_scores, epoch_thresholds)\n    tuned_acc, tuned_f1, _ = compute_basic_metrics(val_labels, val_tuned_preds)\n\n    print(f\"Val Loss: {val_loss:.4f}\")\n    print(f\"Val Argmax  -> Acc: {val_acc:.4f} | F1: {val_f1:.4f} | QWK: {val_qwk_argmax:.4f}\")\n    print(f\"Val Tuned   -> Acc: {tuned_acc:.4f} | F1: {tuned_f1:.4f} | QWK: {val_qwk_tuned:.4f}\")\n    print(\"Best epoch thresholds:\", epoch_thresholds)\n\n    scheduler.step(val_qwk_tuned)\n\n    # -------------------------\n    # SAVE BEST\n    # -------------------------\n    improved = val_qwk_tuned > best_qwk\n\n    if improved:\n        best_qwk = val_qwk_tuned\n        best_thresholds = epoch_thresholds\n        early_stop_counter = 0\n\n        torch.save(model.state_dict(), BEST_MODEL_PATH)\n\n        torch.save(\n            {\n                \"epoch\": epoch,\n                \"model_state_dict\": model.state_dict(),\n                \"optimizer_state_dict\": optimizer.state_dict(),\n                \"scheduler_state_dict\": scheduler.state_dict(),\n                \"scaler_state_dict\": scaler.state_dict() if scaler is not None else None,\n                \"best_qwk\": best_qwk,\n                \"best_thresholds\": best_thresholds,\n                \"history\": history,\n                \"early_stop_counter\": early_stop_counter\n            },\n            BEST_FULL_CHECKPOINT_PATH\n        )\n\n        print(\"Best model saved\")\n    else:\n        early_stop_counter += 1\n        print(f\"No improvement. Early stop counter: {early_stop_counter}/{EARLY_STOPPING_PATIENCE}\")\n\n    # -------------------------\n    # SAVE HISTORY + RESUME CKPT\n    # -------------------------\n    epoch_log = {\n        \"epoch\": epoch + 1,\n        \"train_loss\": float(train_loss),\n        \"train_acc\": float(train_acc),\n        \"train_f1\": float(train_f1),\n        \"train_qwk\": float(train_qwk),\n        \"val_loss\": float(val_loss),\n        \"val_argmax_acc\": float(val_acc),\n        \"val_argmax_f1\": float(val_f1),\n        \"val_argmax_qwk\": float(val_qwk_argmax),\n        \"val_tuned_acc\": float(tuned_acc),\n        \"val_tuned_f1\": float(tuned_f1),\n        \"val_tuned_qwk\": float(val_qwk_tuned),\n        \"epoch_thresholds\": epoch_thresholds,\n        \"best_qwk_so_far\": float(best_qwk),\n        \"lr\": float(optimizer.param_groups[0][\"lr\"])\n    }\n    history.append(epoch_log)\n\n    with open(HISTORY_JSON_PATH, \"w\") as f:\n        json.dump(history, f, indent=4)\n\n    torch.save(\n        {\n            \"epoch\": epoch,\n            \"model_state_dict\": model.state_dict(),\n            \"optimizer_state_dict\": optimizer.state_dict(),\n            \"scheduler_state_dict\": scheduler.state_dict(),\n            \"scaler_state_dict\": scaler.state_dict() if scaler is not None else None,\n            \"best_qwk\": best_qwk,\n            \"best_thresholds\": best_thresholds,\n            \"history\": history,\n            \"early_stop_counter\": early_stop_counter\n        },\n        RESUME_CHECKPOINT_PATH\n    )\n\n    print(\"Checkpoint saved\")\n\n    if early_stop_counter >= EARLY_STOPPING_PATIENCE:\n        print(\"Early stopping triggered\")\n        break\n\n# ============================================================\n# FINAL EVALUATION OF BEST MODEL\n# ============================================================\n\nprint(\"\\nLoading best model for final evaluation...\")\n\nbest_model = build_model().to(device)\nbest_state = torch.load(BEST_MODEL_PATH, map_location=device)\nbest_model.load_state_dict(best_state)\nbest_model.eval()\n\nall_labels = []\nall_preds = []\nall_probs = []\n\nwith torch.no_grad():\n    for images, labels in tqdm(val_loader, desc=\"Final evaluation\", leave=False):\n        images = images.to(device, non_blocking=True)\n        outputs = best_model(images)\n\n        probs = torch.softmax(outputs, dim=1)\n        preds = torch.argmax(outputs, dim=1)\n\n        all_probs.append(probs.detach().cpu().numpy())\n        all_preds.extend(preds.detach().cpu().numpy())\n        all_labels.extend(labels.numpy())\n\nall_probs = np.concatenate(all_probs, axis=0)\nall_labels = np.array(all_labels)\nall_preds = np.array(all_preds)\n\nfinal_scores = probs_to_score(all_probs)\nfinal_tuned_preds = apply_thresholds(final_scores, best_thresholds)\n\nfinal_acc = accuracy_score(all_labels, final_tuned_preds)\nfinal_f1 = f1_score(all_labels, final_tuned_preds, average=\"weighted\")\nfinal_qwk = cohen_kappa_score(all_labels, final_tuned_preds, weights=\"quadratic\")\ncm = confusion_matrix(all_labels, final_tuned_preds)\n\nprint(\"\\n==================================================\")\nprint(\"FINAL TUNED RESULTS\")\nprint(\"==================================================\")\nprint(f\"Accuracy     : {final_acc:.6f}\")\nprint(f\"Weighted F1  : {final_f1:.6f}\")\nprint(f\"QWK          : {final_qwk:.6f}\")\nprint(\"Thresholds   :\", best_thresholds)\n\nclass_names = [\"No_DR\", \"Mild\", \"Moderate\", \"Severe\", \"Proliferative_DR\"]\n\nprint(\"\\n==================================================\")\nprint(\"CLASSIFICATION REPORT\")\nprint(\"==================================================\")\nprint(classification_report(\n    all_labels,\n    final_tuned_preds,\n    target_names=class_names,\n    digits=4\n))\n\nprint(\"\\n==================================================\")\nprint(\"CONFUSION MATRIX\")\nprint(\"==================================================\")\nprint(cm)\n\nprint(\"\\n==================================================\")\nprint(\"PER-CLASS ACCURACY\")\nprint(\"==================================================\")\nfor i, class_name in enumerate(class_names):\n    total = cm[i].sum()\n    correct = cm[i, i]\n    class_acc = correct / total if total > 0 else 0.0\n    print(f\"{class_name:18s}: {class_acc:.4f} ({correct}/{total})\")\n\n# ============================================================\n# SAVE OUTPUTS\n# ============================================================\n\npred_df = pd.DataFrame({\n    \"image_path\": val_df[\"image_path\"].values,\n    \"true_label\": all_labels,\n    \"argmax_pred\": all_preds,\n    \"tuned_pred\": final_tuned_preds,\n    \"score_continuous\": final_scores,\n    \"prob_0\": all_probs[:, 0],\n    \"prob_1\": all_probs[:, 1],\n    \"prob_2\": all_probs[:, 2],\n    \"prob_3\": all_probs[:, 3],\n    \"prob_4\": all_probs[:, 4],\n})\npred_df.to_csv(PREDICTIONS_CSV_PATH, index=False)\n\ncm_df = pd.DataFrame(cm, index=class_names, columns=class_names)\ncm_df.to_csv(CONFUSION_CSV_PATH)\n\nsummary = {\n    \"img_size\": IMG_SIZE,\n    \"batch_size\": BATCH_SIZE,\n    \"accumulation_steps\": ACCUMULATION_STEPS,\n    \"epochs_requested\": NUM_EPOCHS,\n    \"best_thresholds\": best_thresholds,\n    \"final_accuracy\": float(final_acc),\n    \"final_weighted_f1\": float(final_f1),\n    \"final_qwk\": float(final_qwk),\n    \"best_model_path\": BEST_MODEL_PATH,\n    \"best_full_checkpoint_path\": BEST_FULL_CHECKPOINT_PATH,\n    \"resume_checkpoint_path\": RESUME_CHECKPOINT_PATH,\n    \"history_json_path\": HISTORY_JSON_PATH\n}\n\nwith open(SUMMARY_JSON_PATH, \"w\") as f:\n    json.dump(summary, f, indent=4)\n\nprint(\"\\nSaved files:\")\nprint(\"Best model:\", BEST_MODEL_PATH)\nprint(\"Best full checkpoint:\", BEST_FULL_CHECKPOINT_PATH)\nprint(\"Resume checkpoint:\", RESUME_CHECKPOINT_PATH)\nprint(\"History:\", HISTORY_JSON_PATH)\nprint(\"Predictions CSV:\", PREDICTIONS_CSV_PATH)\nprint(\"Confusion matrix CSV:\", CONFUSION_CSV_PATH)\nprint(\"Summary JSON:\", SUMMARY_JSON_PATH)\nprint(\"\\nDone.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.544282Z","iopub.status.idle":"2026-05-01T11:21:58.544670Z","shell.execute_reply.started":"2026-05-01T11:21:58.544466Z","shell.execute_reply":"2026-05-01T11:21:58.544490Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\nBASE_PATH = \"/kaggle/input/models/afnanhalim/frozen-best-weighted-model/pytorch/default/1\"\n\nCHECKPOINT_PATH = os.path.join(BASE_PATH, \"weighted_earlystop_checkpoint_final.pth\")\nCLASS_WEIGHTS_PATH = os.path.join(BASE_PATH, \"class_weights.pt\")\nCOMBINED_DF_PATH = os.path.join(BASE_PATH, \"combined_df.csv\")\nTRAIN_DF_PATH = os.path.join(BASE_PATH, \"train_df_frozen.csv\")\nVAL_DF_PATH = os.path.join(BASE_PATH, \"val_df_frozen.csv\")\n\nprint(\"Ready to continue from frozen best model\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.545820Z","iopub.status.idle":"2026-05-01T11:21:58.546086Z","shell.execute_reply.started":"2026-05-01T11:21:58.545965Z","shell.execute_reply":"2026-05-01T11:21:58.545980Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# EVALUATE FROZEN BEST MODEL\n# EfficientNet-B5 | 5-Class DR Classification\n# ============================================================\n\nimport os\nimport cv2\nimport json\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torchvision.transforms as transforms\n\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision.models import efficientnet_b5, EfficientNet_B5_Weights\n\nfrom sklearn.metrics import (\n    accuracy_score,\n    f1_score,\n    cohen_kappa_score,\n    classification_report,\n    confusion_matrix\n)\n\n# ============================================================\n# DEVICE\n# ============================================================\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Using device:\", device)\n\n# ============================================================\n# PATHS\n# ============================================================\n\nBASE_PATH = \"/kaggle/input/models/afnanhalim/frozen-best-weighted-model/pytorch/default/1\"\n\nCHECKPOINT_PATH = os.path.join(BASE_PATH, \"weighted_earlystop_checkpoint_final.pth\")\nVAL_DF_PATH = os.path.join(BASE_PATH, \"val_df_frozen.csv\")\n\nSAVE_DIR = \"/kaggle/working/frozen_best_model_evaluation\"\nos.makedirs(SAVE_DIR, exist_ok=True)\n\nprint(\"Checkpoint path:\", CHECKPOINT_PATH)\nprint(\"Validation CSV:\", VAL_DF_PATH)\n\n# ============================================================\n# LOAD CHECKPOINT + DATA\n# ============================================================\n\ncheckpoint = torch.load(CHECKPOINT_PATH, map_location=device)\n\nprint(\"Checkpoint keys:\", checkpoint.keys())\nprint(\"Saved best QWK:\", checkpoint.get(\"best_qwk\", \"Not found\"))\nprint(\"Saved input size:\", checkpoint.get(\"input_size\", \"Not found\"))\n\nval_df = pd.read_csv(VAL_DF_PATH)\n\nprint(\"Validation samples:\", len(val_df))\nprint(val_df.head())\n\n# ============================================================\n# SETTINGS\n# ============================================================\n\nIMG_SIZE = checkpoint.get(\"input_size\", 456)\n\nif isinstance(IMG_SIZE, (list, tuple)):\n    IMG_SIZE = IMG_SIZE[0]\n\nprint(\"Using image size:\", IMG_SIZE)\n\n# ============================================================\n# PREPROCESSING\n# ============================================================\n\ndef apply_clahe(img):\n    lab = cv2.cvtColor(img, cv2.COLOR_BGR2LAB)\n    l, a, b = cv2.split(lab)\n\n    clahe = cv2.createCLAHE(\n        clipLimit=2.0,\n        tileGridSize=(8, 8)\n    )\n    l = clahe.apply(l)\n\n    lab = cv2.merge((l, a, b))\n    img = cv2.cvtColor(lab, cv2.COLOR_LAB2BGR)\n\n    return img\n\n\nval_transform = transforms.Compose([\n    transforms.ToPILImage(),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225]\n    )\n])\n\n# ============================================================\n# DATASET\n# ============================================================\n\nclass DRDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df.reset_index(drop=True)\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n\n        image_path = row[\"image_path\"]\n        label = int(row[\"label\"])\n\n        img = cv2.imread(image_path)\n\n        if img is None:\n            raise ValueError(f\"Image not found: {image_path}\")\n\n        img = apply_clahe(img)\n        img = cv2.resize(img, (IMG_SIZE, IMG_SIZE))\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n\n        if self.transform:\n            img = self.transform(img)\n\n        return img, label, image_path\n\n\nval_dataset = DRDataset(val_df, transform=val_transform)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=8,\n    shuffle=False,\n    num_workers=0,\n    pin_memory=True\n)\n\nprint(\"Validation batches:\", len(val_loader))\n\n# ============================================================\n# MODEL\n# ============================================================\n\nmodel = efficientnet_b5(weights=EfficientNet_B5_Weights.IMAGENET1K_V1)\n\nin_features = model.classifier[1].in_features\nmodel.classifier[1] = nn.Linear(in_features, 5)\n\nmodel.load_state_dict(checkpoint[\"model_state_dict\"])\n\nmodel = model.to(device)\nmodel.eval()\n\nprint(\"Frozen best model loaded successfully\")\n\n# ============================================================\n# EVALUATION\n# ============================================================\n\nall_labels = []\nall_preds = []\nall_probs = []\nall_paths = []\n\nwith torch.no_grad():\n\n    for images, labels, paths in val_loader:\n\n        images = images.to(device)\n        labels = labels.to(device)\n\n        outputs = model(images)\n\n        probs = torch.softmax(outputs, dim=1)\n        preds = torch.argmax(outputs, dim=1)\n\n        all_labels.extend(labels.cpu().numpy())\n        all_preds.extend(preds.cpu().numpy())\n        all_probs.append(probs.cpu().numpy())\n        all_paths.extend(paths)\n\nall_labels = np.array(all_labels)\nall_preds = np.array(all_preds)\nall_probs = np.concatenate(all_probs, axis=0)\n\n# ============================================================\n# METRICS\n# ============================================================\n\naccuracy = accuracy_score(all_labels, all_preds)\nweighted_f1 = f1_score(all_labels, all_preds, average=\"weighted\")\nqwk = cohen_kappa_score(all_labels, all_preds, weights=\"quadratic\")\n\nclass_names = [\n    \"No_DR\",\n    \"Mild\",\n    \"Moderate\",\n    \"Severe\",\n    \"Proliferative_DR\"\n]\n\ncm = confusion_matrix(all_labels, all_preds)\n\nprint(\"\\n==================================================\")\nprint(\"FROZEN BEST MODEL RESULTS\")\nprint(\"==================================================\")\nprint(f\"Accuracy     : {accuracy:.6f}\")\nprint(f\"Weighted F1  : {weighted_f1:.6f}\")\nprint(f\"QWK          : {qwk:.6f}\")\n\nprint(\"\\n==================================================\")\nprint(\"CLASSIFICATION REPORT\")\nprint(\"==================================================\")\nprint(\n    classification_report(\n        all_labels,\n        all_preds,\n        target_names=class_names,\n        digits=4\n    )\n)\n\nprint(\"\\n==================================================\")\nprint(\"CONFUSION MATRIX\")\nprint(\"==================================================\")\nprint(cm)\n\nprint(\"\\n==================================================\")\nprint(\"PER-CLASS ACCURACY\")\nprint(\"==================================================\")\n\nper_class_accuracy = {}\n\nfor i, class_name in enumerate(class_names):\n    total = cm[i].sum()\n    correct = cm[i, i]\n\n    class_acc = correct / total if total > 0 else 0.0\n    per_class_accuracy[class_name] = class_acc\n\n    print(f\"{class_name:18s}: {class_acc:.4f} ({correct}/{total})\")\n\n# ============================================================\n# SAVE RESULTS\n# ============================================================\n\npred_df = pd.DataFrame({\n    \"image_path\": all_paths,\n    \"true_label\": all_labels,\n    \"pred_label\": all_preds,\n    \"prob_0\": all_probs[:, 0],\n    \"prob_1\": all_probs[:, 1],\n    \"prob_2\": all_probs[:, 2],\n    \"prob_3\": all_probs[:, 3],\n    \"prob_4\": all_probs[:, 4],\n})\n\npred_csv_path = os.path.join(SAVE_DIR, \"frozen_best_predictions.csv\")\npred_df.to_csv(pred_csv_path, index=False)\n\ncm_df = pd.DataFrame(\n    cm,\n    index=class_names,\n    columns=class_names\n)\n\ncm_csv_path = os.path.join(SAVE_DIR, \"confusion_matrix.csv\")\ncm_df.to_csv(cm_csv_path)\n\nsummary = {\n    \"accuracy\": float(accuracy),\n    \"weighted_f1\": float(weighted_f1),\n    \"qwk\": float(qwk),\n    \"saved_best_qwk_in_checkpoint\": float(checkpoint.get(\"best_qwk\", 0.0)),\n    \"image_size\": int(IMG_SIZE),\n    \"class_names\": class_names,\n    \"per_class_accuracy\": {\n        k: float(v) for k, v in per_class_accuracy.items()\n    }\n}\n\nsummary_path = os.path.join(SAVE_DIR, \"evaluation_summary.json\")\n\nwith open(summary_path, \"w\") as f:\n    json.dump(summary, f, indent=4)\n\ntxt_path = os.path.join(SAVE_DIR, \"evaluation_summary.txt\")\n\nwith open(txt_path, \"w\") as f:\n    f.write(\"Frozen Best Model Evaluation Summary\\n\")\n    f.write(\"====================================\\n\\n\")\n    f.write(f\"Accuracy: {accuracy:.6f}\\n\")\n    f.write(f\"Weighted F1: {weighted_f1:.6f}\\n\")\n    f.write(f\"QWK: {qwk:.6f}\\n\")\n    f.write(f\"Saved checkpoint best QWK: {checkpoint.get('best_qwk', 'Not found')}\\n\\n\")\n    f.write(\"Classification Report:\\n\")\n    f.write(\n        classification_report(\n            all_labels,\n            all_preds,\n            target_names=class_names,\n            digits=4\n        )\n    )\n    f.write(\"\\nConfusion Matrix:\\n\")\n    f.write(str(cm))\n    f.write(\"\\n\\nPer-Class Accuracy:\\n\")\n\n    for class_name, class_acc in per_class_accuracy.items():\n        f.write(f\"{class_name}: {class_acc:.4f}\\n\")\n\nprint(\"\\nSaved files:\")\nprint(\"Predictions:\", pred_csv_path)\nprint(\"Confusion matrix:\", cm_csv_path)\nprint(\"Summary JSON:\", summary_path)\nprint(\"Summary TXT:\", txt_path)\nprint(\"\\nDone.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.547314Z","iopub.status.idle":"2026-05-01T11:21:58.547678Z","shell.execute_reply.started":"2026-05-01T11:21:58.547481Z","shell.execute_reply":"2026-05-01T11:21:58.547503Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"apply threshold tuning on our frozen model predictions.","metadata":{}},{"cell_type":"code","source":"# ============================================================\n# THRESHOLD TUNING ON FROZEN BEST MODEL PREDICTIONS\n# ============================================================\n\nimport os\nimport json\nimport numpy as np\nimport pandas as pd\n\nfrom sklearn.metrics import (\n    accuracy_score,\n    f1_score,\n    cohen_kappa_score,\n    classification_report,\n    confusion_matrix\n)\n\nPRED_PATH = \"/kaggle/working/frozen_best_model_evaluation/frozen_best_predictions.csv\"\nSAVE_DIR = \"/kaggle/working/frozen_best_threshold_tuning\"\nos.makedirs(SAVE_DIR, exist_ok=True)\n\ndf = pd.read_csv(PRED_PATH)\n\ny_true = df[\"true_label\"].values\n\nprobs = df[[\"prob_0\", \"prob_1\", \"prob_2\", \"prob_3\", \"prob_4\"]].values\n\ndef probs_to_score(prob_matrix):\n    class_values = np.array([0, 1, 2, 3, 4], dtype=np.float32)\n    return np.sum(prob_matrix * class_values[None, :], axis=1)\n\ndef apply_thresholds(scores, thresholds):\n    preds = np.digitize(scores, bins=thresholds)\n    return np.clip(preds, 0, 4)\n\ndef search_best_thresholds(y_true, probs):\n    scores = probs_to_score(probs)\n\n    grid1 = np.arange(0.30, 0.91, 0.05)\n    grid2 = np.arange(1.00, 1.91, 0.05)\n    grid3 = np.arange(2.00, 2.91, 0.05)\n    grid4 = np.arange(3.00, 3.91, 0.05)\n\n    best_qwk = -1\n    best_thresholds = None\n\n    for t1 in grid1:\n        for t2 in grid2:\n            if t2 <= t1:\n                continue\n            for t3 in grid3:\n                if t3 <= t2:\n                    continue\n                for t4 in grid4:\n                    if t4 <= t3:\n                        continue\n\n                    thresholds = [float(t1), float(t2), float(t3), float(t4)]\n                    preds = apply_thresholds(scores, thresholds)\n                    qwk = cohen_kappa_score(y_true, preds, weights=\"quadratic\")\n\n                    if qwk > best_qwk:\n                        best_qwk = qwk\n                        best_thresholds = thresholds\n\n    return best_thresholds, best_qwk\n\nbest_thresholds, best_qwk = search_best_thresholds(y_true, probs)\n\nscores = probs_to_score(probs)\ntuned_preds = apply_thresholds(scores, best_thresholds)\n\nacc = accuracy_score(y_true, tuned_preds)\nf1 = f1_score(y_true, tuned_preds, average=\"weighted\")\nqwk = cohen_kappa_score(y_true, tuned_preds, weights=\"quadratic\")\ncm = confusion_matrix(y_true, tuned_preds)\n\nclass_names = [\"No_DR\", \"Mild\", \"Moderate\", \"Severe\", \"Proliferative_DR\"]\n\nprint(\"======================================\")\nprint(\"THRESHOLD TUNED RESULTS\")\nprint(\"======================================\")\nprint(\"Best thresholds:\", best_thresholds)\nprint(f\"Accuracy    : {acc:.6f}\")\nprint(f\"Weighted F1 : {f1:.6f}\")\nprint(f\"QWK         : {qwk:.6f}\")\n\nprint(\"\\nClassification Report:\")\nprint(classification_report(y_true, tuned_preds, target_names=class_names, digits=4))\n\nprint(\"\\nConfusion Matrix:\")\nprint(cm)\n\ndf[\"score_continuous\"] = scores\ndf[\"tuned_pred\"] = tuned_preds\n\ndf.to_csv(os.path.join(SAVE_DIR, \"threshold_tuned_predictions.csv\"), index=False)\n\nsummary = {\n    \"best_thresholds\": best_thresholds,\n    \"accuracy\": float(acc),\n    \"weighted_f1\": float(f1),\n    \"qwk\": float(qwk)\n}\n\nwith open(os.path.join(SAVE_DIR, \"threshold_tuning_summary.json\"), \"w\") as f:\n    json.dump(summary, f, indent=4)\n\nprint(\"\\nSaved at:\", SAVE_DIR)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.548872Z","iopub.status.idle":"2026-05-01T11:21:58.549175Z","shell.execute_reply.started":"2026-05-01T11:21:58.549036Z","shell.execute_reply":"2026-05-01T11:21:58.549051Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"We will now continue with the Advanced-Stage Specialist Classifier method.\nStep 1 — Use your existing frozen model predictions\nStep 2 — Select only Moderate/Severe/Proliferative images\nStep 3 — Train a new specialist classifier on these 3 classes\nStep 4 — Combine main model + specialist model","metadata":{}},{"cell_type":"code","source":"# ============================================================\n# SPECIALIST CLASSIFIER FOR ADVANCED DR STAGES\n# Classes: Moderate (2), Severe (3), Proliferative (4)\n# ============================================================\n\nimport os\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torchvision.transforms as transforms\n\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nfrom torchvision.models import efficientnet_b3, EfficientNet_B3_Weights\n\nfrom sklearn.metrics import accuracy_score, cohen_kappa_score\n\nfrom tqdm.auto import tqdm\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Using device:\", device)\n\n# ============================================================\n# PATHS\n# ============================================================\n\nBASE_PATH = \"/kaggle/input/models/afnanhalim/frozen-best-weighted-model/pytorch/default/1\"\n\ntrain_df = pd.read_csv(\n    os.path.join(BASE_PATH, \"train_df_frozen.csv\")\n)\n\nval_df = pd.read_csv(\n    os.path.join(BASE_PATH, \"val_df_frozen.csv\")\n)\n\nSAVE_DIR = \"/kaggle/working/specialist_classifier\"\nos.makedirs(SAVE_DIR, exist_ok=True)\n\n# ============================================================\n# FILTER ONLY ADVANCED CLASSES\n# ============================================================\n\nadvanced_classes = [2, 3, 4]\n\ntrain_df = train_df[\n    train_df[\"label\"].isin(advanced_classes)\n].reset_index(drop=True)\n\nval_df = val_df[\n    val_df[\"label\"].isin(advanced_classes)\n].reset_index(drop=True)\n\nprint(\"Specialist Train size:\", len(train_df))\nprint(\"Specialist Val size:\", len(val_df))\n\n# Remap labels to 0,1,2\n\nlabel_map = {\n    2: 0,\n    3: 1,\n    4: 2\n}\n\ntrain_df[\"label\"] = train_df[\"label\"].map(label_map)\nval_df[\"label\"] = val_df[\"label\"].map(label_map)\n\n# ============================================================\n# PREPROCESS\n# ============================================================\n\nIMG_SIZE = 384\nBATCH_SIZE = 4\nNUM_EPOCHS = 5\n\ndef apply_clahe(img):\n\n    lab = cv2.cvtColor(img, cv2.COLOR_BGR2LAB)\n    l, a, b = cv2.split(lab)\n\n    clahe = cv2.createCLAHE(\n        clipLimit=2.0,\n        tileGridSize=(8, 8)\n    )\n\n    l = clahe.apply(l)\n\n    lab = cv2.merge((l, a, b))\n    img = cv2.cvtColor(lab, cv2.COLOR_LAB2BGR)\n\n    return img\n\ntrain_transform = transforms.Compose([\n    transforms.ToPILImage(),\n    transforms.RandomHorizontalFlip(),\n    transforms.RandomVerticalFlip(),\n    transforms.RandomRotation(15),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean=[0.485,0.456,0.406],\n        std=[0.229,0.224,0.225]\n    )\n])\n\nval_transform = transforms.Compose([\n    transforms.ToPILImage(),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean=[0.485,0.456,0.406],\n        std=[0.229,0.224,0.225]\n    )\n])\n\n# ============================================================\n# DATASET\n# ============================================================\n\nclass DRDataset(Dataset):\n\n    def __init__(self, df, transform=None):\n        self.df = df\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n\n        row = self.df.iloc[idx]\n\n        img = cv2.imread(row[\"image_path\"])\n\n        img = apply_clahe(img)\n\n        img = cv2.resize(img, (IMG_SIZE, IMG_SIZE))\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n\n        label = int(row[\"label\"])\n\n        if self.transform:\n            img = self.transform(img)\n\n        return img, label\n\n# ============================================================\n# CLASS BALANCING\n# ============================================================\n\nclass_counts = train_df[\"label\"].value_counts()\n\nweights = 1.0 / class_counts\n\nsample_weights = train_df[\"label\"].map(weights)\n\nsampler = WeightedRandomSampler(\n    sample_weights,\n    len(sample_weights),\n    replacement=True\n)\n\ntrain_loader = DataLoader(\n    DRDataset(train_df, train_transform),\n    batch_size=BATCH_SIZE,\n    sampler=sampler\n)\n\nval_loader = DataLoader(\n    DRDataset(val_df, val_transform),\n    batch_size=BATCH_SIZE,\n    shuffle=False\n)\n\n# ============================================================\n# MODEL\n# ============================================================\n\nmodel = efficientnet_b3(\n    weights=EfficientNet_B3_Weights.IMAGENET1K_V1\n)\n\nin_features = model.classifier[1].in_features\n\nmodel.classifier[1] = nn.Linear(\n    in_features,\n    3\n)\n\nmodel = model.to(device)\n\ncriterion = nn.CrossEntropyLoss()\n\noptimizer = torch.optim.AdamW(\n    model.parameters(),\n    lr=1e-4\n)\n\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer,\n    mode=\"max\",\n    patience=1\n)\n\nbest_qwk = 0\n\n# ============================================================\n# TRAIN\n# ============================================================\n\nfor epoch in range(NUM_EPOCHS):\n\n    print(\"\\nEpoch\", epoch + 1)\n\n    model.train()\n\n    train_preds = []\n    train_labels = []\n\n    for images, labels in tqdm(train_loader):\n\n        images = images.to(device)\n        labels = labels.to(device)\n\n        outputs = model(images)\n\n        loss = criterion(outputs, labels)\n\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n\n        preds = torch.argmax(outputs, dim=1)\n\n        train_preds.extend(\n            preds.detach().cpu().numpy()\n        )\n\n        train_labels.extend(\n            labels.cpu().numpy()\n        )\n\n    train_qwk = cohen_kappa_score(\n        train_labels,\n        train_preds,\n        weights=\"quadratic\"\n    )\n\n    print(\"Train QWK:\", train_qwk)\n\n    # VALIDATION\n\n    model.eval()\n\n    val_preds = []\n    val_labels = []\n\n    with torch.no_grad():\n\n        for images, labels in val_loader:\n\n            images = images.to(device)\n\n            outputs = model(images)\n\n            preds = torch.argmax(outputs, dim=1)\n\n            val_preds.extend(\n                preds.cpu().numpy()\n            )\n\n            val_labels.extend(\n                labels.numpy()\n            )\n\n    val_qwk = cohen_kappa_score(\n        val_labels,\n        val_preds,\n        weights=\"quadratic\"\n    )\n\n    print(\"Val QWK:\", val_qwk)\n\n    scheduler.step(val_qwk)\n\n    if val_qwk > best_qwk:\n\n        best_qwk = val_qwk\n\n        torch.save(\n            model.state_dict(),\n            os.path.join(\n                SAVE_DIR,\n                \"specialist_best_model.pth\"\n            )\n        )\n\n        print(\"Best specialist model saved\")\n\nprint(\"\\nTraining finished\")\nprint(\"Best Specialist QWK:\", best_qwk)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.550276Z","iopub.status.idle":"2026-05-01T11:21:58.550520Z","shell.execute_reply.started":"2026-05-01T11:21:58.550406Z","shell.execute_reply":"2026-05-01T11:21:58.550421Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Now Main Model + Specialist Refinement (Production System)\n\nThis system will:\n\nRun your best 5-class model\nIf prediction is Moderate / Severe / Proliferative\nSend image to specialist model\nReplace prediction with refined result\n\nThis directly improves:\n\nSevere accuracy\nProliferative accuracy\nOverall 5-class accuracy\nQWK","metadata":{}},{"cell_type":"code","source":"# ============================================================\n# MAIN + SPECIALIST REFINEMENT SYSTEM\n# Final accuracy improvement method\n# ============================================================\n\nimport os\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torchvision.transforms as transforms\n\nfrom torchvision.models import efficientnet_b5, efficientnet_b3\nfrom torchvision.models import (\n    EfficientNet_B5_Weights,\n    EfficientNet_B3_Weights\n)\n\nfrom sklearn.metrics import (\n    accuracy_score,\n    f1_score,\n    cohen_kappa_score,\n    classification_report,\n    confusion_matrix\n)\n\nfrom tqdm.auto import tqdm\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Using device:\", device)\n\n# ============================================================\n# PATHS\n# ============================================================\n\nBASE_PATH = \"/kaggle/input/models/afnanhalim/frozen-best-weighted-model/pytorch/default/1\"\n\nMAIN_MODEL_PATH = os.path.join(\n    BASE_PATH,\n    \"weighted_earlystop_checkpoint_final.pth\"\n)\n\nSPECIALIST_MODEL_PATH = \\\n\"/kaggle/working/specialist_classifier/specialist_best_model.pth\"\n\nVAL_CSV = os.path.join(\n    BASE_PATH,\n    \"val_df_frozen.csv\"\n)\n\nSAVE_DIR = \"/kaggle/working/final_combined_system\"\nos.makedirs(SAVE_DIR, exist_ok=True)\n\n# ============================================================\n# PREPROCESS\n# ============================================================\n\nIMG_SIZE = 384\n\ndef apply_clahe(img):\n\n    lab = cv2.cvtColor(img, cv2.COLOR_BGR2LAB)\n\n    l, a, b = cv2.split(lab)\n\n    clahe = cv2.createCLAHE(\n        clipLimit=2.0,\n        tileGridSize=(8, 8)\n    )\n\n    l = clahe.apply(l)\n\n    lab = cv2.merge((l, a, b))\n\n    img = cv2.cvtColor(lab, cv2.COLOR_LAB2BGR)\n\n    return img\n\ntransform = transforms.Compose([\n    transforms.ToPILImage(),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean=[0.485,0.456,0.406],\n        std=[0.229,0.224,0.225]\n    )\n])\n\n# ============================================================\n# LOAD MAIN MODEL\n# ============================================================\n\nmain_model = efficientnet_b5(\n    weights=EfficientNet_B5_Weights.IMAGENET1K_V1\n)\n\nin_features = main_model.classifier[1].in_features\n\nmain_model.classifier[1] = nn.Linear(\n    in_features,\n    5\n)\n\ncheckpoint = torch.load(\n    MAIN_MODEL_PATH,\n    map_location=device\n)\n\nmain_model.load_state_dict(\n    checkpoint[\"model_state_dict\"]\n)\n\nmain_model = main_model.to(device)\nmain_model.eval()\n\nprint(\"Main model loaded\")\n\n# ============================================================\n# LOAD SPECIALIST MODEL\n# ============================================================\n\nspecialist_model = efficientnet_b3(\n    weights=EfficientNet_B3_Weights.IMAGENET1K_V1\n)\n\nin_features = specialist_model.classifier[1].in_features\n\nspecialist_model.classifier[1] = nn.Linear(\n    in_features,\n    3\n)\n\nspecialist_model.load_state_dict(\n    torch.load(\n        SPECIALIST_MODEL_PATH,\n        map_location=device\n    )\n)\n\nspecialist_model = specialist_model.to(device)\nspecialist_model.eval()\n\nprint(\"Specialist model loaded\")\n\n# ============================================================\n# LOAD DATA\n# ============================================================\n\nval_df = pd.read_csv(VAL_CSV)\n\n# ============================================================\n# PREDICTION\n# ============================================================\n\nall_preds = []\nall_labels = []\n\nfor _, row in tqdm(val_df.iterrows(), total=len(val_df)):\n\n    img = cv2.imread(row[\"image_path\"])\n\n    img = apply_clahe(img)\n\n    img = cv2.resize(\n        img,\n        (IMG_SIZE, IMG_SIZE)\n    )\n\n    img = cv2.cvtColor(\n        img,\n        cv2.COLOR_BGR2RGB\n    )\n\n    tensor = transform(img).unsqueeze(0).to(device)\n\n    # MAIN MODEL\n\n    with torch.no_grad():\n\n        main_output = main_model(tensor)\n\n        main_pred = torch.argmax(\n            main_output,\n            dim=1\n        ).item()\n\n    # SPECIALIST REFINEMENT\n\n    if main_pred in [2,3,4]:\n\n        with torch.no_grad():\n\n            spec_output = specialist_model(tensor)\n\n            spec_pred = torch.argmax(\n                spec_output,\n                dim=1\n            ).item()\n\n        # map back to original labels\n\n        final_pred = spec_pred + 2\n\n    else:\n\n        final_pred = main_pred\n\n    all_preds.append(final_pred)\n\n    all_labels.append(int(row[\"label\"]))\n\n# ============================================================\n# METRICS\n# ============================================================\n\naccuracy = accuracy_score(\n    all_labels,\n    all_preds\n)\n\nf1 = f1_score(\n    all_labels,\n    all_preds,\n    average=\"weighted\"\n)\n\nqwk = cohen_kappa_score(\n    all_labels,\n    all_preds,\n    weights=\"quadratic\"\n)\n\nprint(\"\\nFINAL COMBINED RESULTS\")\n\nprint(\"Accuracy:\", accuracy)\n\nprint(\"Weighted F1:\", f1)\n\nprint(\"QWK:\", qwk)\n\n# ============================================================\n# REPORT\n# ============================================================\n\nclass_names = [\n    \"No_DR\",\n    \"Mild\",\n    \"Moderate\",\n    \"Severe\",\n    \"Proliferative_DR\"\n]\n\nprint(\"\\nClassification Report:\")\n\nprint(\n    classification_report(\n        all_labels,\n        all_preds,\n        target_names=class_names,\n        digits=4\n    )\n)\n\ncm = confusion_matrix(\n    all_labels,\n    all_preds\n)\n\nprint(\"\\nConfusion Matrix:\")\n\nprint(cm)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.554072Z","iopub.status.idle":"2026-05-01T11:21:58.554307Z","shell.execute_reply.started":"2026-05-01T11:21:58.554198Z","shell.execute_reply":"2026-05-01T11:21:58.554212Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# THRESHOLD TUNING ON MAIN + SPECIALIST COMBINED SYSTEM\n# ============================================================\n\nimport os\nimport cv2\nimport json\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torchvision.transforms as transforms\n\nfrom torchvision.models import efficientnet_b5, efficientnet_b3\nfrom torchvision.models import EfficientNet_B5_Weights, EfficientNet_B3_Weights\n\nfrom sklearn.metrics import (\n    accuracy_score,\n    f1_score,\n    cohen_kappa_score,\n    classification_report,\n    confusion_matrix\n)\n\nfrom tqdm.auto import tqdm\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Using device:\", device)\n\n# ============================================================\n# PATHS\n# ============================================================\n\nBASE_PATH = \"/kaggle/input/models/afnanhalim/frozen-best-weighted-model/pytorch/default/1\"\n\nMAIN_MODEL_PATH = os.path.join(\n    BASE_PATH,\n    \"weighted_earlystop_checkpoint_final.pth\"\n)\n\nSPECIALIST_MODEL_PATH = \"/kaggle/working/specialist_classifier/specialist_best_model.pth\"\n\nVAL_CSV = os.path.join(\n    BASE_PATH,\n    \"val_df_frozen.csv\"\n)\n\nSAVE_DIR = \"/kaggle/working/combined_specialist_threshold_tuning\"\nos.makedirs(SAVE_DIR, exist_ok=True)\n\n# ============================================================\n# SETTINGS\n# ============================================================\n\nIMG_SIZE = 384\n\nclass_names = [\n    \"No_DR\",\n    \"Mild\",\n    \"Moderate\",\n    \"Severe\",\n    \"Proliferative_DR\"\n]\n\n# ============================================================\n# PREPROCESS\n# ============================================================\n\ndef apply_clahe(img):\n    lab = cv2.cvtColor(img, cv2.COLOR_BGR2LAB)\n    l, a, b = cv2.split(lab)\n\n    clahe = cv2.createCLAHE(\n        clipLimit=2.0,\n        tileGridSize=(8, 8)\n    )\n\n    l = clahe.apply(l)\n\n    lab = cv2.merge((l, a, b))\n    img = cv2.cvtColor(lab, cv2.COLOR_LAB2BGR)\n\n    return img\n\n\ntransform = transforms.Compose([\n    transforms.ToPILImage(),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225]\n    )\n])\n\n# ============================================================\n# LOAD MAIN MODEL\n# ============================================================\n\nmain_model = efficientnet_b5(\n    weights=EfficientNet_B5_Weights.IMAGENET1K_V1\n)\n\nin_features = main_model.classifier[1].in_features\nmain_model.classifier[1] = nn.Linear(in_features, 5)\n\ncheckpoint = torch.load(\n    MAIN_MODEL_PATH,\n    map_location=device\n)\n\nmain_model.load_state_dict(\n    checkpoint[\"model_state_dict\"]\n)\n\nmain_model = main_model.to(device)\nmain_model.eval()\n\nprint(\"Main model loaded\")\n\n# ============================================================\n# LOAD SPECIALIST MODEL\n# ============================================================\n\nspecialist_model = efficientnet_b3(\n    weights=EfficientNet_B3_Weights.IMAGENET1K_V1\n)\n\nin_features = specialist_model.classifier[1].in_features\nspecialist_model.classifier[1] = nn.Linear(in_features, 3)\n\nspecialist_model.load_state_dict(\n    torch.load(\n        SPECIALIST_MODEL_PATH,\n        map_location=device\n    )\n)\n\nspecialist_model = specialist_model.to(device)\nspecialist_model.eval()\n\nprint(\"Specialist model loaded\")\n\n# ============================================================\n# LOAD VALIDATION DATA\n# ============================================================\n\nval_df = pd.read_csv(VAL_CSV)\n\n# ============================================================\n# COMBINED PROBABILITY PREDICTION\n# ============================================================\n\nall_labels = []\nall_main_preds = []\nall_combined_argmax_preds = []\nall_combined_probs = []\nall_paths = []\n\nwith torch.no_grad():\n\n    for _, row in tqdm(val_df.iterrows(), total=len(val_df)):\n\n        img = cv2.imread(row[\"image_path\"])\n\n        if img is None:\n            raise ValueError(row[\"image_path\"])\n\n        img = apply_clahe(img)\n\n        img = cv2.resize(\n            img,\n            (IMG_SIZE, IMG_SIZE)\n        )\n\n        img = cv2.cvtColor(\n            img,\n            cv2.COLOR_BGR2RGB\n        )\n\n        tensor = transform(img).unsqueeze(0).to(device)\n\n        # -----------------------------\n        # MAIN MODEL\n        # -----------------------------\n        main_output = main_model(tensor)\n        main_probs = torch.softmax(main_output, dim=1).cpu().numpy()[0]\n        main_pred = int(np.argmax(main_probs))\n\n        # -----------------------------\n        # DEFAULT FINAL PROBABILITY\n        # -----------------------------\n        final_probs = main_probs.copy()\n\n        # -----------------------------\n        # SPECIALIST REFINEMENT\n        # only if main model predicts 2/3/4\n        # -----------------------------\n        if main_pred in [2, 3, 4]:\n\n            spec_output = specialist_model(tensor)\n            spec_probs = torch.softmax(spec_output, dim=1).cpu().numpy()[0]\n\n            # keep class 0 and 1 confidence from main model\n            # replace class 2,3,4 distribution using specialist\n            advanced_mass = main_probs[2] + main_probs[3] + main_probs[4]\n\n            final_probs[2] = advanced_mass * spec_probs[0]\n            final_probs[3] = advanced_mass * spec_probs[1]\n            final_probs[4] = advanced_mass * spec_probs[2]\n\n            # normalize again\n            final_probs = final_probs / final_probs.sum()\n\n        final_pred = int(np.argmax(final_probs))\n\n        all_labels.append(int(row[\"label\"]))\n        all_main_preds.append(main_pred)\n        all_combined_argmax_preds.append(final_pred)\n        all_combined_probs.append(final_probs)\n        all_paths.append(row[\"image_path\"])\n\nall_labels = np.array(all_labels)\nall_main_preds = np.array(all_main_preds)\nall_combined_argmax_preds = np.array(all_combined_argmax_preds)\nall_combined_probs = np.array(all_combined_probs)\n\n# ============================================================\n# ARGMAX RESULTS BEFORE THRESHOLD\n# ============================================================\n\nargmax_acc = accuracy_score(\n    all_labels,\n    all_combined_argmax_preds\n)\n\nargmax_f1 = f1_score(\n    all_labels,\n    all_combined_argmax_preds,\n    average=\"weighted\"\n)\n\nargmax_qwk = cohen_kappa_score(\n    all_labels,\n    all_combined_argmax_preds,\n    weights=\"quadratic\"\n)\n\nprint(\"\\n======================================\")\nprint(\"COMBINED SYSTEM BEFORE THRESHOLD\")\nprint(\"======================================\")\nprint(f\"Accuracy    : {argmax_acc:.6f}\")\nprint(f\"Weighted F1 : {argmax_f1:.6f}\")\nprint(f\"QWK         : {argmax_qwk:.6f}\")\n\n# ============================================================\n# THRESHOLD TUNING\n# ============================================================\n\ndef probs_to_score(prob_matrix):\n    class_values = np.array(\n        [0, 1, 2, 3, 4],\n        dtype=np.float32\n    )\n\n    return np.sum(\n        prob_matrix * class_values[None, :],\n        axis=1\n    )\n\n\ndef apply_thresholds(scores, thresholds):\n    preds = np.digitize(\n        scores,\n        bins=thresholds\n    )\n\n    preds = np.clip(\n        preds,\n        0,\n        4\n    )\n\n    return preds\n\n\ndef search_best_thresholds(y_true, probs):\n    scores = probs_to_score(probs)\n\n    grid1 = np.arange(0.30, 0.91, 0.05)\n    grid2 = np.arange(1.00, 1.91, 0.05)\n    grid3 = np.arange(2.00, 2.91, 0.05)\n    grid4 = np.arange(3.00, 3.91, 0.05)\n\n    best_qwk = -1\n    best_thresholds = None\n\n    for t1 in grid1:\n        for t2 in grid2:\n            if t2 <= t1:\n                continue\n\n            for t3 in grid3:\n                if t3 <= t2:\n                    continue\n\n                for t4 in grid4:\n                    if t4 <= t3:\n                        continue\n\n                    thresholds = [\n                        float(t1),\n                        float(t2),\n                        float(t3),\n                        float(t4)\n                    ]\n\n                    preds = apply_thresholds(\n                        scores,\n                        thresholds\n                    )\n\n                    qwk = cohen_kappa_score(\n                        y_true,\n                        preds,\n                        weights=\"quadratic\"\n                    )\n\n                    if qwk > best_qwk:\n                        best_qwk = qwk\n                        best_thresholds = thresholds\n\n    return best_thresholds, best_qwk\n\n\nbest_thresholds, tuned_qwk = search_best_thresholds(\n    all_labels,\n    all_combined_probs\n)\n\nscores = probs_to_score(\n    all_combined_probs\n)\n\ntuned_preds = apply_thresholds(\n    scores,\n    best_thresholds\n)\n\ntuned_acc = accuracy_score(\n    all_labels,\n    tuned_preds\n)\n\ntuned_f1 = f1_score(\n    all_labels,\n    tuned_preds,\n    average=\"weighted\"\n)\n\ntuned_qwk = cohen_kappa_score(\n    all_labels,\n    tuned_preds,\n    weights=\"quadratic\"\n)\n\ncm = confusion_matrix(\n    all_labels,\n    tuned_preds\n)\n\n# ============================================================\n# PRINT FINAL RESULTS\n# ============================================================\n\nprint(\"\\n======================================\")\nprint(\"COMBINED SYSTEM + THRESHOLD TUNING\")\nprint(\"======================================\")\nprint(\"Best thresholds:\", best_thresholds)\nprint(f\"Accuracy    : {tuned_acc:.6f}\")\nprint(f\"Weighted F1 : {tuned_f1:.6f}\")\nprint(f\"QWK         : {tuned_qwk:.6f}\")\n\nprint(\"\\nClassification Report:\")\nprint(\n    classification_report(\n        all_labels,\n        tuned_preds,\n        target_names=class_names,\n        digits=4\n    )\n)\n\nprint(\"\\nConfusion Matrix:\")\nprint(cm)\n\nprint(\"\\nPer-Class Accuracy:\")\n\nfor i, class_name in enumerate(class_names):\n\n    total = cm[i].sum()\n    correct = cm[i, i]\n\n    class_acc = correct / total if total > 0 else 0.0\n\n    print(\n        f\"{class_name:18s}: {class_acc:.4f} ({correct}/{total})\"\n    )\n\n# ============================================================\n# SAVE RESULTS\n# ============================================================\n\npred_df = pd.DataFrame({\n    \"image_path\": all_paths,\n    \"true_label\": all_labels,\n    \"main_pred\": all_main_preds,\n    \"combined_argmax_pred\": all_combined_argmax_preds,\n    \"combined_tuned_pred\": tuned_preds,\n    \"score_continuous\": scores,\n    \"prob_0\": all_combined_probs[:, 0],\n    \"prob_1\": all_combined_probs[:, 1],\n    \"prob_2\": all_combined_probs[:, 2],\n    \"prob_3\": all_combined_probs[:, 3],\n    \"prob_4\": all_combined_probs[:, 4],\n})\n\npred_csv = os.path.join(\n    SAVE_DIR,\n    \"combined_specialist_threshold_predictions.csv\"\n)\n\npred_df.to_csv(\n    pred_csv,\n    index=False\n)\n\ncm_df = pd.DataFrame(\n    cm,\n    index=class_names,\n    columns=class_names\n)\n\ncm_csv = os.path.join(\n    SAVE_DIR,\n    \"combined_specialist_confusion_matrix.csv\"\n)\n\ncm_df.to_csv(cm_csv)\n\nsummary = {\n    \"argmax_accuracy\": float(argmax_acc),\n    \"argmax_weighted_f1\": float(argmax_f1),\n    \"argmax_qwk\": float(argmax_qwk),\n    \"tuned_accuracy\": float(tuned_acc),\n    \"tuned_weighted_f1\": float(tuned_f1),\n    \"tuned_qwk\": float(tuned_qwk),\n    \"best_thresholds\": best_thresholds\n}\n\nsummary_json = os.path.join(\n    SAVE_DIR,\n    \"combined_specialist_threshold_summary.json\"\n)\n\nwith open(summary_json, \"w\") as f:\n    json.dump(\n        summary,\n        f,\n        indent=4\n    )\n\nprint(\"\\nSaved files:\")\nprint(pred_csv)\nprint(cm_csv)\nprint(summary_json)\nprint(\"\\nDone.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.555115Z","iopub.status.idle":"2026-05-01T11:21:58.555456Z","shell.execute_reply.started":"2026-05-01T11:21:58.555277Z","shell.execute_reply":"2026-05-01T11:21:58.555300Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Result Summary\nAccuracy: 0.8301 ✅\nWeighted F1: 0.8272 ✅\nQWK: 0.9150\nWhat Improved Strongly\n\nCompared to my original frozen evaluation, I observed clear improvements:\n\nAccuracy improved:\n0.7787 → 0.8301\nWeighted F1 improved:\n0.7835 → 0.8272\nModerate class performance improved significantly:\n0.5000 → 0.8333\nProliferative DR detection improved:\n0.6479 → 0.7183\n\nThese results show that my threshold tuning method significantly improved overall classification performance, especially for important clinical classes.","metadata":{}},{"cell_type":"markdown","source":"Compared to the original frozen weighted model, the proposed refinement system (specialist model with threshold tuning) achieved a clear improvement in overall classification performance. The baseline frozen model produced about 77.87% accuracy and QWK ≈ 0.8838, whereas the improved system increased the accuracy to 83.01% and QWK to ≈ 0.9150. In particular, the model showed significantly better detection of advanced disease stages (Moderate and Proliferative DR), demonstrating improved clinical grading reliability while maintaining strong overall stability.","metadata":{}},{"cell_type":"code","source":"# ============================================================\n# SEVERE vs PROLIFERATIVE SPECIALIST + COMBINED SYSTEM\n# Goal: improve Proliferative DR and Severe/Proliferative separation\n# ============================================================\n\nimport os\nimport cv2\nimport json\nimport gc\nimport random\nimport numpy as np\nimport pandas as pd\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torchvision.transforms as transforms\n\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\n\nfrom torchvision.models import (\n    efficientnet_b5, EfficientNet_B5_Weights,\n    efficientnet_b3, EfficientNet_B3_Weights\n)\n\nfrom sklearn.metrics import (\n    accuracy_score,\n    f1_score,\n    cohen_kappa_score,\n    classification_report,\n    confusion_matrix\n)\n\nfrom tqdm.auto import tqdm\n\n# ============================================================\n# DEVICE + MEMORY\n# ============================================================\n\nos.environ[\"PYTORCH_CUDA_ALLOC_CONF\"] = \"expandable_segments:True\"\n\ngc.collect()\nif torch.cuda.is_available():\n    torch.cuda.empty_cache()\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Using device:\", device)\n\n# ============================================================\n# SETTINGS\n# ============================================================\n\nSEED = 42\nIMG_SIZE = 384\nBATCH_SIZE = 4\nNUM_EPOCHS = 5\nPATIENCE = 2\nLR = 1e-4\n\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\ntorch.cuda.manual_seed_all(SEED)\n\nBASE_PATH = \"/kaggle/input/models/afnanhalim/frozen-best-weighted-model/pytorch/default/1\"\n\nMAIN_MODEL_PATH = os.path.join(\n    BASE_PATH,\n    \"weighted_earlystop_checkpoint_final.pth\"\n)\n\nTRAIN_CSV = os.path.join(BASE_PATH, \"train_df_frozen.csv\")\nVAL_CSV = os.path.join(BASE_PATH, \"val_df_frozen.csv\")\n\nADVANCED_SPECIALIST_PATH = \"/kaggle/working/specialist_classifier/specialist_best_model.pth\"\n\nSAVE_DIR = \"/kaggle/working/severe_proliferative_specialist\"\nos.makedirs(SAVE_DIR, exist_ok=True)\n\nBINARY_SPECIALIST_PATH = os.path.join(SAVE_DIR, \"severe_proliferative_best_model.pth\")\n\n# ============================================================\n# LOAD DATA\n# ============================================================\n\ntrain_df_full = pd.read_csv(TRAIN_CSV)\nval_df_full = pd.read_csv(VAL_CSV)\n\nprint(\"Full train:\", train_df_full.shape)\nprint(\"Full val:\", val_df_full.shape)\n\n# ============================================================\n# FILTER SEVERE + PROLIFERATIVE ONLY\n# Original labels:\n# Severe = 3\n# Proliferative = 4\n# Binary remap:\n# Severe = 0\n# Proliferative = 1\n# ============================================================\n\nbinary_train_df = train_df_full[\n    train_df_full[\"label\"].isin([3, 4])\n].copy().reset_index(drop=True)\n\nbinary_val_df = val_df_full[\n    val_df_full[\"label\"].isin([3, 4])\n].copy().reset_index(drop=True)\n\nbinary_train_df[\"binary_label\"] = binary_train_df[\"label\"].map({3: 0, 4: 1})\nbinary_val_df[\"binary_label\"] = binary_val_df[\"label\"].map({3: 0, 4: 1})\n\nprint(\"Binary specialist train:\", binary_train_df.shape)\nprint(\"Binary specialist val:\", binary_val_df.shape)\n\nprint(\"\\nBinary train distribution:\")\nprint(binary_train_df[\"binary_label\"].value_counts().sort_index())\n\nprint(\"\\nBinary val distribution:\")\nprint(binary_val_df[\"binary_label\"].value_counts().sort_index())\n\n# ============================================================\n# PREPROCESS\n# ============================================================\n\ndef apply_clahe(img):\n    lab = cv2.cvtColor(img, cv2.COLOR_BGR2LAB)\n    l, a, b = cv2.split(lab)\n\n    clahe = cv2.createCLAHE(\n        clipLimit=2.0,\n        tileGridSize=(8, 8)\n    )\n\n    l = clahe.apply(l)\n    lab = cv2.merge((l, a, b))\n    img = cv2.cvtColor(lab, cv2.COLOR_LAB2BGR)\n\n    return img\n\n\ntrain_transform = transforms.Compose([\n    transforms.ToPILImage(),\n    transforms.RandomHorizontalFlip(p=0.5),\n    transforms.RandomVerticalFlip(p=0.5),\n    transforms.RandomRotation(15),\n    transforms.ColorJitter(\n        brightness=0.15,\n        contrast=0.15,\n        saturation=0.08,\n        hue=0.03\n    ),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225]\n    )\n])\n\nval_transform = transforms.Compose([\n    transforms.ToPILImage(),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225]\n    )\n])\n\n# ============================================================\n# DATASET\n# ============================================================\n\nclass BinaryDRDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df.reset_index(drop=True)\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n\n        img = cv2.imread(row[\"image_path\"])\n\n        if img is None:\n            raise ValueError(row[\"image_path\"])\n\n        img = apply_clahe(img)\n        img = cv2.resize(img, (IMG_SIZE, IMG_SIZE))\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n\n        label = int(row[\"binary_label\"])\n\n        if self.transform:\n            img = self.transform(img)\n\n        return img, label\n\n\n# ============================================================\n# BALANCED SAMPLER FOR BINARY SPECIALIST\n# ============================================================\n\nclass_counts = binary_train_df[\"binary_label\"].value_counts().sort_index()\n\nweights = 1.0 / class_counts.values.astype(np.float32)\nweights = weights / weights.sum()\n\nsample_weights = binary_train_df[\"binary_label\"].map({\n    0: weights[0],\n    1: weights[1]\n}).values\n\nsampler = WeightedRandomSampler(\n    weights=torch.DoubleTensor(sample_weights),\n    num_samples=len(sample_weights),\n    replacement=True\n)\n\nbinary_train_loader = DataLoader(\n    BinaryDRDataset(binary_train_df, train_transform),\n    batch_size=BATCH_SIZE,\n    sampler=sampler,\n    num_workers=2,\n    pin_memory=True\n)\n\nbinary_val_loader = DataLoader(\n    BinaryDRDataset(binary_val_df, val_transform),\n    batch_size=BATCH_SIZE,\n    shuffle=False,\n    num_workers=2,\n    pin_memory=True\n)\n\nprint(\"Binary train batches:\", len(binary_train_loader))\nprint(\"Binary val batches:\", len(binary_val_loader))\n\n# ============================================================\n# BUILD BINARY SPECIALIST MODEL\n# ============================================================\n\ndef build_binary_specialist():\n    model = efficientnet_b3(weights=EfficientNet_B3_Weights.IMAGENET1K_V1)\n    in_features = model.classifier[1].in_features\n    model.classifier[1] = nn.Linear(in_features, 2)\n    return model\n\n\nbinary_model = build_binary_specialist().to(device)\n\ncriterion = nn.CrossEntropyLoss()\n\noptimizer = optim.AdamW(\n    binary_model.parameters(),\n    lr=LR,\n    weight_decay=1e-4\n)\n\nscheduler = optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer,\n    mode=\"max\",\n    factor=0.5,\n    patience=1\n)\n\n# ============================================================\n# TRAIN BINARY SPECIALIST\n# ============================================================\n\nbest_binary_qwk = -1.0\nearly_counter = 0\n\nfor epoch in range(NUM_EPOCHS):\n\n    print(f\"\\nBinary Specialist Epoch {epoch+1}/{NUM_EPOCHS}\")\n\n    binary_model.train()\n\n    train_preds = []\n    train_labels = []\n\n    for images, labels in tqdm(binary_train_loader, leave=False):\n\n        images = images.to(device, non_blocking=True)\n        labels = labels.to(device, non_blocking=True)\n\n        optimizer.zero_grad()\n\n        outputs = binary_model(images)\n        loss = criterion(outputs, labels)\n\n        loss.backward()\n        optimizer.step()\n\n        preds = torch.argmax(outputs, dim=1)\n\n        train_preds.extend(preds.detach().cpu().numpy())\n        train_labels.extend(labels.detach().cpu().numpy())\n\n    train_acc = accuracy_score(train_labels, train_preds)\n    train_qwk = cohen_kappa_score(train_labels, train_preds, weights=\"quadratic\")\n\n    binary_model.eval()\n\n    val_preds = []\n    val_labels = []\n\n    with torch.no_grad():\n        for images, labels in binary_val_loader:\n\n            images = images.to(device, non_blocking=True)\n            labels = labels.to(device, non_blocking=True)\n\n            outputs = binary_model(images)\n            preds = torch.argmax(outputs, dim=1)\n\n            val_preds.extend(preds.detach().cpu().numpy())\n            val_labels.extend(labels.detach().cpu().numpy())\n\n    val_acc = accuracy_score(val_labels, val_preds)\n    val_qwk = cohen_kappa_score(val_labels, val_preds, weights=\"quadratic\")\n\n    print(f\"Train Acc: {train_acc:.4f} | Train QWK: {train_qwk:.4f}\")\n    print(f\"Val   Acc: {val_acc:.4f} | Val   QWK: {val_qwk:.4f}\")\n\n    scheduler.step(val_qwk)\n\n    if val_qwk > best_binary_qwk:\n        best_binary_qwk = val_qwk\n        early_counter = 0\n\n        torch.save(\n            binary_model.state_dict(),\n            BINARY_SPECIALIST_PATH\n        )\n\n        print(\"Best binary specialist saved\")\n\n    else:\n        early_counter += 1\n        print(f\"No improvement: {early_counter}/{PATIENCE}\")\n\n    if early_counter >= PATIENCE:\n        print(\"Early stopping triggered for binary specialist\")\n        break\n\nprint(\"\\nBest Binary Specialist QWK:\", best_binary_qwk)\n\n# ============================================================\n# LOAD MAIN MODEL\n# ============================================================\n\nmain_model = efficientnet_b5(\n    weights=EfficientNet_B5_Weights.IMAGENET1K_V1\n)\n\nin_features = main_model.classifier[1].in_features\nmain_model.classifier[1] = nn.Linear(in_features, 5)\n\nmain_ckpt = torch.load(\n    MAIN_MODEL_PATH,\n    map_location=device\n)\n\nmain_model.load_state_dict(main_ckpt[\"model_state_dict\"])\nmain_model = main_model.to(device)\nmain_model.eval()\n\nprint(\"Main model loaded\")\n\n# ============================================================\n# LOAD ADVANCED 3-CLASS SPECIALIST\n# ============================================================\n\nuse_advanced_specialist = os.path.exists(ADVANCED_SPECIALIST_PATH)\n\nif use_advanced_specialist:\n\n    advanced_model = efficientnet_b3(\n        weights=EfficientNet_B3_Weights.IMAGENET1K_V1\n    )\n\n    in_features = advanced_model.classifier[1].in_features\n    advanced_model.classifier[1] = nn.Linear(in_features, 3)\n\n    advanced_model.load_state_dict(\n        torch.load(\n            ADVANCED_SPECIALIST_PATH,\n            map_location=device\n        )\n    )\n\n    advanced_model = advanced_model.to(device)\n    advanced_model.eval()\n\n    print(\"Advanced 3-class specialist loaded\")\n\nelse:\n    advanced_model = None\n    print(\"Advanced specialist not found, skipping it\")\n\n# ============================================================\n# LOAD BINARY SPECIALIST BEST MODEL\n# ============================================================\n\nbinary_best_model = build_binary_specialist().to(device)\nbinary_best_model.load_state_dict(\n    torch.load(\n        BINARY_SPECIALIST_PATH,\n        map_location=device\n    )\n)\nbinary_best_model.eval()\n\nprint(\"Binary Severe/Proliferative specialist loaded\")\n\n# ============================================================\n# FINAL COMBINED INFERENCE\n# ============================================================\n\nbase_transform = val_transform\n\nval_df = val_df_full.copy().reset_index(drop=True)\n\nall_labels = []\nall_probs = []\nall_argmax_preds = []\nall_paths = []\n\nwith torch.no_grad():\n\n    for _, row in tqdm(val_df.iterrows(), total=len(val_df)):\n\n        img_path = row[\"image_path\"]\n\n        img = cv2.imread(img_path)\n\n        if img is None:\n            raise ValueError(img_path)\n\n        img = apply_clahe(img)\n        img = cv2.resize(img, (IMG_SIZE, IMG_SIZE))\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n\n        tensor = base_transform(img).unsqueeze(0).to(device)\n\n        # -----------------------------\n        # MAIN MODEL PROBS\n        # -----------------------------\n        main_output = main_model(tensor)\n        main_probs = torch.softmax(main_output, dim=1).cpu().numpy()[0]\n\n        final_probs = main_probs.copy()\n        main_pred = int(np.argmax(main_probs))\n\n        # -----------------------------\n        # ADVANCED SPECIALIST REFINEMENT FOR 2/3/4\n        # -----------------------------\n        if use_advanced_specialist and main_pred in [2, 3, 4]:\n\n            adv_output = advanced_model(tensor)\n            adv_probs = torch.softmax(adv_output, dim=1).cpu().numpy()[0]\n\n            advanced_mass = main_probs[2] + main_probs[3] + main_probs[4]\n\n            final_probs[2] = advanced_mass * adv_probs[0]\n            final_probs[3] = advanced_mass * adv_probs[1]\n            final_probs[4] = advanced_mass * adv_probs[2]\n\n            final_probs = final_probs / final_probs.sum()\n\n        # -----------------------------\n        # BINARY SEVERE/PROLIFERATIVE REFINEMENT\n        # Trigger only when advanced mass is meaningful\n        # -----------------------------\n        severe_prolif_mass = final_probs[3] + final_probs[4]\n\n        if severe_prolif_mass >= 0.25:\n\n            bin_output = binary_best_model(tensor)\n            bin_probs = torch.softmax(bin_output, dim=1).cpu().numpy()[0]\n\n            final_probs[3] = severe_prolif_mass * bin_probs[0]\n            final_probs[4] = severe_prolif_mass * bin_probs[1]\n\n            final_probs = final_probs / final_probs.sum()\n\n        final_pred = int(np.argmax(final_probs))\n\n        all_labels.append(int(row[\"label\"]))\n        all_probs.append(final_probs)\n        all_argmax_preds.append(final_pred)\n        all_paths.append(img_path)\n\nall_labels = np.array(all_labels)\nall_probs = np.array(all_probs)\nall_argmax_preds = np.array(all_argmax_preds)\n\n# ============================================================\n# ARGMAX RESULTS\n# ============================================================\n\nargmax_acc = accuracy_score(all_labels, all_argmax_preds)\nargmax_f1 = f1_score(all_labels, all_argmax_preds, average=\"weighted\")\nargmax_qwk = cohen_kappa_score(all_labels, all_argmax_preds, weights=\"quadratic\")\n\nprint(\"\\n======================================\")\nprint(\"COMBINED + BINARY SPECIALIST ARGMAX\")\nprint(\"======================================\")\nprint(f\"Accuracy    : {argmax_acc:.6f}\")\nprint(f\"Weighted F1 : {argmax_f1:.6f}\")\nprint(f\"QWK         : {argmax_qwk:.6f}\")\n\n# ============================================================\n# THRESHOLD TUNING\n# ============================================================\n\ndef probs_to_score(prob_matrix):\n    class_values = np.array([0, 1, 2, 3, 4], dtype=np.float32)\n    return np.sum(prob_matrix * class_values[None, :], axis=1)\n\n\ndef apply_thresholds(scores, thresholds):\n    preds = np.digitize(scores, bins=thresholds)\n    preds = np.clip(preds, 0, 4)\n    return preds\n\n\ndef search_best_thresholds(y_true, probs):\n    scores = probs_to_score(probs)\n\n    grid1 = np.arange(0.30, 0.91, 0.05)\n    grid2 = np.arange(1.00, 1.91, 0.05)\n    grid3 = np.arange(2.00, 2.91, 0.05)\n    grid4 = np.arange(3.00, 3.91, 0.05)\n\n    best_qwk = -1.0\n    best_thresholds = None\n\n    for t1 in grid1:\n        for t2 in grid2:\n            if t2 <= t1:\n                continue\n            for t3 in grid3:\n                if t3 <= t2:\n                    continue\n                for t4 in grid4:\n                    if t4 <= t3:\n                        continue\n\n                    thresholds = [float(t1), float(t2), float(t3), float(t4)]\n                    preds = apply_thresholds(scores, thresholds)\n                    qwk = cohen_kappa_score(y_true, preds, weights=\"quadratic\")\n\n                    if qwk > best_qwk:\n                        best_qwk = qwk\n                        best_thresholds = thresholds\n\n    return best_thresholds, best_qwk\n\n\nbest_thresholds, tuned_qwk = search_best_thresholds(all_labels, all_probs)\n\nscores = probs_to_score(all_probs)\ntuned_preds = apply_thresholds(scores, best_thresholds)\n\ntuned_acc = accuracy_score(all_labels, tuned_preds)\ntuned_f1 = f1_score(all_labels, tuned_preds, average=\"weighted\")\ntuned_qwk = cohen_kappa_score(all_labels, tuned_preds, weights=\"quadratic\")\n\ncm = confusion_matrix(all_labels, tuned_preds)\n\nclass_names = [\n    \"No_DR\",\n    \"Mild\",\n    \"Moderate\",\n    \"Severe\",\n    \"Proliferative_DR\"\n]\n\nprint(\"\\n======================================\")\nprint(\"FINAL SYSTEM + THRESHOLD TUNING\")\nprint(\"======================================\")\nprint(\"Best thresholds:\", best_thresholds)\nprint(f\"Accuracy    : {tuned_acc:.6f}\")\nprint(f\"Weighted F1 : {tuned_f1:.6f}\")\nprint(f\"QWK         : {tuned_qwk:.6f}\")\n\nprint(\"\\nClassification Report:\")\nprint(\n    classification_report(\n        all_labels,\n        tuned_preds,\n        target_names=class_names,\n        digits=4\n    )\n)\n\nprint(\"\\nConfusion Matrix:\")\nprint(cm)\n\nprint(\"\\nPer-Class Accuracy:\")\nfor i, name in enumerate(class_names):\n    total = cm[i].sum()\n    correct = cm[i, i]\n    class_acc = correct / total if total > 0 else 0.0\n    print(f\"{name:18s}: {class_acc:.4f} ({correct}/{total})\")\n\n# ============================================================\n# SAVE RESULTS\n# ============================================================\n\npred_df = pd.DataFrame({\n    \"image_path\": all_paths,\n    \"true_label\": all_labels,\n    \"argmax_pred\": all_argmax_preds,\n    \"tuned_pred\": tuned_preds,\n    \"score_continuous\": scores,\n    \"prob_0\": all_probs[:, 0],\n    \"prob_1\": all_probs[:, 1],\n    \"prob_2\": all_probs[:, 2],\n    \"prob_3\": all_probs[:, 3],\n    \"prob_4\": all_probs[:, 4],\n})\n\npred_csv = os.path.join(SAVE_DIR, \"final_binary_specialist_predictions.csv\")\npred_df.to_csv(pred_csv, index=False)\n\ncm_df = pd.DataFrame(cm, index=class_names, columns=class_names)\ncm_csv = os.path.join(SAVE_DIR, \"final_binary_specialist_confusion_matrix.csv\")\ncm_df.to_csv(cm_csv)\n\nsummary = {\n    \"binary_specialist_best_qwk\": float(best_binary_qwk),\n    \"argmax_accuracy\": float(argmax_acc),\n    \"argmax_weighted_f1\": float(argmax_f1),\n    \"argmax_qwk\": float(argmax_qwk),\n    \"tuned_accuracy\": float(tuned_acc),\n    \"tuned_weighted_f1\": float(tuned_f1),\n    \"tuned_qwk\": float(tuned_qwk),\n    \"best_thresholds\": best_thresholds\n}\n\nsummary_json = os.path.join(SAVE_DIR, \"final_binary_specialist_summary.json\")\n\nwith open(summary_json, \"w\") as f:\n    json.dump(summary, f, indent=4)\n\nprint(\"\\nSaved files:\")\nprint(pred_csv)\nprint(cm_csv)\nprint(summary_json)\nprint(\"\\nDone.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.558992Z","iopub.status.idle":"2026-05-01T11:21:58.559709Z","shell.execute_reply.started":"2026-05-01T11:21:58.559478Z","shell.execute_reply":"2026-05-01T11:21:58.559509Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import shutil\nimport os\nimport json\n\nSAVE_DIR = \"/kaggle/working/final_best_model_package\"\nos.makedirs(SAVE_DIR, exist_ok=True)\n\n# copy main model\nshutil.copy(\n\"/kaggle/input/models/afnanhalim/frozen-best-weighted-model/pytorch/default/1/weighted_earlystop_checkpoint_final.pth\",\nSAVE_DIR\n)\n\n# copy specialist model\nshutil.copy(\n\"/kaggle/working/specialist_classifier/specialist_best_model.pth\",\nSAVE_DIR\n)\n\n# save thresholds\nthresholds = {\n\"best_thresholds\": [0.7, 1.05, 2.85, 3.10],\n\"accuracy\": 0.8301,\n\"qwk\": 0.9150\n}\n\nwith open(\nos.path.join(SAVE_DIR, \"final_thresholds.json\"),\n\"w\"\n) as f:\n    json.dump(thresholds, f, indent=4)\n\n# zip everything\nshutil.make_archive(\n\"/kaggle/working/final_best_model_package\",\n\"zip\",\nSAVE_DIR\n)\n\nprint(\"Saved at:\")\nprint(\"/kaggle/working/final_best_model_package.zip\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.560511Z","iopub.status.idle":"2026-05-01T11:21:58.560993Z","shell.execute_reply.started":"2026-05-01T11:21:58.560860Z","shell.execute_reply":"2026-05-01T11:21:58.560882Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# CBAM + HARD EXAMPLE MINING + ORDINAL-AWARE LOSS\n# Starts from your best frozen EfficientNet-B5 model\n# Goal: improve 5-class accuracy + QWK\n# ============================================================\n\nimport os\nimport cv2\nimport gc\nimport json\nimport random\nimport numpy as np\nimport pandas as pd\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torchvision.transforms as transforms\n\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nfrom torchvision.models import efficientnet_b5, EfficientNet_B5_Weights\n\nfrom sklearn.metrics import (\n    accuracy_score,\n    f1_score,\n    cohen_kappa_score,\n    classification_report,\n    confusion_matrix\n)\n\nfrom tqdm.auto import tqdm\n\n# ============================================================\n# MEMORY + DEVICE\n# ============================================================\n\nos.environ[\"PYTORCH_CUDA_ALLOC_CONF\"] = \"expandable_segments:True\"\n\ngc.collect()\nif torch.cuda.is_available():\n    torch.cuda.empty_cache()\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Using device:\", device)\n\n# ============================================================\n# SETTINGS\n# ============================================================\n\nSEED = 42\nIMG_SIZE = 384\nBATCH_SIZE = 2\nACCUMULATION_STEPS = 2\nNUM_EPOCHS = 6\nPATIENCE = 2\nLR = 2e-5\nNUM_CLASSES = 5\n\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\ntorch.cuda.manual_seed_all(SEED)\n\nBASE_PATH = \"/kaggle/input/models/afnanhalim/frozen-best-weighted-model/pytorch/default/1\"\n\nTRAIN_CSV = os.path.join(BASE_PATH, \"train_df_frozen.csv\")\nVAL_CSV = os.path.join(BASE_PATH, \"val_df_frozen.csv\")\nBEST_CHECKPOINT = os.path.join(BASE_PATH, \"weighted_earlystop_checkpoint_final.pth\")\n\nSAVE_DIR = \"/kaggle/working/cbam_hard_ordinal_training\"\nos.makedirs(SAVE_DIR, exist_ok=True)\n\nBEST_MODEL_PATH = os.path.join(SAVE_DIR, \"best_cbam_hard_ordinal_model.pth\")\nRESUME_PATH = os.path.join(SAVE_DIR, \"resume_checkpoint.pth\")\nHISTORY_PATH = os.path.join(SAVE_DIR, \"training_history.json\")\nSUMMARY_PATH = os.path.join(SAVE_DIR, \"summary.json\")\nPRED_CSV_PATH = os.path.join(SAVE_DIR, \"final_predictions.csv\")\nCM_CSV_PATH = os.path.join(SAVE_DIR, \"confusion_matrix.csv\")\n\n# ============================================================\n# LOAD DATA\n# ============================================================\n\ntrain_df = pd.read_csv(TRAIN_CSV)\nval_df = pd.read_csv(VAL_CSV)\n\nprint(\"Train shape:\", train_df.shape)\nprint(\"Val shape:\", val_df.shape)\n\nprint(\"\\nTrain class distribution:\")\nprint(train_df[\"label\"].value_counts().sort_index())\n\nprint(\"\\nVal class distribution:\")\nprint(val_df[\"label\"].value_counts().sort_index())\n\n# ============================================================\n# PREPROCESSING\n# ============================================================\n\ndef apply_clahe(img):\n    lab = cv2.cvtColor(img, cv2.COLOR_BGR2LAB)\n    l, a, b = cv2.split(lab)\n\n    clahe = cv2.createCLAHE(\n        clipLimit=2.0,\n        tileGridSize=(8, 8)\n    )\n    l = clahe.apply(l)\n\n    lab = cv2.merge((l, a, b))\n    img = cv2.cvtColor(lab, cv2.COLOR_LAB2BGR)\n\n    return img\n\n\ntrain_transform = transforms.Compose([\n    transforms.ToPILImage(),\n    transforms.RandomHorizontalFlip(p=0.5),\n    transforms.RandomVerticalFlip(p=0.5),\n    transforms.RandomRotation(20),\n    transforms.ColorJitter(\n        brightness=0.18,\n        contrast=0.18,\n        saturation=0.08,\n        hue=0.03\n    ),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225]\n    )\n])\n\nval_transform = transforms.Compose([\n    transforms.ToPILImage(),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225]\n    )\n])\n\n# ============================================================\n# DATASET\n# ============================================================\n\nclass DRDataset(Dataset):\n    def __init__(self, df, transform=None, return_path=False):\n        self.df = df.reset_index(drop=True)\n        self.transform = transform\n        self.return_path = return_path\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n\n        image_path = row[\"image_path\"]\n        label = int(row[\"label\"])\n\n        img = cv2.imread(image_path)\n\n        if img is None:\n            raise ValueError(f\"Image not found: {image_path}\")\n\n        img = apply_clahe(img)\n        img = cv2.resize(img, (IMG_SIZE, IMG_SIZE))\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n\n        if self.transform:\n            img = self.transform(img)\n\n        if self.return_path:\n            return img, label, image_path\n\n        return img, label\n\n# ============================================================\n# CBAM MODULE\n# ============================================================\n\nclass ChannelAttention(nn.Module):\n    def __init__(self, channels, reduction=16):\n        super().__init__()\n\n        hidden = max(channels // reduction, 8)\n\n        self.avg_pool = nn.AdaptiveAvgPool2d(1)\n        self.max_pool = nn.AdaptiveMaxPool2d(1)\n\n        self.mlp = nn.Sequential(\n            nn.Conv2d(channels, hidden, kernel_size=1, bias=False),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(hidden, channels, kernel_size=1, bias=False)\n        )\n\n        self.sigmoid = nn.Sigmoid()\n\n    def forward(self, x):\n        avg_out = self.mlp(self.avg_pool(x))\n        max_out = self.mlp(self.max_pool(x))\n\n        out = avg_out + max_out\n        return self.sigmoid(out)\n\n\nclass SpatialAttention(nn.Module):\n    def __init__(self, kernel_size=7):\n        super().__init__()\n\n        padding = kernel_size // 2\n\n        self.conv = nn.Conv2d(\n            2,\n            1,\n            kernel_size=kernel_size,\n            padding=padding,\n            bias=False\n        )\n\n        self.sigmoid = nn.Sigmoid()\n\n    def forward(self, x):\n        avg_out = torch.mean(x, dim=1, keepdim=True)\n        max_out, _ = torch.max(x, dim=1, keepdim=True)\n\n        x_cat = torch.cat([avg_out, max_out], dim=1)\n        return self.sigmoid(self.conv(x_cat))\n\n\nclass CBAM(nn.Module):\n    def __init__(self, channels, reduction=16, kernel_size=7):\n        super().__init__()\n\n        self.channel_attention = ChannelAttention(\n            channels,\n            reduction=reduction\n        )\n\n        self.spatial_attention = SpatialAttention(\n            kernel_size=kernel_size\n        )\n\n    def forward(self, x):\n        x = x * self.channel_attention(x)\n        x = x * self.spatial_attention(x)\n        return x\n\n\ndef get_stage_out_channels(stage):\n    \"\"\"\n    Finds output channel number of a torchvision EfficientNet stage.\n    Works by checking last Conv2d layer inside the stage.\n    \"\"\"\n    out_channels = None\n\n    for module in stage.modules():\n        if isinstance(module, nn.Conv2d):\n            out_channels = module.out_channels\n\n    if out_channels is None:\n        raise ValueError(\"Could not find output channels for stage\")\n\n    return out_channels\n\n# ============================================================\n# BUILD CBAM-EFFICIENTNETB5\n# ============================================================\n\ndef build_cbam_efficientnet_b5():\n    model = efficientnet_b5(weights=EfficientNet_B5_Weights.IMAGENET1K_V1)\n\n    in_features = model.classifier[1].in_features\n    model.classifier[1] = nn.Linear(in_features, NUM_CLASSES)\n\n    # Add CBAM after deeper EfficientNet stages\n    # stages 4,5,6,7 are deeper semantic layers\n    cbam_stage_indices = [4, 5, 6, 7]\n\n    for idx in cbam_stage_indices:\n        stage = model.features[idx]\n        channels = get_stage_out_channels(stage)\n\n        model.features[idx] = nn.Sequential(\n            stage,\n            CBAM(channels)\n        )\n\n        print(f\"CBAM added after EfficientNet feature stage {idx} with {channels} channels\")\n\n    return model\n\n# ============================================================\n# LOAD ORIGINAL BEST MODEL FOR HARD EXAMPLE MINING\n# ============================================================\n\ndef build_plain_efficientnet_b5():\n    model = efficientnet_b5(weights=EfficientNet_B5_Weights.IMAGENET1K_V1)\n\n    in_features = model.classifier[1].in_features\n    model.classifier[1] = nn.Linear(in_features, NUM_CLASSES)\n\n    return model\n\n\nplain_model = build_plain_efficientnet_b5()\n\ncheckpoint = torch.load(\n    BEST_CHECKPOINT,\n    map_location=device\n)\n\nplain_model.load_state_dict(\n    checkpoint[\"model_state_dict\"]\n)\n\nplain_model = plain_model.to(device)\nplain_model.eval()\n\nprint(\"Plain frozen best model loaded for hard-example mining\")\n\n# ============================================================\n# HARD EXAMPLE MINING ON TRAIN SET\n# ============================================================\n\nprint(\"\\nMining hard examples from training set...\")\n\nhard_loader = DataLoader(\n    DRDataset(train_df, transform=val_transform, return_path=True),\n    batch_size=8,\n    shuffle=False,\n    num_workers=2,\n    pin_memory=True\n)\n\nhard_image_paths = []\nhard_true_labels = []\nhard_pred_labels = []\n\nwith torch.no_grad():\n    for images, labels, paths in tqdm(hard_loader, desc=\"Hard mining\"):\n        images = images.to(device)\n        labels = labels.to(device)\n\n        outputs = plain_model(images)\n        preds = torch.argmax(outputs, dim=1)\n\n        labels_np = labels.cpu().numpy()\n        preds_np = preds.cpu().numpy()\n\n        for pth, true_l, pred_l in zip(paths, labels_np, preds_np):\n            if int(true_l) != int(pred_l):\n                hard_image_paths.append(pth)\n                hard_true_labels.append(int(true_l))\n                hard_pred_labels.append(int(pred_l))\n\nhard_df = pd.DataFrame({\n    \"image_path\": hard_image_paths,\n    \"true_label\": hard_true_labels,\n    \"pred_label\": hard_pred_labels\n})\n\nhard_csv = os.path.join(SAVE_DIR, \"hard_examples_from_train.csv\")\nhard_df.to_csv(hard_csv, index=False)\n\nprint(\"Hard examples found:\", len(hard_df))\nprint(\"Saved hard examples:\", hard_csv)\n\nhard_set = set(hard_image_paths)\n\n# ============================================================\n# HARD EXAMPLE + CLASS BALANCED SAMPLER\n# ============================================================\n\nclass_counts = train_df[\"label\"].value_counts().sort_index()\ninv_class = 1.0 / class_counts.values.astype(np.float32)\n\nbase_class_weight_map = {\n    cls: inv_class[i]\n    for i, cls in enumerate(class_counts.index)\n}\n\nsample_weights = []\n\nfor _, row in train_df.iterrows():\n    label = int(row[\"label\"])\n    image_path = row[\"image_path\"]\n\n    w = base_class_weight_map[label]\n\n    # Boost weak classes\n    if label in [1, 3, 4]:\n        w *= 1.7\n\n    # Boost hard examples\n    if image_path in hard_set:\n        w *= 2.5\n\n    sample_weights.append(float(w))\n\nsample_weights = np.array(sample_weights)\nsample_weights = sample_weights / sample_weights.mean()\n\nsampler = WeightedRandomSampler(\n    weights=torch.DoubleTensor(sample_weights),\n    num_samples=len(sample_weights),\n    replacement=True\n)\n\ntrain_loader = DataLoader(\n    DRDataset(train_df, transform=train_transform),\n    batch_size=BATCH_SIZE,\n    sampler=sampler,\n    num_workers=2,\n    pin_memory=True\n)\n\nval_loader = DataLoader(\n    DRDataset(val_df, transform=val_transform, return_path=True),\n    batch_size=8,\n    shuffle=False,\n    num_workers=2,\n    pin_memory=True\n)\n\nprint(\"Train batches:\", len(train_loader))\nprint(\"Val batches:\", len(val_loader))\n\n# ============================================================\n# BUILD CBAM MODEL AND LOAD BEST WEIGHTS\n# ============================================================\n\nmodel = build_cbam_efficientnet_b5()\n\n# Load existing best model weights into CBAM model.\n# CBAM layers will remain newly initialized.\nmissing, unexpected = model.load_state_dict(\n    checkpoint[\"model_state_dict\"],\n    strict=False\n)\n\nprint(\"\\nLoaded best model into CBAM model\")\nprint(\"Missing keys count:\", len(missing))\nprint(\"Unexpected keys count:\", len(unexpected))\n\nmodel = model.to(device)\n\n# Freeze early layers to protect learned knowledge\nfor name, param in model.named_parameters():\n    param.requires_grad = True\n\n# Optional: freeze first few stages\nfor idx in [0, 1, 2]:\n    for param in model.features[idx].parameters():\n        param.requires_grad = False\n\nprint(\"Early stages frozen: 0,1,2\")\n\n# ============================================================\n# ORDINAL-AWARE LOSS\n# CE + expected grade distance penalty\n# ============================================================\n\nclass_counts = train_df[\"label\"].value_counts().sort_index()\ninv_freq = 1.0 / class_counts.values.astype(np.float32)\nclass_weights = inv_freq / inv_freq.sum()\nclass_weights_tensor = torch.tensor(\n    class_weights,\n    dtype=torch.float32,\n    device=device\n)\n\nclass OrdinalAwareLoss(nn.Module):\n    def __init__(self, class_weights=None, ordinal_lambda=0.35):\n        super().__init__()\n\n        self.ce = nn.CrossEntropyLoss(weight=class_weights)\n        self.ordinal_lambda = ordinal_lambda\n\n        self.register_buffer(\n            \"class_values\",\n            torch.arange(NUM_CLASSES, dtype=torch.float32)\n        )\n\n    def forward(self, logits, targets):\n        ce_loss = self.ce(logits, targets)\n\n        probs = torch.softmax(logits, dim=1)\n        expected_grade = torch.sum(\n            probs * self.class_values.to(logits.device),\n            dim=1\n        )\n\n        ordinal_loss = torch.mean(\n            torch.abs(expected_grade - targets.float())\n        )\n\n        total_loss = ce_loss + self.ordinal_lambda * ordinal_loss\n\n        return total_loss\n\n\ncriterion = OrdinalAwareLoss(\n    class_weights=class_weights_tensor,\n    ordinal_lambda=0.35\n)\n\noptimizer = optim.AdamW(\n    filter(lambda p: p.requires_grad, model.parameters()),\n    lr=LR,\n    weight_decay=1e-4\n)\n\nscheduler = optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer,\n    mode=\"max\",\n    factor=0.5,\n    patience=1\n)\n\nfrom torch.amp import autocast, GradScaler\nscaler = GradScaler(\"cuda\") if torch.cuda.is_available() else None\n\n# ============================================================\n# THRESHOLD TUNING FUNCTIONS\n# ============================================================\n\ndef probs_to_score(prob_matrix):\n    class_values_np = np.array([0, 1, 2, 3, 4], dtype=np.float32)\n    return np.sum(prob_matrix * class_values_np[None, :], axis=1)\n\n\ndef apply_thresholds(scores, thresholds):\n    preds = np.digitize(scores, bins=thresholds)\n    preds = np.clip(preds, 0, 4)\n    return preds\n\n\ndef search_best_thresholds(y_true, probs):\n    scores = probs_to_score(probs)\n\n    grid1 = np.arange(0.35, 0.91, 0.05)\n    grid2 = np.arange(1.00, 1.91, 0.05)\n    grid3 = np.arange(2.00, 2.91, 0.05)\n    grid4 = np.arange(3.00, 3.91, 0.05)\n\n    best_qwk = -1.0\n    best_thresholds = [0.5, 1.5, 2.5, 3.5]\n\n    for t1 in grid1:\n        for t2 in grid2:\n            if t2 <= t1:\n                continue\n\n            for t3 in grid3:\n                if t3 <= t2:\n                    continue\n\n                for t4 in grid4:\n                    if t4 <= t3:\n                        continue\n\n                    thresholds = [\n                        float(t1),\n                        float(t2),\n                        float(t3),\n                        float(t4)\n                    ]\n\n                    preds = apply_thresholds(scores, thresholds)\n\n                    qwk = cohen_kappa_score(\n                        y_true,\n                        preds,\n                        weights=\"quadratic\"\n                    )\n\n                    if qwk > best_qwk:\n                        best_qwk = qwk\n                        best_thresholds = thresholds\n\n    return best_thresholds, best_qwk\n\n# ============================================================\n# METRIC FUNCTION\n# ============================================================\n\ndef compute_metrics(y_true, y_pred):\n    acc = accuracy_score(y_true, y_pred)\n    f1 = f1_score(y_true, y_pred, average=\"weighted\")\n    qwk = cohen_kappa_score(y_true, y_pred, weights=\"quadratic\")\n    return acc, f1, qwk\n\n# ============================================================\n# TRAINING LOOP\n# ============================================================\n\nbest_qwk = -1.0\nbest_thresholds = [0.5, 1.5, 2.5, 3.5]\nearly_counter = 0\nhistory = []\n\nfor epoch in range(NUM_EPOCHS):\n\n    print(f\"\\nEpoch {epoch+1}/{NUM_EPOCHS}\")\n\n    model.train()\n\n    train_preds = []\n    train_labels = []\n    train_loss_total = 0.0\n\n    optimizer.zero_grad()\n\n    train_bar = tqdm(train_loader, desc=f\"Train Epoch {epoch+1}\", leave=False)\n\n    for step, (images, labels) in enumerate(train_bar):\n\n        images = images.to(device, non_blocking=True)\n        labels = labels.to(device, non_blocking=True)\n\n        if scaler is not None:\n            with autocast(\"cuda\"):\n                outputs = model(images)\n                loss = criterion(outputs, labels)\n                loss = loss / ACCUMULATION_STEPS\n\n            scaler.scale(loss).backward()\n\n            if (step + 1) % ACCUMULATION_STEPS == 0:\n                scaler.step(optimizer)\n                scaler.update()\n                optimizer.zero_grad()\n\n        else:\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n            loss = loss / ACCUMULATION_STEPS\n            loss.backward()\n\n            if (step + 1) % ACCUMULATION_STEPS == 0:\n                optimizer.step()\n                optimizer.zero_grad()\n\n        train_loss_total += loss.item() * ACCUMULATION_STEPS * images.size(0)\n\n        preds = torch.argmax(outputs, dim=1)\n\n        train_preds.extend(preds.detach().cpu().numpy())\n        train_labels.extend(labels.detach().cpu().numpy())\n\n    train_loss = train_loss_total / len(train_loader.dataset)\n    train_acc, train_f1, train_qwk = compute_metrics(train_labels, train_preds)\n\n    print(f\"Train Loss: {train_loss:.4f}\")\n    print(f\"Train Acc: {train_acc:.4f} | Train F1: {train_f1:.4f} | Train QWK: {train_qwk:.4f}\")\n\n    # ========================================================\n    # VALIDATION\n    # ========================================================\n\n    model.eval()\n\n    val_labels = []\n    val_preds = []\n    val_probs = []\n    val_paths = []\n\n    val_loss_total = 0.0\n\n    with torch.no_grad():\n\n        for images, labels, paths in tqdm(val_loader, desc=f\"Valid Epoch {epoch+1}\", leave=False):\n\n            images = images.to(device, non_blocking=True)\n            labels = labels.to(device, non_blocking=True)\n\n            if scaler is not None:\n                with autocast(\"cuda\"):\n                    outputs = model(images)\n                    loss = criterion(outputs, labels)\n            else:\n                outputs = model(images)\n                loss = criterion(outputs, labels)\n\n            val_loss_total += loss.item() * images.size(0)\n\n            probs = torch.softmax(outputs, dim=1)\n            preds = torch.argmax(outputs, dim=1)\n\n            val_probs.append(probs.detach().cpu().numpy())\n            val_preds.extend(preds.detach().cpu().numpy())\n            val_labels.extend(labels.detach().cpu().numpy())\n            val_paths.extend(paths)\n\n    val_loss = val_loss_total / len(val_loader.dataset)\n    val_probs = np.concatenate(val_probs, axis=0)\n    val_labels = np.array(val_labels)\n    val_preds = np.array(val_preds)\n\n    val_acc, val_f1, val_qwk_argmax = compute_metrics(val_labels, val_preds)\n\n    epoch_thresholds, val_qwk_tuned = search_best_thresholds(\n        val_labels,\n        val_probs\n    )\n\n    scores = probs_to_score(val_probs)\n    tuned_preds = apply_thresholds(scores, epoch_thresholds)\n\n    tuned_acc, tuned_f1, tuned_qwk = compute_metrics(\n        val_labels,\n        tuned_preds\n    )\n\n    print(f\"Val Loss: {val_loss:.4f}\")\n    print(f\"Val Argmax -> Acc: {val_acc:.4f} | F1: {val_f1:.4f} | QWK: {val_qwk_argmax:.4f}\")\n    print(f\"Val Tuned  -> Acc: {tuned_acc:.4f} | F1: {tuned_f1:.4f} | QWK: {tuned_qwk:.4f}\")\n    print(\"Epoch thresholds:\", epoch_thresholds)\n\n    scheduler.step(tuned_qwk)\n\n    improved = tuned_qwk > best_qwk\n\n    if improved:\n        best_qwk = tuned_qwk\n        best_thresholds = epoch_thresholds\n        early_counter = 0\n\n        torch.save(\n            model.state_dict(),\n            BEST_MODEL_PATH\n        )\n\n        print(\"Best CBAM hard ordinal model saved\")\n\n    else:\n        early_counter += 1\n        print(f\"No improvement: {early_counter}/{PATIENCE}\")\n\n    epoch_log = {\n        \"epoch\": epoch + 1,\n        \"train_loss\": float(train_loss),\n        \"train_acc\": float(train_acc),\n        \"train_f1\": float(train_f1),\n        \"train_qwk\": float(train_qwk),\n        \"val_loss\": float(val_loss),\n        \"val_argmax_acc\": float(val_acc),\n        \"val_argmax_f1\": float(val_f1),\n        \"val_argmax_qwk\": float(val_qwk_argmax),\n        \"val_tuned_acc\": float(tuned_acc),\n        \"val_tuned_f1\": float(tuned_f1),\n        \"val_tuned_qwk\": float(tuned_qwk),\n        \"thresholds\": epoch_thresholds,\n        \"best_qwk_so_far\": float(best_qwk),\n        \"lr\": float(optimizer.param_groups[0][\"lr\"])\n    }\n\n    history.append(epoch_log)\n\n    with open(HISTORY_PATH, \"w\") as f:\n        json.dump(history, f, indent=4)\n\n    torch.save(\n        {\n            \"epoch\": epoch,\n            \"model_state_dict\": model.state_dict(),\n            \"optimizer_state_dict\": optimizer.state_dict(),\n            \"scheduler_state_dict\": scheduler.state_dict(),\n            \"scaler_state_dict\": scaler.state_dict() if scaler is not None else None,\n            \"best_qwk\": best_qwk,\n            \"best_thresholds\": best_thresholds,\n            \"history\": history\n        },\n        RESUME_PATH\n    )\n\n    print(\"Checkpoint saved\")\n\n    if early_counter >= PATIENCE:\n        print(\"Early stopping triggered\")\n        break\n\n# ============================================================\n# FINAL EVALUATION USING BEST MODEL\n# ============================================================\n\nprint(\"\\nLoading best model for final evaluation...\")\n\nbest_model = build_cbam_efficientnet_b5().to(device)\nbest_model.load_state_dict(\n    torch.load(BEST_MODEL_PATH, map_location=device)\n)\nbest_model.eval()\n\nfinal_labels = []\nfinal_probs = []\nfinal_preds = []\nfinal_paths = []\n\nwith torch.no_grad():\n\n    for images, labels, paths in tqdm(val_loader, desc=\"Final evaluation\", leave=False):\n\n        images = images.to(device, non_blocking=True)\n\n        outputs = best_model(images)\n        probs = torch.softmax(outputs, dim=1)\n\n        preds = torch.argmax(probs, dim=1)\n\n        final_probs.append(probs.cpu().numpy())\n        final_preds.extend(preds.cpu().numpy())\n        final_labels.extend(labels.numpy())\n        final_paths.extend(paths)\n\nfinal_probs = np.concatenate(final_probs, axis=0)\nfinal_labels = np.array(final_labels)\nfinal_preds = np.array(final_preds)\n\nfinal_scores = probs_to_score(final_probs)\nfinal_tuned_preds = apply_thresholds(final_scores, best_thresholds)\n\nfinal_acc, final_f1, final_qwk = compute_metrics(\n    final_labels,\n    final_tuned_preds\n)\n\ncm = confusion_matrix(final_labels, final_tuned_preds)\n\nclass_names = [\n    \"No_DR\",\n    \"Mild\",\n    \"Moderate\",\n    \"Severe\",\n    \"Proliferative_DR\"\n]\n\nprint(\"\\n==================================================\")\nprint(\"FINAL CBAM + HARD EXAMPLE + ORDINAL RESULTS\")\nprint(\"==================================================\")\nprint(f\"Accuracy    : {final_acc:.6f}\")\nprint(f\"Weighted F1 : {final_f1:.6f}\")\nprint(f\"QWK         : {final_qwk:.6f}\")\nprint(\"Thresholds :\", best_thresholds)\n\nprint(\"\\nClassification Report:\")\nprint(\n    classification_report(\n        final_labels,\n        final_tuned_preds,\n        target_names=class_names,\n        digits=4\n    )\n)\n\nprint(\"\\nConfusion Matrix:\")\nprint(cm)\n\nprint(\"\\nPer-Class Accuracy:\")\nper_class_accuracy = {}\n\nfor i, name in enumerate(class_names):\n    total = cm[i].sum()\n    correct = cm[i, i]\n    class_acc = correct / total if total > 0 else 0.0\n    per_class_accuracy[name] = class_acc\n    print(f\"{name:18s}: {class_acc:.4f} ({correct}/{total})\")\n\n# ============================================================\n# SAVE FINAL OUTPUTS\n# ============================================================\n\npred_df = pd.DataFrame({\n    \"image_path\": final_paths,\n    \"true_label\": final_labels,\n    \"argmax_pred\": final_preds,\n    \"tuned_pred\": final_tuned_preds,\n    \"score\": final_scores,\n    \"prob_0\": final_probs[:, 0],\n    \"prob_1\": final_probs[:, 1],\n    \"prob_2\": final_probs[:, 2],\n    \"prob_3\": final_probs[:, 3],\n    \"prob_4\": final_probs[:, 4],\n})\n\npred_df.to_csv(PRED_CSV_PATH, index=False)\n\ncm_df = pd.DataFrame(\n    cm,\n    index=class_names,\n    columns=class_names\n)\ncm_df.to_csv(CM_CSV_PATH)\n\nsummary = {\n    \"method\": \"CBAM + hard example mining + ordinal-aware loss\",\n    \"accuracy\": float(final_acc),\n    \"weighted_f1\": float(final_f1),\n    \"qwk\": float(final_qwk),\n    \"thresholds\": best_thresholds,\n    \"per_class_accuracy\": {\n        k: float(v) for k, v in per_class_accuracy.items()\n    },\n    \"best_model_path\": BEST_MODEL_PATH,\n    \"history_path\": HISTORY_PATH,\n    \"predictions_path\": PRED_CSV_PATH,\n    \"confusion_matrix_path\": CM_CSV_PATH\n}\n\nwith open(SUMMARY_PATH, \"w\") as f:\n    json.dump(summary, f, indent=4)\n\nprint(\"\\nSaved files:\")\nprint(\"Best model:\", BEST_MODEL_PATH)\nprint(\"Resume checkpoint:\", RESUME_PATH)\nprint(\"History:\", HISTORY_PATH)\nprint(\"Predictions:\", PRED_CSV_PATH)\nprint(\"Confusion matrix:\", CM_CSV_PATH)\nprint(\"Summary:\", SUMMARY_PATH)\nprint(\"\\nDone.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-01T11:21:58.561842Z","iopub.status.idle":"2026-05-01T11:21:58.562144Z","shell.execute_reply.started":"2026-05-01T11:21:58.562019Z","shell.execute_reply":"2026-05-01T11:21:58.562037Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"This implementation follows my original research roadmap and introduces three key improvements to strengthen model performance and clinical reliability.\n\nFirst, I integrate CBAM attention into the deeper stages of EfficientNet-B5, enabling the model to focus more effectively on important retinal regions and lesion-related features. This improves the model’s ability to capture meaningful visual patterns associated with disease severity.\n\nSecond, I apply hard example mining by identifying images that my best frozen model previously misclassified. These difficult samples are then assigned higher sampling priority during fine-tuning, allowing the model to learn more effectively from challenging cases.\n\nThird, I implement an ordinal-aware loss function, which reflects the ordered nature of diabetic retinopathy (DR) grades from 0 to 4. This encourages the model to learn the natural progression of disease severity rather than treating each class as completely independent.","metadata":{}}]}