{"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,"isSourceIdPinned":false}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# ==========================================================\n# Early Detection & Severity Classification of DR\n# ==========================================================\n\nimport os, cv2, numpy as np, pandas as pd, matplotlib.pyplot as plt, seaborn as sns, tensorflow as tf\nfrom sklearn.model_selection import StratifiedKFold\nfrom sklearn.metrics import confusion_matrix, classification_report, roc_auc_score, roc_curve, accuracy_score\nfrom sklearn.utils.class_weight import compute_class_weight\nfrom tensorflow.keras import layers, models\nfrom tensorflow.keras.callbacks import ModelCheckpoint, EarlyStopping, ReduceLROnPlateau\nfrom tensorflow.keras.utils import to_categorical\nfrom tqdm import tqdm\n\n# ----------------------------------------------------------\n# GPU CHECK\n# ----------------------------------------------------------\nprint(\"TensorFlow:\", tf.__version__)\nprint(\"GPUs:\", tf.config.list_physical_devices('GPU'))\n\n# Multi-GPU Strategy\nstrategy = tf.distribute.MirroredStrategy()\nprint(\"Number of devices:\", strategy.num_replicas_in_sync)\n\n# ----------------------------------------------------------\n# PATHS\n# ----------------------------------------------------------\nCSV_PATH = \"/kaggle/input/competitions/aptos2019-blindness-detection/train.csv\"\nIMG_DIR  = \"/kaggle/input/competitions/aptos2019-blindness-detection/train_images\"\n\nIMG_SIZE = 299\nSEED = 42\nnp.random.seed(SEED)\ntf.random.set_seed(SEED)\n\n# ----------------------------------------------------------\n# LOAD DATA\n# ----------------------------------------------------------\ndf = pd.read_csv(CSV_PATH)\ndf[\"path\"] = df[\"id_code\"].apply(lambda x: f\"{IMG_DIR}/{x}.png\")\n\nprint(df.head())\nprint(df[\"diagnosis\"].value_counts())\n\n# ----------------------------------------------------------\n# PREPROCESSING\n# ----------------------------------------------------------\ndef crop_fundus(img):\n    gray = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY)\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        c = max(contours, key=cv2.contourArea)\n        x, y, w, h = cv2.boundingRect(c)\n        img = img[y:y+h, x:x+w]\n\n    return img\n\n\ndef preprocess_image(path):\n    img = cv2.imread(path)\n    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n\n    # crop black borders\n    img = crop_fundus(img)\n\n    # resize\n    img = cv2.resize(img, (IMG_SIZE, IMG_SIZE))\n\n    # CLAHE\n    lab = cv2.cvtColor(img, cv2.COLOR_RGB2LAB)\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_LAB2RGB)\n\n    # Gaussian blur\n    img = cv2.GaussianBlur(img, (3,3), 0)\n\n    # normalize\n    img = img.astype(\"float32\") / 255.0\n    return img\n\n# ----------------------------------------------------------\n# LOAD IMAGES INTO MEMORY\n# ----------------------------------------------------------\nX = np.zeros((len(df), IMG_SIZE, IMG_SIZE, 3), dtype=np.float32)\ny = df[\"diagnosis\"].values\n\nfor i, path in enumerate(tqdm(df[\"path\"])):\n    X[i] = preprocess_image(path)\n\n# ----------------------------------------------------------\n# LABELS\n# ----------------------------------------------------------\n# Binary: 0 vs DR\ny_binary = np.where(y == 0, 0, 1)\n\n# Early Detection: 0=No DR, 1=Early(1,2), 2=Severe(3,4)\ndef map_early(v):\n    if v == 0:\n        return 0\n    elif v in [1,2]:\n        return 1\n    else:\n        return 2\n\ny_early = np.array([map_early(v) for v in y])\n\n# Severity: Original 5 classes\ny_severity = y.copy()\n\n# ----------------------------------------------------------\n# AUGMENTATION\n# ----------------------------------------------------------\naugmenter = tf.keras.Sequential([\n    layers.RandomFlip(\"horizontal\"),\n    layers.RandomRotation(0.12),\n    layers.RandomZoom(0.12),\n    layers.RandomContrast(0.15),\n    layers.RandomTranslation(0.05, 0.05)\n])\n\n# ----------------------------------------------------------\n# PAPER CNN ARCHITECTURE\n# ----------------------------------------------------------\ndef build_model(num_classes):\n    inputs = layers.Input(shape=(IMG_SIZE, IMG_SIZE, 3))\n\n    x = layers.Conv2D(32, (5,5), padding=\"same\")(inputs)\n    x = layers.PReLU()(x)\n    x = layers.MaxPooling2D()(x)\n\n    x = layers.Conv2D(64, (5,5), activation=\"relu\", padding=\"same\")(x)\n    x = layers.MaxPooling2D()(x)\n\n    x = layers.Conv2D(128, (5,5), activation=\"relu\", padding=\"same\")(x)\n    x = layers.MaxPooling2D()(x)\n\n    x = layers.Conv2D(256, (5,5), activation=\"relu\", padding=\"same\")(x)\n    x = layers.MaxPooling2D()(x)\n\n    x = layers.GlobalAveragePooling2D()(x)\n\n    x = layers.Dense(256, activation=\"relu\")(x)\n    x = layers.Dropout(0.3)(x)\n\n    x = layers.Dense(128, activation=\"relu\")(x)\n    x = layers.Dropout(0.3)(x)\n\n    x = layers.Dense(64, activation=\"relu\")(x)\n\n    outputs = layers.Dense(num_classes, activation=\"softmax\")(x)\n\n    model = models.Model(inputs, outputs)\n    return model\n\n# ----------------------------------------------------------\n# TRAIN FUNCTION\n# ----------------------------------------------------------\ndef run_experiment(X, y_labels, num_classes, epochs, batch_size, lr, name):\n\n    print(f\"\\n========== {name} ==========\")\n\n    skf = StratifiedKFold(n_splits=5, shuffle=True, random_state=SEED)\n    scores = []\n\n    all_true = []\n    all_pred = []\n\n    for fold, (train_idx, val_idx) in enumerate(skf.split(X, y_labels), 1):\n\n        print(f\"\\n--- Fold {fold} ---\")\n\n        X_train, X_val = X[train_idx], X[val_idx]\n        y_train, y_val = y_labels[train_idx], y_labels[val_idx]\n\n        y_train_cat = to_categorical(y_train, num_classes)\n        y_val_cat   = to_categorical(y_val, num_classes)\n\n        # Class weights\n        weights = compute_class_weight(\n            class_weight=\"balanced\",\n            classes=np.unique(y_train),\n            y=y_train\n        )\n        class_weights = dict(enumerate(weights))\n\n        # Better weights only for severity classification\n        if name == \"Severity\":\n    class_weights = {\n        0: 1.0,\n        1: 1.4,\n        2: 1.3,\n        3: 2.0,\n        4: 2.5\n    }\n\n        # tf.data\n        train_ds = tf.data.Dataset.from_tensor_slices((X_train, y_train_cat))\n        train_ds = train_ds.shuffle(1024).batch(batch_size)\n        train_ds = train_ds.map(lambda a,b: (augmenter(a, training=True), b))\n        train_ds = train_ds.prefetch(tf.data.AUTOTUNE)\n\n        val_ds = tf.data.Dataset.from_tensor_slices((X_val, y_val_cat))\n        val_ds = val_ds.batch(batch_size).prefetch(tf.data.AUTOTUNE)\n\n        # Multi-GPU model build\n        with strategy.scope():\n            model = build_model(num_classes)\n            model.compile(\n                optimizer=tf.keras.optimizers.Adam(learning_rate=lr),\n                loss=tf.keras.losses.CategoricalCrossentropy(label_smoothing=0.1),\n                metrics=[\"accuracy\"]\n            )\n\n        callbacks = [\n    ModelCheckpoint(f\"{name}_fold{fold}.keras\", save_best_only=True),\n    EarlyStopping(\n        monitor=\"val_accuracy\",\n        patience=8,\n        restore_best_weights=True\n    ),\n    ReduceLROnPlateau(\n        monitor=\"val_accuracy\",\n        factor=0.5,\n        patience=3,\n        min_lr=1e-7\n    )\n]\n\n        model.fit(\n            train_ds,\n            validation_data=val_ds,\n            epochs=epochs,\n            class_weight=class_weights,\n            callbacks=callbacks,\n            verbose=1\n        )\n\n        preds = model.predict(val_ds, verbose=0)\n        pred_labels = np.argmax(preds, axis=1)\n\n        acc = accuracy_score(y_val, pred_labels)\n        scores.append(acc)\n\n        all_true.extend(y_val)\n        all_pred.extend(pred_labels)\n\n        print(\"Fold Accuracy:\", acc)\n\n        # Confusion Matrix\n        cm = confusion_matrix(y_val, pred_labels)\n        plt.figure(figsize=(6,5))\n        sns.heatmap(cm, annot=True, fmt=\"d\", cmap=\"Blues\")\n        plt.title(f\"{name} Fold {fold}\")\n        plt.show()\n\n        # ROC for binary\n        if num_classes == 2:\n            auc = roc_auc_score(y_val, preds[:,1])\n            fpr, tpr, _ = roc_curve(y_val, preds[:,1])\n\n            plt.figure(figsize=(6,5))\n            plt.plot(fpr, tpr, label=f\"AUC = {auc:.3f}\")\n            plt.plot([0,1],[0,1],'--')\n            plt.legend()\n            plt.title(\"ROC Curve\")\n            plt.show()\n\n    print(\"\\nAverage Accuracy:\", np.mean(scores))\n    print(\"\\nClassification Report:\")\n    print(classification_report(all_true, all_pred))\n\n# ----------------------------------------------------------\n# RUN ALL 3 EXPERIMENTS\n# ----------------------------------------------------------\n\n# 1 Binary Classification\nrun_experiment(\n    X, y_binary,\n    num_classes=2,\n    epochs=30,\n    batch_size=32,\n    lr=1e-4,\n    name=\"Binary\"\n)\n\n# 2 Early Detection\nrun_experiment(\n    X, y_early,\n    num_classes=3,\n    epochs=40,\n    batch_size=16,\n    lr=1e-3,\n    name=\"Early\"\n)\n\n# 3 Severity Classification\nrun_experiment(\n    X, y_severity,\n    num_classes=5,\n    epochs=40,\n    batch_size=8,\n    lr=1e-4,\n    name=\"Severity\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T07:57:33.852973Z","iopub.execute_input":"2026-04-17T07:57:33.853449Z"}},"outputs":[],"execution_count":null}]}