{"cells":[{"cell_type":"markdown","metadata":{},"source":"# Fundus Gate: Binary Image Classifier (Fundus vs. Non-Fundus)\n### Pre-screening Input Validation Gate for Diabetic Retinopathy Diagnostic Pipeline\n**Thesis:** Nghiên cứu Vision Transformer và Ứng dụng Phân loại Mức độ Bệnh Tiểu đường từ Ảnh Đáy Võng mạc Mắt  \n**Architecture:** Pretrained MobileNetV3Small (Lightweight, Web/Mobile Optimized)  \n**Input:** 224x224x3  \n**Classes:** `0 = non_fundus`, `1 = fundus`"},{"cell_type":"code","execution_count":null,"metadata":{"trusted":true},"outputs":[],"source":"import os\nimport glob\nimport json\nimport random\nimport time\nimport shutil\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom PIL import Image\n\nimport tensorflow as tf\nfrom tensorflow.keras import layers, models, optimizers, callbacks\nfrom sklearn.model_selection import train_test_split, GroupShuffleSplit\nfrom sklearn.metrics import (\n    accuracy_score, precision_score, recall_score, f1_score,\n    roc_auc_score, roc_curve, confusion_matrix\n)\n\n# 1. Reproducibility\nSEED = 42\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntf.random.set_seed(SEED)\nos.environ['PYTHONHASHSEED'] = str(SEED)\nos.environ['TF_DETERMINISTIC_OPS'] = '1'\n\n# 2. Hardware & Mixed Precision Configuration\nprint(f\"TensorFlow Version: {tf.__version__}\")\ngpus = tf.config.list_physical_devices('GPU')\nif gpus:\n    for gpu in gpus:\n        tf.config.experimental.set_memory_growth(gpu, True)\n    tf.keras.mixed_precision.set_global_policy('mixed_float16')\n    print(f\"GPUs available: {[g.name for g in gpus]}\")\n    print(f\"Mixed Precision Policy: {tf.keras.mixed_precision.global_policy().name}\")\nelse:\n    print(\"WARNING: No GPU detected, running on CPU.\")"},{"cell_type":"code","execution_count":null,"metadata":{"trusted":true},"outputs":[],"source":"print(\"=\" * 60)\nprint(\"1. DIRECT DATASET DISCOVERY & FAST DEDUPLICATION\")\nprint(\"=\" * 60)\n\nstart_scan_time = time.time()\nvalid_exts = ('.jpg', '.jpeg', '.png', '.bmp')\n\ndef get_images_from_dataset(owner_slug, max_images=None):\n    candidates = [\n        f\"/kaggle/input/datasets/{owner_slug}\",\n        f\"/kaggle/input/{owner_slug.split('/')[-1]}\",\n        f\"/kaggle/input/{owner_slug}\",\n        f\"/kaggle/input/competitions/{owner_slug.split('/')[-1]}\"\n    ]\n    ds_root = None\n    for c in candidates:\n        if os.path.exists(c):\n            ds_root = c\n            break\n            \n    if ds_root is None:\n        slug = owner_slug.split('/')[-1]\n        for root, dirs, _ in os.walk(\"/kaggle/input\"):\n            if slug in dirs:\n                ds_root = os.path.join(root, slug)\n                break\n                \n    if ds_root is None:\n        print(f\"  [!] Directory not found for: {owner_slug}\")\n        return []\n        \n    imgs = set()\n    for root, _, files in os.walk(ds_root):\n        for f in files:\n            if f.lower().endswith(valid_exts):\n                imgs.add(os.path.join(root, f))\n                if max_images and len(imgs) >= max_images:\n                    return list(imgs)\n    print(f\"  -> {owner_slug} ({ds_root}): {len(imgs):,} images found\")\n    return list(imgs)\n\n# 1. Fundus (Positive, EyePACS)\nprint(\"Collecting Fundus images...\")\nfundus_raw = get_images_from_dataset(\"tanlikesmath/diabetic-retinopathy-resized\", max_images=16000)\nif len(fundus_raw) == 0:\n    fundus_raw = get_images_from_dataset(\"diabetic-retinopathy-detection\", max_images=16000)\nfundus_candidates = list(set(fundus_raw))\nprint(f\"Total Unique Fundus: {len(fundus_candidates):,}\")\n\n# 2. Group A: Generic Negatives\nprint(\"\\nCollecting Generic Negatives...\")\ngeneric_natural = get_images_from_dataset(\"prasunroy/natural-images\", max_images=3000)\ngeneric_intel = get_images_from_dataset(\"puneet6060/intel-image-classification\", max_images=3000)\ngeneric_docs = get_images_from_dataset(\"ritvik1909/document-classification-dataset\", max_images=1000)\ngeneric_candidates = list(set(generic_natural + generic_intel + generic_docs))\nprint(f\"Total Unique Generic Negatives: {len(generic_candidates):,}\")\n\n# 3. Group B: Medical Negatives (retinal OCT, chest X-ray, brain MRI)\nprint(\"\\nCollecting Medical Negatives...\")\nmedical_oct = get_images_from_dataset(\"shakilrana/octdl-retinal-oct-images-dataset\", max_images=3000)\nmedical_xray = get_images_from_dataset(\"thomasdubail/chest-pneumonia-256x256\", max_images=3000)\nmedical_mri = get_images_from_dataset(\"sartajbhuvaji/brain-tumor-classification-mri\", max_images=1500)\nmedical_candidates = list(set(medical_oct + medical_xray + medical_mri))\nprint(f\"Total Unique Medical Negatives: {len(medical_candidates):,}\")\n\n# 4. Group C: Ocular Hard Negatives (CASIA iris, eye-dataset close-ups, cataract anterior, eye-disease)\nprint(\"\\nCollecting Ocular Hard Negatives...\")\nocular_iris = get_images_from_dataset(\"monareyhanii/casia-iris-syn\", max_images=3000)\nocular_eye = get_images_from_dataset(\"kayvanshah/eye-dataset\", max_images=2000)\nocular_cataract = get_images_from_dataset(\"akshayramakrishnan28/cataract-classification-dataset\", max_images=1500)\nocular_diseases = get_images_from_dataset(\"kondwani/eye-disease-dataset\", max_images=1000)\nocular_candidates = list(set(ocular_iris + ocular_eye + ocular_cataract + ocular_diseases))\nprint(f\"Total Unique Ocular Hard Negatives: {len(ocular_candidates):,}\")\n\nprint(f\"\\nDataset scan completed in {time.time() - start_scan_time:.2f}s!\")"},{"cell_type":"code","execution_count":null,"metadata":{"trusted":true},"outputs":[],"source":"print(\"=\" * 60)\nprint(\"2. DATASET BALANCING, INTEGRITY & CHALLENGE TEST ISOLATION\")\nprint(\"=\" * 60)\n\ndef filter_valid_images(path_list, max_sample=None):\n    valid = []\n    shuffled = list(set(path_list))\n    random.shuffle(shuffled)\n    for p in shuffled:\n        try:\n            if os.path.getsize(p) > 500:\n                valid.append(p)\n                if max_sample and len(valid) >= max_sample:\n                    break\n        except Exception:\n            continue\n    return valid\n\nTARGET_PER_GROUP = 3500\nvalid_generic = filter_valid_images(generic_candidates, TARGET_PER_GROUP)\nvalid_medical = filter_valid_images(medical_candidates, TARGET_PER_GROUP)\nvalid_ocular_all = filter_valid_images(ocular_candidates, TARGET_PER_GROUP + 1000)\n\nprint(f\"Valid pool: Generic={len(valid_generic):,}, Medical={len(valid_medical):,}, Ocular={len(valid_ocular_all):,}\")\n\n# 1. STRICT ISOLATION OF OCULAR HARD NEGATIVES CHALLENGE TEST\nCHALLENGE_SIZE = min(1000, max(500, int(len(valid_ocular_all) * 0.20)))\nrandom.shuffle(valid_ocular_all)\nchallenge_ocular_paths = valid_ocular_all[:CHALLENGE_SIZE]\nstandard_ocular_paths = valid_ocular_all[CHALLENGE_SIZE:CHALLENGE_SIZE + TARGET_PER_GROUP]\n\nchallenge_test_df = pd.DataFrame({\n    'filepath': challenge_ocular_paths,\n    'label': 0,\n    'category': 'ocular_hard_challenge',\n    'patient_id': ['non_fundus'] * len(challenge_ocular_paths)\n})\nprint(f\"Strictly Isolated Ocular Challenge Test Set: {len(challenge_test_df):,} images\")\nprint(\"  (Zero exposure during model training or validation)\")\n\n# 2. Build Standard Negatives\ngeneric_df = pd.DataFrame({\n    'filepath': valid_generic,\n    'label': 0,\n    'category': 'generic',\n    'patient_id': ['non_fundus'] * len(valid_generic)\n})\n\nmedical_df = pd.DataFrame({\n    'filepath': valid_medical,\n    'label': 0,\n    'category': 'medical',\n    'patient_id': ['non_fundus'] * len(valid_medical)\n})\n\nocular_df = pd.DataFrame({\n    'filepath': standard_ocular_paths,\n    'label': 0,\n    'category': 'ocular_hard',\n    'patient_id': ['non_fundus'] * len(standard_ocular_paths)\n})\n\ntotal_negatives_count = len(generic_df) + len(medical_df) + len(ocular_df)\nprint(f\"Total Standard Non-Fundus Pool: {total_negatives_count:,} images\")\n\n# 3. Sample Equal Number of Fundus Images for 50/50 Balance\nTARGET_FUNDUS = total_negatives_count\nvalid_fundus = filter_valid_images(fundus_candidates, TARGET_FUNDUS)\n\nfundus_df = pd.DataFrame({\n    'filepath': valid_fundus,\n    'label': 1,\n    'category': 'fundus'\n})\nfundus_df['patient_id'] = fundus_df['filepath'].apply(\n    lambda x: os.path.basename(x).split('.')[0].split('_')[0]\n)\n\nprint(f\"Standard Fundus Pool: {len(fundus_df):,} images ({fundus_df['patient_id'].nunique():,} unique patients)\")\nprint(f\"Balanced Dataset Ratio: {len(fundus_df)/(len(fundus_df) + total_negatives_count)*100:.1f}% Fundus / {total_negatives_count/(len(fundus_df) + total_negatives_count)*100:.1f}% Non-Fundus\")"},{"cell_type":"code","execution_count":null,"metadata":{"trusted":true},"outputs":[],"source":"print(\"=\" * 60)\nprint(\"3. SPLIT TRAIN (70%) / VAL (15%) / TEST (15%) & DATA INTEGRITY VERIFICATION\")\nprint(\"=\" * 60)\n\n# 1. Split Fundus using GroupShuffleSplit on patient_id to prevent data leakage\n# Train: 70%, Val: 15%, Test: 15%\ngss_train = GroupShuffleSplit(n_splits=1, train_size=0.70, random_state=SEED)\ntrain_idx, temp_idx = next(gss_train.split(fundus_df, groups=fundus_df['patient_id']))\nfundus_train = fundus_df.iloc[train_idx].copy()\nfundus_temp = fundus_df.iloc[temp_idx].copy()\n\n# Split remaining 30% equally into Val (15%) and Test (15%)\ngss_val = GroupShuffleSplit(n_splits=1, train_size=0.50, random_state=SEED)\nval_idx, test_idx = next(gss_val.split(fundus_temp, groups=fundus_temp['patient_id']))\nfundus_val = fundus_temp.iloc[val_idx].copy()\nfundus_test = fundus_temp.iloc[test_idx].copy()\n\n# Verify zero patient overlap in fundus\np_train = set(fundus_train['patient_id'])\np_val = set(fundus_val['patient_id'])\np_test = set(fundus_test['patient_id'])\nassert len(p_train.intersection(p_val)) == 0, \"LEAKAGE DETECTED between Train and Val!\"\nassert len(p_train.intersection(p_test)) == 0, \"LEAKAGE DETECTED between Train and Test!\"\nassert len(p_val.intersection(p_test)) == 0, \"LEAKAGE DETECTED between Val and Test!\"\nprint(\"Zero patient leakage verified across Fundus Train, Val, and Test.\")\n\n# 2. Split Negatives with Stratification on Category (Generic, Medical, Ocular)\nnegatives_df = pd.concat([generic_df, medical_df, ocular_df], ignore_index=True)\nnegatives_df = negatives_df.drop_duplicates(subset=['filepath']).reset_index(drop=True)\n\nneg_train, neg_temp = train_test_split(\n    negatives_df, train_size=0.70, stratify=negatives_df['category'], random_state=SEED\n)\nneg_val, neg_test = train_test_split(\n    neg_temp, train_size=0.50, stratify=neg_temp['category'], random_state=SEED\n)\n\n# 3. Combine into final Train, Val, and Test DataFrames\ntrain_df = pd.concat([fundus_train, neg_train], ignore_index=True).sample(frac=1.0, random_state=SEED).reset_index(drop=True)\nval_df = pd.concat([fundus_val, neg_val], ignore_index=True).sample(frac=1.0, random_state=SEED).reset_index(drop=True)\ntest_df = pd.concat([fundus_test, neg_test], ignore_index=True).sample(frac=1.0, random_state=SEED).reset_index(drop=True)\n\n# 4. Strict Filepath Disjointness Verification\nall_train_paths = set(train_df['filepath'])\nall_val_paths = set(val_df['filepath'])\nall_test_paths = set(test_df['filepath'])\nall_challenge_paths = set(challenge_test_df['filepath'])\n\nassert len(all_train_paths.intersection(all_val_paths)) == 0, \"Train-Val overlap detected!\"\nassert len(all_train_paths.intersection(all_test_paths)) == 0, \"Train-Test overlap detected!\"\nassert len(all_val_paths.intersection(all_test_paths)) == 0, \"Val-Test overlap detected!\"\nassert len(all_challenge_paths.intersection(all_train_paths)) == 0, \"Challenge-Train overlap detected!\"\nassert len(all_challenge_paths.intersection(all_val_paths)) == 0, \"Challenge-Val overlap detected!\"\nassert len(all_challenge_paths.intersection(all_test_paths)) == 0, \"Challenge-Test overlap detected!\"\nprint(\"Zero image overlap verified between all partitions including Challenge Test.\")\n\n# 5. Display Breakdown Table\nsplit_summary = pd.DataFrame({\n    'Train (70%)': train_df['category'].value_counts(),\n    'Val (15%)': val_df['category'].value_counts(),\n    'Test (15%)': test_df['category'].value_counts()\n}).fillna(0).astype(int)\nsplit_summary.loc['TOTAL'] = split_summary.sum()\nprint(\"\\n--- SPLIT COMPOSITION SUMMARY ---\")\nprint(split_summary.to_string())\nprint(f\"\\nChallenge Test (Ocular Hard Holdout): {len(challenge_test_df):,} images\")"},{"cell_type":"code","execution_count":null,"metadata":{"trusted":true},"outputs":[],"source":"print(\"=\" * 60)\nprint(\"4. BUILDING OPTIMIZED TF.DATA PIPELINE WITH LIGHT AUGMENTATION\")\nprint(\"=\" * 60)\n\nIMG_SIZE = (224, 224)\nBATCH_SIZE = 64\nAUTOTUNE = tf.data.AUTOTUNE\n\ndef load_and_preprocess_image(path, label):\n    img_bytes = tf.io.read_file(path)\n    img = tf.io.decode_image(img_bytes, channels=3, expand_animations=False)\n    img = tf.image.resize(img, IMG_SIZE)\n    img = tf.cast(img, tf.float32)\n    img = tf.keras.applications.mobilenet_v3.preprocess_input(img)\n    label = tf.cast(label, tf.float32)\n    return img, label\n\n# Light augmentation layer: horizontal flip, slight rotation, zoom, contrast\naugmentation_layer = tf.keras.Sequential([\n    layers.RandomFlip(\"horizontal\"),\n    layers.RandomRotation(factor=0.05),     # ~18 degrees\n    layers.RandomZoom(height_factor=0.05, width_factor=0.05),\n    layers.RandomContrast(factor=0.1)\n], name=\"light_augmentation\")\n\ndef augment_image(img, label):\n    img = augmentation_layer(img, training=True)\n    return img, label\n\ndef create_tf_dataset(df, is_training=False):\n    paths = df['filepath'].values\n    labels = df['label'].values\n    \n    ds = tf.data.Dataset.from_tensor_slices((paths, labels))\n    ds = ds.map(load_and_preprocess_image, num_parallel_calls=AUTOTUNE)\n    ds = ds.cache()\n    \n    if is_training:\n        ds = ds.shuffle(buffer_size=min(len(df), 4096), seed=SEED, reshuffle_each_iteration=True)\n        ds = ds.map(augment_image, num_parallel_calls=AUTOTUNE)\n        \n    ds = ds.batch(BATCH_SIZE)\n    ds = ds.prefetch(buffer_size=AUTOTUNE)\n    return ds\n\ntrain_ds = create_tf_dataset(train_df, is_training=True)\nval_ds = create_tf_dataset(val_df, is_training=False)\ntest_ds = create_tf_dataset(test_df, is_training=False)\nchallenge_ds = create_tf_dataset(challenge_test_df, is_training=False)\n\n# Sanity check: inspect first batch\nsample_batch_imgs, sample_batch_labels = next(iter(train_ds))\nprint(f\"Batch Image Shape: {sample_batch_imgs.shape} (Expected: ({BATCH_SIZE}, 224, 224, 3))\")\nprint(f\"Batch Label Shape: {sample_batch_labels.shape}\")\nprint(f\"Batch Label Distribution: {int(tf.reduce_sum(sample_batch_labels).numpy())} Fundus, {BATCH_SIZE - int(tf.reduce_sum(sample_batch_labels).numpy())} Non-Fundus\")\n\n# Plot 6 sanity check samples\nfig, axes = plt.subplots(1, 6, figsize=(15, 3))\nfor i in range(6):\n    img = sample_batch_imgs[i].numpy().astype(np.float32)\n    img_disp = (img + 1.0) / 2.0\n    lbl = int(sample_batch_labels[i].numpy())\n    axes[i].imshow(np.clip(img_disp, 0.0, 1.0).astype(np.float32))\n    axes[i].set_title(f\"{'FUNDUS (1)' if lbl==1 else 'NON-FUNDUS (0)'}\")\n    axes[i].axis('off')\nplt.tight_layout()\nplt.savefig('/kaggle/working/data_sanity_check.png', dpi=150)\nplt.close(fig)\nprint(\"Saved data sanity check plot to /kaggle/working/data_sanity_check.png\")"},{"cell_type":"code","execution_count":null,"metadata":{"trusted":true},"outputs":[],"source":"print(\"=\" * 60)\nprint(\"5. BUILDING MOBILENETV3SMALL CLASSIFIER\")\nprint(\"=\" * 60)\n\ndef build_fundus_gate_model():\n    base_model = tf.keras.applications.MobileNetV3Small(\n        input_shape=(224, 224, 3),\n        include_top=False,\n        weights='imagenet'\n    )\n    base_model.trainable = False  # Initially frozen for Stage 1\n    \n    inputs = layers.Input(shape=(224, 224, 3), name='input_image')\n    x = base_model(inputs, training=False)\n    x = layers.GlobalAveragePooling2D(name='global_avg_pool')(x)\n    x = layers.Dropout(0.3, name='dropout_regularization')(x)\n    # Explicit float32 dtype on final output layer for numerical stability in mixed precision\n    outputs = layers.Dense(1, activation='sigmoid', dtype='float32', name='fundus_gate_prob')(x)\n    \n    model = models.Model(inputs=inputs, outputs=outputs, name='FundusGate_MobileNetV3Small')\n    return model, base_model\n\nmodel, base_model = build_fundus_gate_model()\nmodel.summary()\n\ntrainable_params = sum(int(np.prod(v.shape)) for v in model.trainable_weights)\nnon_trainable_params = sum(int(np.prod(v.shape)) for v in model.non_trainable_weights)\nprint(f\"Total Parameters: {trainable_params + non_trainable_params:,}\")\nprint(f\"Trainable Parameters (Stage 1): {trainable_params:,}\")\nprint(f\"Frozen Parameters: {non_trainable_params:,}\")"},{"cell_type":"code","execution_count":null,"metadata":{"trusted":true},"outputs":[],"source":"print(\"=\" * 60)\nprint(\"6. STAGE 1 TRAINING: CLASSIFICATION HEAD ONLY (FROZEN BACKBONE)\")\nprint(\"=\" * 60)\n\nmodel.compile(\n    optimizer=optimizers.Adam(learning_rate=1e-3),\n    loss=tf.keras.losses.BinaryCrossentropy(),\n    metrics=[\n        'accuracy',\n        tf.keras.metrics.AUC(name='roc_auc'),\n        tf.keras.metrics.Precision(name='precision'),\n        tf.keras.metrics.Recall(name='recall')\n    ]\n)\n\nSTAGE1_EPOCHS = 3\nstart_time = time.time()\n\nhistory_stage1 = model.fit(\n    train_ds,\n    validation_data=val_ds,\n    epochs=STAGE1_EPOCHS,\n    verbose=1\n)\n\nprint(f\"Stage 1 completed in {time.time() - start_time:.1f}s\")"},{"cell_type":"code","execution_count":null,"metadata":{"trusted":true},"outputs":[],"source":"print(\"=\" * 60)\nprint(\"7. STAGE 2 TRAINING: FINE-TUNING TOP 30% LAYERS\")\nprint(\"=\" * 60)\n\n# Unfreeze top ~30% layers of MobileNetV3Small\nbase_model.trainable = True\nnum_layers = len(base_model.layers)\nunfreeze_from = int(num_layers * 0.70)  # Freeze first 70%, unfreeze last 30%\n\nfor layer in base_model.layers[:unfreeze_from]:\n    layer.trainable = False\nfor layer in base_model.layers[unfreeze_from:]:\n    layer.trainable = True\n\nprint(f\"Total layers in MobileNetV3Small: {num_layers}\")\nprint(f\"Frozen layers: 0 to {unfreeze_from-1}\")\nprint(f\"Unfrozen fine-tuning layers: {unfreeze_from} to {num_layers-1} ({num_layers - unfreeze_from} layers)\")\n\n# Recompile with small learning rate\nmodel.compile(\n    optimizer=optimizers.Adam(learning_rate=1e-4),\n    loss=tf.keras.losses.BinaryCrossentropy(),\n    metrics=[\n        'accuracy',\n        tf.keras.metrics.AUC(name='roc_auc'),\n        tf.keras.metrics.Precision(name='precision'),\n        tf.keras.metrics.Recall(name='recall')\n    ]\n)\n\nSTAGE2_EPOCHS = 3\nbest_model_path = '/kaggle/working/fundus_gate_best.keras'\n\ntraining_callbacks = [\n    callbacks.ModelCheckpoint(\n        filepath=best_model_path,\n        monitor='val_loss',\n        mode='min',\n        save_best_only=True,\n        verbose=1\n    ),\n    callbacks.EarlyStopping(\n        monitor='val_loss',\n        patience=2,\n        restore_best_weights=True,\n        verbose=1\n    ),\n    callbacks.ReduceLROnPlateau(\n        monitor='val_loss',\n        factor=0.5,\n        patience=1,\n        min_lr=1e-6,\n        verbose=1\n    )\n]\n\nstart_time = time.time()\nhistory_stage2 = model.fit(\n    train_ds,\n    validation_data=val_ds,\n    epochs=STAGE2_EPOCHS,\n    callbacks=training_callbacks,\n    verbose=1\n)\nprint(f\"Stage 2 fine-tuning completed in {time.time() - start_time:.1f}s\")\n\n# Load best checkpoint\nif os.path.exists(best_model_path):\n    print(f\"Loading best weights from {best_model_path}\")\n    model = tf.keras.models.load_model(best_model_path)"},{"cell_type":"code","execution_count":null,"metadata":{"trusted":true},"outputs":[],"source":"print(\"=\" * 60)\nprint(\"8. VALIDATION SET THRESHOLD OPTIMIZATION\")\nprint(\"=\" * 60)\n\n# Compute predictions on Validation Set\nprint(\"Generating predictions on Validation Set...\")\nval_preds = model.predict(val_ds, verbose=1).flatten()\nval_labels = val_df['label'].values\n\n# Calculate ROC curve & AUC\nval_auc = roc_auc_score(val_labels, val_preds)\nprint(f\"Validation ROC-AUC: {val_auc:.4f}\")\n\n# Threshold Search:\n# Objective: Fundus Recall >= 95% while minimizing False Positive Rate (FPR) / false acceptance\ncandidate_thresholds = np.linspace(0.01, 0.99, 197)\nthreshold_records = []\n\nfor t in candidate_thresholds:\n    preds_binary = (val_preds >= t).astype(int)\n    rec = recall_score(val_labels, preds_binary, zero_division=0)\n    prec = precision_score(val_labels, preds_binary, zero_division=0)\n    f1 = f1_score(val_labels, preds_binary, zero_division=0)\n    \n    # Specificity & FPR\n    tn, fp, fn, tp = confusion_matrix(val_labels, preds_binary).ravel()\n    spec = tn / (tn + fp) if (tn + fp) > 0 else 0.0\n    fpr = fp / (tn + fp) if (tn + fp) > 0 else 0.0\n    \n    threshold_records.append({\n        'threshold': float(t),\n        'recall': float(rec),\n        'precision': float(prec),\n        'f1': float(f1),\n        'specificity': float(spec),\n        'fpr': float(fpr)\n    })\n\nthresh_df = pd.DataFrame(threshold_records)\n\n# Filter thresholds satisfying Fundus Recall >= 95% (0.950)\nvalid_candidates = thresh_df[thresh_df['recall'] >= 0.950]\n\nif len(valid_candidates) > 0:\n    # Select threshold that minimizes FPR (i.e. minimizes false acceptance)\n    sorted_candidates = valid_candidates.sort_values(by=['fpr', 'f1'], ascending=[True, False])\n    optimal_row = sorted_candidates.iloc[0]\n    best_thresh = float(optimal_row['threshold'])\n    print(f\"Selected Optimal Threshold: {best_thresh:.4f}\")\n    print(f\"  Val Recall: {optimal_row['recall']*100:.2f}% (>= 95% constraint satisfied)\")\n    print(f\"  Val Specificity: {optimal_row['specificity']*100:.2f}%\")\n    print(f\"  Val FPR (False Acceptance Rate): {optimal_row['fpr']*100:.2f}%\")\n    print(f\"  Val F1-Score: {optimal_row['f1']:.4f}\")\nelse:\n    # Fallback to Youden's J index\n    thresh_df['youden_j'] = thresh_df['recall'] + thresh_df['specificity'] - 1.0\n    best_thresh = float(thresh_df.loc[thresh_df['youden_j'].idxmax()]['threshold'])\n    print(f\"Fallback Optimal Threshold: {best_thresh:.4f}\")\n\n# Save threshold.json\nval_row_chosen = thresh_df.iloc[(thresh_df['threshold'] - best_thresh).abs().argsort()[:1]].iloc[0]\nthreshold_meta = {\n    'optimal_threshold': round(best_thresh, 4),\n    'constraint': 'Fundus Recall >= 0.95 and minimum FPR',\n    'validation_metrics': {\n        'roc_auc': round(float(val_auc), 4),\n        'recall': round(float(val_row_chosen['recall']), 4),\n        'specificity': round(float(val_row_chosen['specificity']), 4),\n        'fpr': round(float(val_row_chosen['fpr']), 4),\n        'f1': round(float(val_row_chosen['f1']), 4)\n    }\n}\n\nwith open('/kaggle/working/threshold.json', 'w') as f:\n    json.dump(threshold_meta, f, indent=4)\nprint(\"Saved threshold metadata to /kaggle/working/threshold.json\")\n\n# Plot Threshold Optimization Curves\nplt.figure(figsize=(10, 6))\nplt.plot(thresh_df['threshold'], thresh_df['recall'], label='Fundus Recall (Sensitivity)', color='blue', lw=2)\nplt.plot(thresh_df['threshold'], thresh_df['specificity'], label='Specificity (1 - FPR)', color='green', lw=2)\nplt.plot(thresh_df['threshold'], thresh_df['fpr'], label='FPR (False Acceptance)', color='red', lw=2, linestyle='--')\nplt.axvline(best_thresh, color='black', linestyle=':', label=f'Chosen Threshold = {best_thresh:.2f}', lw=2)\nplt.axhline(0.95, color='gray', linestyle='-.', label='Target Recall >= 95%')\nplt.title('Threshold Optimization on Validation Set', fontsize=14, fontweight='bold')\nplt.xlabel('Classification Threshold', fontsize=12)\nplt.ylabel('Metric Value', fontsize=12)\nplt.grid(True, alpha=0.3)\nplt.legend(loc='lower left', fontsize=10)\nplt.tight_layout()\nplt.savefig('/kaggle/working/threshold_tuning.png', dpi=150)\nplt.close('all')"},{"cell_type":"code","execution_count":null,"metadata":{"trusted":true},"outputs":[],"source":"print(\"=\" * 60)\nprint(f\"9. EVALUATION ON UNSEEN TEST SET (FIXED THRESHOLD = {best_thresh:.4f})\")\nprint(\"=\" * 60)\n\ntest_preds = model.predict(test_ds, verbose=1).flatten()\ntest_labels = test_df['label'].values\ntest_preds_binary = (test_preds >= best_thresh).astype(int)\n\n# Comprehensive Metrics\ntest_acc = accuracy_score(test_labels, test_preds_binary)\ntest_prec = precision_score(test_labels, test_preds_binary, zero_division=0)\ntest_rec = recall_score(test_labels, test_preds_binary, zero_division=0)\ntest_f1 = f1_score(test_labels, test_preds_binary, zero_division=0)\ntest_auc = roc_auc_score(test_labels, test_preds)\n\ntn, fp, fn, tp = confusion_matrix(test_labels, test_preds_binary).ravel()\ntest_spec = tn / (tn + fp) if (tn + fp) > 0 else 0.0\ntest_fpr = fp / (tn + fp) if (tn + fp) > 0 else 0.0\ntest_fnr = fn / (tp + fn) if (tp + fn) > 0 else 0.0\n\nprint(f\"Test Accuracy:    {test_acc*100:.2f}%\")\nprint(f\"Test Precision:   {test_prec*100:.2f}%\")\nprint(f\"Test Recall:      {test_rec*100:.2f}% (Sensitivity)\")\nprint(f\"Test Specificity: {test_spec*100:.2f}% (True Negative Rate)\")\nprint(f\"Test FPR:         {test_fpr*100:.2f}% (False Acceptance Rate)\")\nprint(f\"Test FNR:         {test_fnr*100:.2f}%\")\nprint(f\"Test F1-Score:    {test_f1:.4f}\")\nprint(f\"Test ROC-AUC:     {test_auc:.4f}\")\nprint(f\"Confusion Matrix: TP={tp}, FP={fp}, TN={tn}, FN={fn}\")\n\n# Subgroup Breakdown on Test Set\ntest_df['predicted_prob'] = test_preds\ntest_df['predicted_class'] = test_preds_binary\n\nprint(\"\\n--- FALSE ACCEPTANCE RATE (FPR) BY SUBGROUP ---\")\nfor cat in ['generic', 'medical', 'ocular_hard']:\n    sub_df = test_df[test_df['category'] == cat]\n    if len(sub_df) > 0:\n        sub_fp = (sub_df['predicted_class'] == 1).sum()\n        sub_fpr = sub_fp / len(sub_df)\n        print(f\"  {cat.capitalize()} Negatives: {sub_fp}/{len(sub_df)} false accepted (FPR = {sub_fpr*100:.2f}%)\")\n\n# Plot Confusion Matrix\nplt.figure(figsize=(6, 5))\ncm = np.array([[tn, fp], [fn, tp]])\nsns.heatmap(cm, annot=True, fmt='d', cmap='Blues',\n            xticklabels=['Pred Non-Fundus (0)', 'Pred Fundus (1)'],\n            yticklabels=['Actual Non-Fundus (0)', 'Actual Fundus (1)'])\nplt.title(f'Test Confusion Matrix (Threshold={best_thresh:.2f})', fontsize=12, fontweight='bold')\nplt.tight_layout()\nplt.savefig('/kaggle/working/confusion_matrix.png', dpi=150)\nplt.close('all')"},{"cell_type":"code","execution_count":null,"metadata":{"trusted":true},"outputs":[],"source":"print(\"=\" * 60)\nprint(f\"10. EVALUATION ON OCULAR HARD NEGATIVES CHALLENGE TEST (STRICT HOLDOUT)\")\nprint(\"=\" * 60)\n\nchallenge_preds = model.predict(challenge_ds, verbose=1).flatten()\nchallenge_preds_binary = (challenge_preds >= best_thresh).astype(int)\n\nchallenge_test_df['predicted_prob'] = challenge_preds\nchallenge_test_df['predicted_class'] = challenge_preds_binary\n\ntotal_challenge = len(challenge_test_df)\nchallenge_fp = int((challenge_preds_binary == 1).sum())\nchallenge_tn = int((challenge_preds_binary == 0).sum())\nchallenge_fpr = challenge_fp / total_challenge if total_challenge > 0 else 0.0\nchallenge_spec = challenge_tn / total_challenge if total_challenge > 0 else 0.0\n\nprint(f\"Total Challenge Hard Ocular Samples: {total_challenge:,}\")\nprint(f\"Correctly Rejected (TN):             {challenge_tn:,} ({challenge_spec*100:.2f}%)\")\nprint(f\"False Acceptance (FP):               {challenge_fp:,} ({challenge_fpr*100:.2f}%)\")\nprint(f\"Challenge Specificity:               {challenge_spec*100:.2f}%\")\nprint(f\"Challenge FPR (False Acceptance):    {challenge_fpr*100:.2f}%\")"},{"cell_type":"code","execution_count":null,"metadata":{"trusted":true},"outputs":[],"source":"print(\"=\" * 60)\nprint(\"11. ERROR ANALYSIS: INSPECTING MISCLASSIFIED SAMPLES\")\nprint(\"=\" * 60)\n\nfps_standard = test_df[(test_df['label'] == 0) & (test_df['predicted_class'] == 1)]\nfps_challenge = challenge_test_df[(challenge_test_df['label'] == 0) & (challenge_test_df['predicted_class'] == 1)]\nall_fps = pd.concat([fps_standard, fps_challenge], ignore_index=True)\n\nall_fns = test_df[(test_df['label'] == 1) & (test_df['predicted_class'] == 0)]\n\nprint(f\"Total False Positives in Test + Challenge: {len(all_fps)}\")\nprint(f\"Total False Negatives in Test: {len(all_fns)}\")\n\nfig, axes = plt.subplots(2, 4, figsize=(16, 8))\n\n# Top row: False Positives\naxes[0, 0].set_ylabel(\"False Positives\\n(Non-Fundus → Accepted)\", fontsize=11, fontweight='bold')\nfor i in range(4):\n    ax = axes[0, i]\n    if i < len(all_fps):\n        row = all_fps.iloc[i]\n        try:\n            img = Image.open(row['filepath']).convert('RGB').resize((224, 224))\n            ax.imshow(np.array(img, dtype=np.uint8))\n            ax.set_title(f\"Cat: {row['category']}\\nP(fundus)={row['predicted_prob']:.2f}\", fontsize=10, color='red')\n        except Exception:\n            ax.text(0.5, 0.5, \"Load Error\", ha='center')\n    else:\n        ax.text(0.5, 0.5, \"No Error\", ha='center')\n    ax.axis('off')\n\n# Bottom row: False Negatives\naxes[1, 0].set_ylabel(\"False Negatives\\n(Fundus → Rejected)\", fontsize=11, fontweight='bold')\nfor i in range(4):\n    ax = axes[1, i]\n    if i < len(all_fns):\n        row = all_fns.iloc[i]\n        try:\n            img = Image.open(row['filepath']).convert('RGB').resize((224, 224))\n            ax.imshow(np.array(img, dtype=np.uint8))\n            ax.set_title(f\"Actual: Fundus\\nP(fundus)={row['predicted_prob']:.2f}\", fontsize=10, color='orange')\n        except Exception:\n            ax.text(0.5, 0.5, \"Load Error\", ha='center')\n    else:\n        ax.text(0.5, 0.5, \"No Error\", ha='center')\n    ax.axis('off')\n\nplt.suptitle(f\"Error Analysis on False Acceptances & False Rejections (Threshold={best_thresh:.2f})\", fontsize=14, fontweight='bold')\nplt.tight_layout()\nplt.savefig('/kaggle/working/error_analysis.png', dpi=150)\nplt.close('all')\nprint(\"Saved error analysis plot to /kaggle/working/error_analysis.png\")"},{"cell_type":"code","execution_count":null,"metadata":{"trusted":true},"outputs":[],"source":"print(\"=\" * 60)\nprint(\"12. EXPORTING FINAL MODEL ARTIFACTS (KERAS, TFLITE, TFLITE-FP16)\")\nprint(\"=\" * 60)\n\nexport_dir = \"/kaggle/working\"\nkeras_path = os.path.join(export_dir, \"fundus_gate.keras\")\ntflite_path = os.path.join(export_dir, \"fundus_gate.tflite\")\ntflite_fp16_path = os.path.join(export_dir, \"fundus_gate_float16.tflite\")\n\n# 1. Save Full Keras Model\nmodel.save(keras_path)\nprint(f\"Saved Keras model: {keras_path} ({os.path.getsize(keras_path)/1024/1024:.2f} MB)\")\n\n# 2. Convert to Standard TFLite (Float32) & Float16 Quantized via CPU Subprocess\n# (Prevents multi-GPU Grappler session deadlock on Dual Tesla T4 VMs)\nconvert_script = \"\\n\".join([\n    \"import os, sys\",\n    \"os.environ['CUDA_VISIBLE_DEVICES'] = ''\",\n    \"import tensorflow as tf\",\n    \"\",\n    f\"keras_path = '{keras_path}'\",\n    f\"tflite_path = '{tflite_path}'\",\n    f\"tflite_fp16_path = '{tflite_fp16_path}'\",\n    \"\",\n    \"print('Loading Keras model on CPU for clean TFLite conversion...')\",\n    \"model = tf.keras.models.load_model(keras_path)\",\n    \"run_model = tf.function(lambda x: model(x))\",\n    \"concrete_func = run_model.get_concrete_function(tf.TensorSpec([1, 224, 224, 3], tf.float32))\",\n    \"\",\n    \"print('Converting to Float32 TFLite...')\",\n    \"converter = tf.lite.TFLiteConverter.from_concrete_functions([concrete_func])\",\n    \"tflite_model = converter.convert()\",\n    \"with open(tflite_path, 'wb') as f:\",\n    \"    f.write(tflite_model)\",\n    \"print(f'Saved TFLite model: {tflite_path} ({os.path.getsize(tflite_path)/1024/1024:.2f} MB)')\",\n    \"\",\n    \"print('Converting to Float16 Quantized TFLite...')\",\n    \"converter_fp16 = tf.lite.TFLiteConverter.from_concrete_functions([concrete_func])\",\n    \"converter_fp16.optimizations = [tf.lite.Optimize.DEFAULT]\",\n    \"converter_fp16.target_spec.supported_types = [tf.float16]\",\n    \"tflite_fp16_model = converter_fp16.convert()\",\n    \"with open(tflite_fp16_path, 'wb') as f:\",\n    \"    f.write(tflite_fp16_model)\",\n    \"print(f'Saved TFLite Float16 model: {tflite_fp16_path} ({os.path.getsize(tflite_fp16_path)/1024/1024:.2f} MB)')\",\n    \"print('TFLite CPU Conversion Complete!')\"\n])\n\nscript_path = \"/kaggle/working/convert_tflite.py\"\nwith open(script_path, \"w\") as f:\n    f.write(convert_script)\n\nimport sys, subprocess\nproc = subprocess.run([sys.executable, script_path], capture_output=True, text=True)\nprint(proc.stdout)\nif proc.stderr:\n    print(proc.stderr)\nassert proc.returncode == 0, f\"TFLite conversion failed with returncode {proc.returncode}\"\n\n# 4. Verify TFLite Inference\nprint(\"\\nVerifying TFLite Inference...\")\ninterpreter = tf.lite.Interpreter(model_path=tflite_path)\ninterpreter.allocate_tensors()\ninput_details = interpreter.get_input_details()\noutput_details = interpreter.get_output_details()\n\ndummy_input = np.random.uniform(-1.0, 1.0, size=(1, 224, 224, 3)).astype(np.float32)\ninterpreter.set_tensor(input_details[0]['index'], dummy_input)\ninterpreter.invoke()\ndummy_pred = interpreter.get_tensor(output_details[0]['index'])[0][0]\nprint(f\"TFLite Dummy Inference Output: {dummy_pred:.4f} (Valid Range [0, 1])\")\n\n# 5. Save Complete metrics.json\nfull_metrics = {\n    'model_name': 'FundusGate_MobileNetV3Small',\n    'input_shape': [224, 224, 3],\n    'total_parameters': int(model.count_params()),\n    'optimal_threshold': round(best_thresh, 4),\n    'validation_metrics': {\n        'roc_auc': round(float(val_auc), 4),\n        'recall': round(float(val_row_chosen['recall']), 4),\n        'specificity': round(float(val_row_chosen['specificity']), 4),\n        'fpr': round(float(val_row_chosen['fpr']), 4),\n        'f1': round(float(val_row_chosen['f1']), 4)\n    },\n    'test_metrics': {\n        'accuracy': round(float(test_acc), 4),\n        'precision': round(float(test_prec), 4),\n        'recall': round(float(test_rec), 4),\n        'specificity': round(float(test_spec), 4),\n        'fpr': round(float(test_fpr), 4),\n        'fnr': round(float(test_fnr), 4),\n        'f1': round(float(test_f1), 4),\n        'roc_auc': round(float(test_auc), 4),\n        'confusion_matrix': {\n            'tp': int(tp), 'fp': int(fp), 'tn': int(tn), 'fn': int(fn)\n        }\n    },\n    'challenge_test_metrics': {\n        'total_samples': total_challenge,\n        'true_negatives': challenge_tn,\n        'false_positives': challenge_fp,\n        'challenge_specificity': round(float(challenge_spec), 4),\n        'challenge_fpr': round(float(challenge_fpr), 4)\n    },\n    'artifacts': {\n        'keras_size_mb': round(os.path.getsize(keras_path)/1024/1024, 2),\n        'tflite_size_mb': round(os.path.getsize(tflite_path)/1024/1024, 2),\n        'tflite_fp16_size_mb': round(os.path.getsize(tflite_fp16_path)/1024/1024, 2)\n    }\n}\n\nwith open('/kaggle/working/metrics.json', 'w') as f:\n    json.dump(full_metrics, f, indent=4)\nprint(\"Saved full metrics to /kaggle/working/metrics.json\")\n\n# 6. Final Console Summary Report\nprint(\"\\n\" + \"=\" * 60)\nprint(\"FINAL STANDARDIZED REPORT BLOCK\")\nprint(\"=\" * 60)\nprint(f\"Dataset: EyePACS ({len(fundus_df):,} Fundus) vs 3-Group Negatives ({len(negatives_df):,} Non-Fundus)\")\nprint(f\"Model: MobileNetV3Small (Pretrained ImageNet)\")\nprint(f\"Parameters: {model.count_params():,}\")\nprint(f\"Threshold: {best_thresh:.4f} (Calibrated on Validation set for Recall >= 95%)\")\nprint(f\"Test Recall: {test_rec*100:.2f}%\")\nprint(f\"Test Specificity: {test_spec*100:.2f}%\")\nprint(f\"Test FPR: {test_fpr*100:.2f}%\")\nprint(f\"Challenge FPR: {challenge_fpr*100:.2f}%\")\nprint(f\"Model size: {os.path.getsize(keras_path)/1024/1024:.2f} MB\")\nprint(f\"TFLite size: {os.path.getsize(tflite_path)/1024/1024:.2f} MB (FP16: {os.path.getsize(tflite_fp16_path)/1024/1024:.2f} MB)\")\nprint(\"=\" * 60)"}],"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.10.12"},"kaggle":{"accelerator":"gpu","dataSources":[],"dockerImageVersionId":30699,"isGpuEnabled":true,"isInternetEnabled":true,"language":"python","sourceType":"notebook"}},"nbformat":4,"nbformat_minor":4}