{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":24800,"datasetId":1042002,"databundleVersionId":1831594},{"sourceType":"datasetVersion","sourceId":18613,"datasetId":5839,"databundleVersionId":18613},{"sourceType":"modelInstanceVersion","sourceId":554469,"databundleVersionId":13582352,"modelInstanceId":422333}],"dockerImageVersionId":31040,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Lung Cancer Detection using ViT and DenseNet on NIH and VinDr datasets","metadata":{}},{"cell_type":"markdown","source":"### Imports & global setup","metadata":{}},{"cell_type":"code","source":"import os\nimport random\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\nimport numpy as np\nimport pandas as pd\nimport cv2\nimport pydicom\nfrom PIL import Image\nimport matplotlib.pyplot as plt\nimport matplotlib.cm as cm\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import (\n    roc_auc_score, average_precision_score,\n    precision_recall_curve, roc_curve,\n    classification_report, confusion_matrix, ConfusionMatrixDisplay\n)\n\n# TensorFlow / Keras imports \nimport tensorflow as tf\nimport tf_keras as keras\nfrom tf_keras import layers\nfrom tf_keras.callbacks import EarlyStopping\n\n\n# Hugging Face\nfrom transformers import TFViTModel, TFAutoModel","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-03T17:18:14.049187Z","iopub.execute_input":"2025-09-03T17:18:14.049828Z","iopub.status.idle":"2025-09-03T17:18:39.681506Z","shell.execute_reply.started":"2025-09-03T17:18:14.049802Z","shell.execute_reply":"2025-09-03T17:18:39.680738Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"- Pulls in standard PyData + TF stack.\n- warnings.filterwarnings(\"ignore\") hides noisy library warnings (useful in Kaggle logs).","metadata":{}},{"cell_type":"code","source":"# ---- CONFIG ----\nIMG_SIZE = 224\nSAVE_DIR = \"/kaggle/working/processed_data\"\nVINDR_PATH = \"/kaggle/input/vinbigdata-chest-xray-abnormalities-detection\"\nNIH_PATH = \"/kaggle/input/data\"\nos.makedirs(SAVE_DIR, exist_ok=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-03T17:18:39.682894Z","iopub.execute_input":"2025-09-03T17:18:39.683766Z","iopub.status.idle":"2025-09-03T17:18:39.687112Z","shell.execute_reply.started":"2025-09-03T17:18:39.683746Z","shell.execute_reply":"2025-09-03T17:18:39.68652Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"- Sets a common image size (224) that works with common ImageNet models.\n- Defines dataset locations (Kaggle input mounts).\n- Ensures a working directory exists for any artifacts.","metadata":{}},{"cell_type":"code","source":"# ---- LOAD DATA ----\nvindr_df = pd.read_csv(os.path.join(VINDR_PATH, \"train.csv\"))\nnih_df = pd.read_csv(os.path.join(NIH_PATH, \"Data_Entry_2017.csv\"))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-03T17:18:39.687948Z","iopub.execute_input":"2025-09-03T17:18:39.688175Z","iopub.status.idle":"2025-09-03T17:18:40.058686Z","shell.execute_reply.started":"2025-09-03T17:18:39.688155Z","shell.execute_reply":"2025-09-03T17:18:40.058117Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"- Loads each dataset’s metadata:\n> - NIH: Data_Entry_2017.csv holds file names and comma-separated labels.\n> - VinDr: train.csv holds infomations including image_id and class_name.","metadata":{}},{"cell_type":"markdown","source":"### Build NIH image path map","metadata":{}},{"cell_type":"code","source":"# ---- BUILD NIH IMAGE MAP ----\ndef build_nih_path_map(base_path):\n    path_map = {}\n    for subdir in os.listdir(base_path):\n        if subdir.startswith(\"images_\"):\n            img_dir = os.path.join(base_path, subdir, \"images\")\n            for file in os.listdir(img_dir):\n                path_map[file] = os.path.join(img_dir, file)\n    return path_map","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-03T17:18:40.059379Z","iopub.execute_input":"2025-09-03T17:18:40.059574Z","iopub.status.idle":"2025-09-03T17:18:40.064301Z","shell.execute_reply.started":"2025-09-03T17:18:40.059558Z","shell.execute_reply":"2025-09-03T17:18:40.063753Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"- NIH on Kaggle typically comes split into folders like images_001, images_002, etc., each containing an images subfolder.\n- This function walks those folders and creates a dict:\nkey = image filename (e.g., 00000001_000.png), value = absolute path.\n- You’ll use the map to quickly resolve any NIH “Image Index” to a file path.","metadata":{}},{"cell_type":"code","source":"nih_image_map = build_nih_path_map(NIH_PATH)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-03T17:18:40.066297Z","iopub.execute_input":"2025-09-03T17:18:40.066963Z","iopub.status.idle":"2025-09-03T17:18:41.245303Z","shell.execute_reply.started":"2025-09-03T17:18:40.066933Z","shell.execute_reply":"2025-09-03T17:18:41.244739Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"- Builds that lookup once.","metadata":{}},{"cell_type":"markdown","source":"### NIH binary labels (nodule/mass vs no finding)","metadata":{}},{"cell_type":"code","source":"# ---- NIH BINARY LABELS ----\nnih_pos = nih_df[nih_df[\"Finding Labels\"].str.contains(\"Nodule|Mass\")]\nnih_neg = nih_df[nih_df[\"Finding Labels\"] == \"No Finding\"]\nnih_neg_sampled = nih_neg.sample(n=len(nih_pos), random_state=42)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-03T17:18:41.245934Z","iopub.execute_input":"2025-09-03T17:18:41.246196Z","iopub.status.idle":"2025-09-03T17:18:41.320379Z","shell.execute_reply.started":"2025-09-03T17:18:41.246173Z","shell.execute_reply":"2025-09-03T17:18:41.319611Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"- Positive = any record whose label string contains “Nodule” or “Mass”. (str.contains uses regex “Nodule|Mass”.)\n- Negative = records labeled exactly “No Finding”.\n- Then class balance: downsample negatives to match the number of positives (prevents a biased classifier).","metadata":{}},{"cell_type":"code","source":"nih_binary_df = pd.concat([nih_pos, nih_neg_sampled])\nnih_binary_df[\"label\"] = nih_binary_df[\"Finding Labels\"].apply(lambda x: 1 if \"Nodule\" in x or \"Mass\" in x else 0)\nnih_binary_df = nih_binary_df.sample(frac=1, random_state=42).reset_index(drop=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-03T17:18:41.321219Z","iopub.execute_input":"2025-09-03T17:18:41.321772Z","iopub.status.idle":"2025-09-03T17:18:41.342944Z","shell.execute_reply.started":"2025-09-03T17:18:41.321745Z","shell.execute_reply":"2025-09-03T17:18:41.342245Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"- Merges pos/neg, assigns a binary label (1 for nodule/mass, 0 otherwise), and shuffles.\n\n**Why this matters:** NIH has multi-label findings; you’re collapsing them to a single binary clinical target and balancing classes for stable training.","metadata":{}},{"cell_type":"markdown","source":"### Build VinDr image path map","metadata":{}},{"cell_type":"code","source":"# ---- BUILD VINDR IMAGE MAP ----\ndef build_vindr_path_map(base_path):\n    img_dir = os.path.join(base_path, \"train\")\n    return {file.replace(\".dicom\", \"\"): os.path.join(img_dir, file)\n            for file in os.listdir(img_dir) if file.endswith(\".dicom\")}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-03T17:18:41.343566Z","iopub.execute_input":"2025-09-03T17:18:41.343791Z","iopub.status.idle":"2025-09-03T17:18:41.34777Z","shell.execute_reply.started":"2025-09-03T17:18:41.343774Z","shell.execute_reply":"2025-09-03T17:18:41.347164Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"- VinDr images are DICOM files in train/.\n- Creates a dict from image_id (filename without .dicom) → full path.","metadata":{}},{"cell_type":"code","source":"vindr_image_map = build_vindr_path_map(VINDR_PATH)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-03T17:18:41.348439Z","iopub.execute_input":"2025-09-03T17:18:41.348638Z","iopub.status.idle":"2025-09-03T17:18:41.760747Z","shell.execute_reply.started":"2025-09-03T17:18:41.348615Z","shell.execute_reply":"2025-09-03T17:18:41.75987Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"- Builds that lookup once.","metadata":{}},{"cell_type":"markdown","source":"### VinDr binary labels (nodule/mass vs no finding)","metadata":{}},{"cell_type":"code","source":"# ---- VINDR BINARY LABELS ----\nvindr_binary_df = vindr_df[vindr_df[\"class_name\"].isin([\"Nodule/Mass\", \"No finding\"])].copy()\nvindr_binary_df[\"label\"] = vindr_binary_df[\"class_name\"].apply(lambda x: 1 if x == \"Nodule/Mass\" else 0)\nvindr_binary_df = vindr_binary_df.drop_duplicates(\"image_id\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-03T17:18:41.761624Z","iopub.execute_input":"2025-09-03T17:18:41.761855Z","iopub.status.idle":"2025-09-03T17:18:41.789563Z","shell.execute_reply.started":"2025-09-03T17:18:41.761838Z","shell.execute_reply":"2025-09-03T17:18:41.789049Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"- Filters the annotation table to only two classes, then converts to label ∈ {0,1}.\n- drop_duplicates(\"image_id\") ensures exactly one row per image, since VinDr often has multiple findings per image.","metadata":{}},{"cell_type":"code","source":"pos = vindr_binary_df[vindr_binary_df[\"label\"] == 1]\nneg = vindr_binary_df[vindr_binary_df[\"label\"] == 0].sample(n=len(pos), random_state=42)\nvindr_binary_df = pd.concat([pos, neg]).sample(frac=1, random_state=42).reset_index(drop=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-03T17:18:41.790474Z","iopub.execute_input":"2025-09-03T17:18:41.790761Z","iopub.status.idle":"2025-09-03T17:18:41.798868Z","shell.execute_reply.started":"2025-09-03T17:18:41.790737Z","shell.execute_reply":"2025-09-03T17:18:41.798328Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"- Balances VinDr the same way as NIH: equal positives and negatives, then shuffle.","metadata":{}},{"cell_type":"markdown","source":"### Safety checks","metadata":{}},{"cell_type":"code","source":"# Safety checks (optional)\nassert {'Image Index','Finding Labels'}.issubset(nih_binary_df.columns)\nassert {'image_id','class_name','label'}.issubset(vindr_binary_df.columns)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-03T17:18:41.799486Z","iopub.execute_input":"2025-09-03T17:18:41.79975Z","iopub.status.idle":"2025-09-03T17:18:41.813574Z","shell.execute_reply.started":"2025-09-03T17:18:41.799727Z","shell.execute_reply":"2025-09-03T17:18:41.812831Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"- Early sanity: expected columns exist after filtering.","metadata":{}},{"cell_type":"markdown","source":"### Training config","metadata":{}},{"cell_type":"code","source":"# ---- CONFIG ----\nIMG_SIZE = 224\nBATCH_SIZE = 128\nEPOCHS = 30  # bump if you have more time/TPU\nSEED = 42\nAUTO = tf.data.AUTOTUNE","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-03T17:18:41.814305Z","iopub.execute_input":"2025-09-03T17:18:41.814539Z","iopub.status.idle":"2025-09-03T17:18:41.827906Z","shell.execute_reply.started":"2025-09-03T17:18:41.814518Z","shell.execute_reply":"2025-09-03T17:18:41.82729Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"- Hyperparameters for the input pipeline & training.\n- AUTO lets TF tune prefetch & parallelism.\n- Note: BATCH_SIZE=128 might be aggressive for DICOMs/ViT depending on GPU memory.","metadata":{}},{"cell_type":"markdown","source":"### Helpers: splitting & tf.data","metadata":{}},{"cell_type":"markdown","source":"#### Stratified split","metadata":{}},{"cell_type":"code","source":"def stratified_split(df, label_col='label', test_size=0.2, seed=SEED):\n    train_df, val_df = train_test_split(df, test_size=test_size, stratify=df[label_col], random_state=seed)\n    return train_df.reset_index(drop=True), val_df.reset_index(drop=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-03T17:18:41.830428Z","iopub.execute_input":"2025-09-03T17:18:41.830629Z","iopub.status.idle":"2025-09-03T17:18:41.84386Z","shell.execute_reply.started":"2025-09-03T17:18:41.830615Z","shell.execute_reply":"2025-09-03T17:18:41.843109Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"- Preserves class balance in both train/val by stratifying on the label.","metadata":{}},{"cell_type":"markdown","source":"#### Reading NIH JPGs","metadata":{}},{"cell_type":"code","source":"def _read_nih_image(path):\n    img = cv2.imread(path, cv2.IMREAD_GRAYSCALE).astype(np.float32)\n    img = cv2.resize(img, (IMG_SIZE, IMG_SIZE), interpolation=cv2.INTER_AREA)\n    img = (img - np.mean(img)) / (np.std(img) + 1e-6)   # z-score normalization\n    return np.stack([img, img, img], axis=-1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-03T17:18:41.844777Z","iopub.execute_input":"2025-09-03T17:18:41.845026Z","iopub.status.idle":"2025-09-03T17:18:41.85827Z","shell.execute_reply.started":"2025-09-03T17:18:41.845003Z","shell.execute_reply":"2025-09-03T17:18:41.857739Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"- Loads as grayscale (single channel).\n- Resizes to 224×224, normalizes to [0,1].\n- Replicates the single channel to 3 so models pretrained on RGB can be used without changing weights.","metadata":{}},{"cell_type":"markdown","source":"#### NIH generator","metadata":{}},{"cell_type":"code","source":"def nih_generator(df, image_map, file_col='Image Index', label_col='label'):\n    for _, row in df.iterrows():\n        yield _read_nih_image(image_map[row[file_col]]), np.array(row[label_col], dtype=np.int32)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-03T17:18:41.859089Z","iopub.execute_input":"2025-09-03T17:18:41.859323Z","iopub.status.idle":"2025-09-03T17:18:41.873343Z","shell.execute_reply.started":"2025-09-03T17:18:41.859304Z","shell.execute_reply":"2025-09-03T17:18:41.872739Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"- Python generator that yields (image, label) pairs for NIH rows using the path map.","metadata":{}},{"cell_type":"markdown","source":"#### Reading VinDr DICOMs","metadata":{}},{"cell_type":"code","source":"def _read_vindr_dicom(path):\n    arr = pydicom.dcmread(path).pixel_array.astype(np.float32)\n    arr = cv2.resize(arr, (IMG_SIZE, IMG_SIZE), interpolation=cv2.INTER_AREA)\n    arr = (arr - np.mean(arr)) / (np.std(arr) + 1e-6)   # z-score normalization\n    return np.stack([arr, arr, arr], axis=-1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-03T17:18:41.873918Z","iopub.execute_input":"2025-09-03T17:18:41.874137Z","iopub.status.idle":"2025-09-03T17:18:41.887823Z","shell.execute_reply.started":"2025-09-03T17:18:41.874122Z","shell.execute_reply":"2025-09-03T17:18:41.887103Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"- Reads DICOM pixel data (often 12–16-bit), converts to float, and min-max normalizes to [0,1].\n- Resizes and replicates to 3 channels (same reason as NIH).","metadata":{}},{"cell_type":"markdown","source":"#### VinDr generator","metadata":{}},{"cell_type":"code","source":"def vindr_generator(df, image_map, file_col='image_id', label_col='label'):\n    for _, row in df.iterrows():\n        yield _read_vindr_dicom(image_map[row[file_col]]), np.array(row[label_col], dtype=np.int32)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-03T17:18:41.888826Z","iopub.execute_input":"2025-09-03T17:18:41.888997Z","iopub.status.idle":"2025-09-03T17:18:41.909935Z","shell.execute_reply.started":"2025-09-03T17:18:41.888977Z","shell.execute_reply":"2025-09-03T17:18:41.90933Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### Wrap generators as tf.data","metadata":{}},{"cell_type":"code","source":"def make_tfdata_from_generator(gen_fn, df, **kwargs):\n    sig = (tf.TensorSpec((IMG_SIZE, IMG_SIZE, 3), tf.float32), tf.TensorSpec((), tf.int32))\n    return tf.data.Dataset.from_generator(lambda: gen_fn(df, **kwargs), output_signature=sig)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-03T17:18:41.910751Z","iopub.execute_input":"2025-09-03T17:18:41.910977Z","iopub.status.idle":"2025-09-03T17:18:41.923936Z","shell.execute_reply.started":"2025-09-03T17:18:41.910962Z","shell.execute_reply":"2025-09-03T17:18:41.923308Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"- Converts a Python generator into a TensorFlow Dataset with an explicit signature:\n> - Image: (224, 224, 3), float32\n> - Label: scalar int32\n- This lets TF build an efficient pipeline while still using your custom readers.","metadata":{}},{"cell_type":"markdown","source":"#### Prepare dataset (augment, batch, prefetch)","metadata":{}},{"cell_type":"code","source":"def prepare_dataset(ds, training=True):\n    if training:\n        ds = ds.shuffle(2048, seed=SEED)\n    return ds.batch(BATCH_SIZE).prefetch(AUTO)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-03T17:18:41.924564Z","iopub.execute_input":"2025-09-03T17:18:41.924823Z","iopub.status.idle":"2025-09-03T17:18:41.939833Z","shell.execute_reply.started":"2025-09-03T17:18:41.924802Z","shell.execute_reply":"2025-09-03T17:18:41.939199Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Visualization helper","metadata":{}},{"cell_type":"code","source":"# -----------------------------\n# Visualization\n# -----------------------------\n\ndef visualize_samples(df, image_map, read_fn, title, file_col):\n    samples = df.sample(5, random_state=SEED)\n    plt.figure(figsize=(15, 3))\n    for i, (_, row) in enumerate(samples.iterrows()):\n        img = read_fn(image_map[row[file_col]])\n        plt.subplot(1, 5, i+1)\n        plt.imshow(img[..., 0], cmap='gray')\n        plt.title(f\"Label: {row['label']}\")\n        plt.axis('off')\n    plt.suptitle(title)\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-03T17:18:41.940546Z","iopub.execute_input":"2025-09-03T17:18:41.94081Z","iopub.status.idle":"2025-09-03T17:18:41.954905Z","shell.execute_reply.started":"2025-09-03T17:18:41.940789Z","shell.execute_reply":"2025-09-03T17:18:41.954243Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"nih_train_df, nih_val_df = stratified_split(nih_binary_df)\nvisualize_samples(nih_train_df, nih_image_map, _read_nih_image, 'NIH Training Samples', 'Image Index')\n\nvindr_train_df, vindr_val_df = stratified_split(vindr_binary_df)\nvisualize_samples(vindr_train_df, vindr_image_map, _read_vindr_dicom, 'VinDr Training Samples', 'image_id')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-03T17:18:41.95551Z","iopub.execute_input":"2025-09-03T17:18:41.955745Z","iopub.status.idle":"2025-09-03T17:18:48.66659Z","shell.execute_reply.started":"2025-09-03T17:18:41.955724Z","shell.execute_reply":"2025-09-03T17:18:48.665577Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"- Randomly samples 5 images, reads them with the provided read_fn (NIH or VinDr), and shows the first channel (grayscale) with the binary label.\n- Quick visual sanity check that:\n\n> - Paths resolve,\n> - Normalization looks right,\n> - Labels are balanced and make sense.","metadata":{}},{"cell_type":"markdown","source":"### Models","metadata":{}},{"cell_type":"markdown","source":"#### DenseNet Model","metadata":{}},{"cell_type":"code","source":"def build_densenet_model_transfer(img_size=224, \n                                  unfreeze_layers=30, \n                                  dropout=0.2, \n                                  learning_rate=1e-5):\n    inp = keras.Input((img_size, img_size, 3))\n\n    base = keras.applications.DenseNet169(\n        include_top=False,\n        weights='imagenet',\n        input_tensor=inp,\n        pooling='avg'\n    )\n\n    # Freeze all layers first\n    base.trainable = False\n\n    # Unfreeze last N layers (skip BatchNorm for stability)\n    for layer in base.layers[-unfreeze_layers:]:\n        if not isinstance(layer, layers.BatchNormalization):\n            layer.trainable = True\n\n    x = layers.Dropout(dropout)(base.output)\n    out = layers.Dense(1, activation='sigmoid')(x)\n\n    model = keras.Model(inp, out, name='DenseNet169_binary')\n\n    model.compile(\n        optimizer=keras.optimizers.Adam(learning_rate),\n        loss='binary_crossentropy',\n        metrics=[\n            keras.metrics.Recall(name='recall'),\n            keras.metrics.AUC(name='auc'),\n            keras.metrics.AUC(name='pr_auc', curve='PR'),\n            'accuracy'\n        ]\n    )\n    return model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-03T17:18:48.667499Z","iopub.execute_input":"2025-09-03T17:18:48.667743Z","iopub.status.idle":"2025-09-03T17:18:48.673897Z","shell.execute_reply.started":"2025-09-03T17:18:48.667718Z","shell.execute_reply":"2025-09-03T17:18:48.673226Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### Vision Transformer Model","metadata":{}},{"cell_type":"code","source":"# def build_vit_model_transfer(hf_model_name=\"google/vit-base-patch16-224-in21k\",\n#                              img_size=224,\n#                              unfreeze_layers=4,\n#                              dropout=0.2,\n#                              base_lr=1e-4,\n#                              weight_decay=1e-4,\n#                              warmup_epochs=2,\n#                              total_epochs=20,\n#                              steps_per_epoch=100):\n\n#     inp = keras.Input((img_size, img_size, 3), dtype=tf.float32, name=\"pixel_values\")\n\n#     vit_backbone = TFAutoModel.from_pretrained(hf_model_name)\n#     vit_backbone.trainable = False\n\n#     if hasattr(vit_backbone, \"vit\"):\n#         encoder_layers = vit_backbone.vit.encoder.layer\n#         for layer in encoder_layers[-unfreeze_layers:]:\n#             layer.trainable = True\n\n#     hidden_size = vit_backbone.config.hidden_size\n\n#     def vit_forward(tensor):\n#         tensor = tf.transpose(tensor, perm=[0, 3, 1, 2])  # NHWC → NCHW\n#         outputs = vit_backbone(pixel_values=tensor)\n#         if hasattr(outputs, \"pooler_output\") and outputs.pooler_output is not None:\n#             return outputs.pooler_output\n#         else:\n#             return tf.reduce_mean(outputs.last_hidden_state, axis=1)\n\n#     pooled = layers.Lambda(vit_forward, name=\"vit_backbone\", output_shape=(hidden_size,))(inp)\n\n#     h = layers.Dropout(dropout)(pooled)\n#     h = layers.Dense(256, activation=tf.nn.gelu)(h)\n#     h = layers.Dropout(dropout)(h)\n#     out = layers.Dense(1, activation=\"sigmoid\", name=\"out\")(h)\n\n#     model = keras.Model(inputs=inp, outputs=out, name=\"ViT_transfer_medical\")\n\n#     # ---- Native TF Cosine Decay ----\n#     total_steps = steps_per_epoch * total_epochs\n#     lr_schedule = tf.keras.optimizers.schedules.CosineDecay(\n#         initial_learning_rate=base_lr,\n#         decay_steps=total_steps\n#     )\n\n#     optimizer = tf.keras.optimizers.Adam(\n#         learning_rate=lr_schedule\n#     )\n\n#     model.compile(\n#         optimizer=optimizer,\n#         loss=\"binary_crossentropy\",\n#         metrics=[\n#             keras.metrics.Recall(name=\"recall\"),\n#             keras.metrics.AUC(name=\"auc\"),\n#             keras.metrics.AUC(name=\"pr_auc\", curve=\"PR\"),\n#             \"accuracy\"\n#         ]\n#     )\n\n#     return model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-03T17:18:48.674543Z","iopub.execute_input":"2025-09-03T17:18:48.674805Z","iopub.status.idle":"2025-09-03T17:18:48.694741Z","shell.execute_reply.started":"2025-09-03T17:18:48.67478Z","shell.execute_reply":"2025-09-03T17:18:48.694227Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# 5. Plotting helpers\n# =========================\ndef plot_training_curves(history, metrics=None):\n    hist = history.history\n    if metrics is None:\n        metrics = [m for m in hist.keys() if not m.startswith('val_') and m != 'loss']\n\n    plt.figure(figsize=(15, 4))\n    plt.subplot(1, len(metrics) + 1, 1)\n    plt.plot(hist['loss'], label='Train Loss')\n    plt.plot(hist['val_loss'], label='Val Loss')\n    plt.xlabel('Epoch'); plt.ylabel('Loss'); plt.title('Loss'); plt.legend()\n\n    for i, metric in enumerate(metrics, start=2):\n        plt.subplot(1, len(metrics) + 1, i)\n        plt.plot(hist[metric], label=f'Train {metric}')\n        plt.plot(hist[f'val_{metric}'], label=f'Val {metric}')\n        plt.xlabel('Epoch'); plt.ylabel(metric)\n        plt.title(metric.capitalize()); plt.legend()\n\n    plt.tight_layout()\n    plt.show()\n\ndef predict_on_tf_dataset(model, ds):\n    y_true, y_pred = [], []\n    for batch_x, batch_y in ds.unbatch().batch(256):\n        preds = model.predict(batch_x, verbose=0)\n        y_true.append(batch_y.numpy().ravel())\n        y_pred.append(preds.ravel())\n    return np.concatenate(y_true), np.concatenate(y_pred)\n\ndef plot_roc_pr_curves(y_true, y_pred):\n    fpr, tpr, _ = roc_curve(y_true, y_pred)\n    roc_auc = roc_auc_score(y_true, y_pred)\n    prec, rec, _ = precision_recall_curve(y_true, y_pred)\n    pr_auc = average_precision_score(y_true, y_pred)\n\n    print(f\"Validation ROC AUC: {roc_auc:.4f}\")\n    print(f\"Validation PR-AUC: {pr_auc:.4f}\")\n\n    plt.figure(figsize=(12,5))\n    plt.subplot(1,2,1)\n    plt.plot(fpr, tpr, label=f\"ROC AUC = {roc_auc:.4f}\")\n    plt.plot([0,1],[0,1], '--', alpha=0.4)\n    plt.xlabel(\"False Positive Rate\"); plt.ylabel(\"True Positive Rate\")\n    plt.title(\"ROC Curve\"); plt.legend()\n\n    plt.subplot(1,2,2)\n    plt.plot(rec, prec, label=f\"PR AUC = {pr_auc:.4f}\")\n    plt.xlabel(\"Recall\"); plt.ylabel(\"Precision\")\n    plt.title(\"Precision-Recall Curve\"); plt.legend()\n\n    plt.tight_layout(); plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-03T17:18:48.695558Z","iopub.execute_input":"2025-09-03T17:18:48.695808Z","iopub.status.idle":"2025-09-03T17:18:48.71784Z","shell.execute_reply.started":"2025-09-03T17:18:48.695794Z","shell.execute_reply":"2025-09-03T17:18:48.717198Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def evaluate_model(y_true, y_pred, threshold=0.5):\n    \"\"\"\n    Extended evaluation: Classification report,\n    specificity, confusion matrix, calibration curve\n    \"\"\"\n\n    # --- Thresholding ---\n    y_pred_bin = (y_pred >= threshold).astype(int)\n\n    # --- Classification report (precision, recall, F1) ---\n    print(\"\\nClassification Report:\")\n    print(classification_report(y_true, y_pred_bin, digits=4))\n\n    # --- Specificity ---\n    tn, fp, fn, tp = confusion_matrix(y_true, y_pred_bin).ravel()\n    specificity = tn / (tn + fp)\n    print(f\"Specificity: {specificity:.4f}\")\n\n    # --- Confusion Matrix ---\n    disp = ConfusionMatrixDisplay.from_predictions(\n        y_true, y_pred_bin,\n        display_labels=[\"No Cancer\", \"Cancer\"],\n        cmap=\"Blues\"\n    )\n    disp.ax_.set_title(\"Confusion Matrix\")\n    plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-03T17:18:48.718566Z","iopub.execute_input":"2025-09-03T17:18:48.718821Z","iopub.status.idle":"2025-09-03T17:18:48.737562Z","shell.execute_reply.started":"2025-09-03T17:18:48.718802Z","shell.execute_reply":"2025-09-03T17:18:48.736885Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ======= DenseNet on VinDr =======\n# print(\"\\n=== Training DenseNet on VinDr ===\")\ntrain_df, val_df = stratified_split(vindr_binary_df, test_size=0.2)\ntrain_ds = make_tfdata_from_generator(vindr_generator, train_df, image_map=vindr_image_map)\nval_ds = make_tfdata_from_generator(vindr_generator, val_df, image_map=vindr_image_map)\ntrain_ds = prepare_dataset(train_ds, training=True)\nval_ds = prepare_dataset(val_ds, training=False)\n\n# densenet_vindr = build_densenet_model_transfer(unfreeze_layers=100, learning_rate=0.00001)\n\n# # Early stopping callback\n# early_stop = EarlyStopping(\n#     monitor=\"val_loss\", \n#     patience=3, \n#     restore_best_weights=True,\n#     verbose=1\n# )\n\n# history_dn_vindr = densenet_vindr.fit(\n#     train_ds,\n#     validation_data=val_ds,\n#     epochs=EPOCHS,\n#     callbacks=[early_stop]\n# )\n\n# plot_training_curves(history_dn_vindr, metrics=['accuracy','recall','auc','pr_auc'])\n\n# # Save the trained model\n# MODEL_PATH = \"densenet_vindr_model.h5\"\n# densenet_vindr.save(MODEL_PATH)\n# MODEL_PATH = \"densenet_vindr_model\"\n# densenet_vindr.save(MODEL_PATH)  # Saves in TensorFlow SavedModel format\n\n# print(f\"Model saved to {MODEL_PATH}\")\n\n# Load the model back\nMODEL_PATH = \"/kaggle/input/densenet_vindr_model/tensorflow2/default/1/densenet_vindr_model.h5\"\n# Rebuild model from code\ndensenet_vindr = build_densenet_model_transfer(unfreeze_layers=100, learning_rate=1e-5)\n\n# Load only weights\ndensenet_vindr.load_weights(MODEL_PATH)\nprint(\"Model reloaded successfully.\")\n\ny_true, y_pred = predict_on_tf_dataset(densenet_vindr, val_ds)\n\n# ROC & PR curves\nplot_roc_pr_curves(y_true, y_pred)\n\n# Evaluate across thresholds\nfor thr in [0.3, 0.4, 0.5, 0.6, 0.7]:\n    print(f\"\\nThreshold {thr}\")\n    evaluate_model(y_true, y_pred, threshold=thr)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-03T17:18:48.738326Z","iopub.execute_input":"2025-09-03T17:18:48.739001Z","iopub.status.idle":"2025-09-03T17:24:10.198555Z","shell.execute_reply.started":"2025-09-03T17:18:48.738978Z","shell.execute_reply":"2025-09-03T17:24:10.197874Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### Add Explainability with Grad-CAM","metadata":{}},{"cell_type":"markdown","source":"Grad-CAM highlights the image regions most influential for the model’s prediction.","metadata":{}},{"cell_type":"code","source":"def make_gradcam_heatmap(img_array, model, last_conv_layer_name):\n    \"\"\"\n    img_array: np.array of shape (1, H, W, 3), preprocessed image\n    model: trained keras model\n    last_conv_layer_name: name of last conv layer\n    \"\"\"\n    # Build a sub-model that maps input -> (conv outputs, predictions)\n    grad_model = keras.models.Model(\n        inputs=model.input,\n        outputs=[model.get_layer(last_conv_layer_name).output, model.output]\n    )\n\n    with tf.GradientTape() as tape:\n        conv_outputs, predictions = grad_model(img_array)\n        pred_class = predictions[:, 0]  # binary classification\n\n    # Compute gradients of class score wrt conv feature maps\n    grads = tape.gradient(pred_class, conv_outputs)\n    pooled_grads = tf.reduce_mean(grads, axis=(0, 1, 2))\n\n    conv_outputs = conv_outputs[0]\n    heatmap = conv_outputs @ pooled_grads[..., tf.newaxis]\n    heatmap = tf.squeeze(heatmap)\n\n    # Normalize\n    heatmap = tf.maximum(heatmap, 0) / (tf.reduce_max(heatmap) + 1e-8)\n    return heatmap.numpy()\n\ndef overlay_heatmap(img, heatmap, alpha=0.4, cmap='jet'):\n    heatmap = np.uint8(255 * heatmap)\n    cmap = cm.get_cmap(cmap)\n    colored = cmap(np.arange(256))[:, :3]\n    heatmap_colored = colored[heatmap]\n\n    heatmap_colored = Image.fromarray((heatmap_colored * 255).astype(np.uint8)).resize(img.size)\n    overlayed = Image.blend(img.convert(\"RGBA\"), heatmap_colored.convert(\"RGBA\"), alpha)\n    return overlayed\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-03T17:24:10.199408Z","iopub.execute_input":"2025-09-03T17:24:10.199884Z","iopub.status.idle":"2025-09-03T17:24:10.207172Z","shell.execute_reply.started":"2025-09-03T17:24:10.199858Z","shell.execute_reply":"2025-09-03T17:24:10.206503Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(type(densenet_vindr.input))\nprint(densenet_vindr.input)\nprint(densenet_vindr.layers[0])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-03T17:24:10.207926Z","iopub.execute_input":"2025-09-03T17:24:10.208331Z","iopub.status.idle":"2025-09-03T17:24:10.229377Z","shell.execute_reply.started":"2025-09-03T17:24:10.208308Z","shell.execute_reply":"2025-09-03T17:24:10.228748Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_gradcam_samples(model, val_ds, last_conv_layer_name, n_samples=6, threshold=0.5):\n    \"\"\"\n    Plot Original / Grad-CAM / Overlay for n_samples from val_ds,\n    including True label and Predicted label (with probability).\n\n    Args:\n        model: trained Keras model (binary output with sigmoid).\n        val_ds: tf.data.Dataset yielding (image, label).\n        last_conv_layer_name: str, name of last conv layer for Grad-CAM.\n        n_samples: int, number of samples to visualize.\n        threshold: float, decision threshold for binary class (default 0.5).\n    \"\"\"\n    # Collect exactly n_samples regardless of dataset batch size\n    sample_imgs, sample_labels = [], []\n    for img, label in val_ds.unbatch().take(n_samples):\n        sample_imgs.append(img.numpy())\n        sample_labels.append(int(label.numpy()))\n\n    rows, cols = len(sample_imgs), 3  # one row per sample: Original | Heatmap | Overlay\n    plt.figure(figsize=(5 * cols, 4 * rows))\n\n    for idx, (img, true_label) in enumerate(zip(sample_imgs, sample_labels)):\n        # Prepare input\n        x_in = np.expand_dims(img, axis=0)\n\n        # Predict probability and class\n        prob = model(x_in, training=False).numpy().squeeze()\n        prob = float(np.array(prob).squeeze())  # robust squeeze\n        pred_label = int(prob >= threshold)\n\n        # Grad-CAM heatmap\n        heatmap = make_gradcam_heatmap(x_in, model, last_conv_layer_name)\n\n        # Overlay\n        img_pil = Image.fromarray((img * 255).astype(np.uint8))\n        overlayed_img = overlay_heatmap(img_pil, heatmap)\n\n        # --- Plotting ---\n        # Original\n        ax = plt.subplot(rows, cols, idx * 3 + 1)\n        ax.imshow(img)\n        ax.set_title(f\"Original\\nTrue={true_label}\")\n        ax.axis(\"off\")\n\n        # Heatmap\n        ax = plt.subplot(rows, cols, idx * 3 + 2)\n        ax.imshow(heatmap, cmap=\"jet\")\n        ax.set_title(\"Grad-CAM Heatmap\")\n        ax.axis(\"off\")\n\n        # Overlay with labels\n        ax = plt.subplot(rows, cols, idx * 3 + 3)\n        ax.imshow(overlayed_img)\n        ax.set_title(f\"Overlay\\nTrue={true_label} | Pred={pred_label} (p={prob:.2f})\")\n        ax.axis(\"off\")\n\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-03T17:43:15.164279Z","iopub.execute_input":"2025-09-03T17:43:15.165055Z","iopub.status.idle":"2025-09-03T17:43:15.173534Z","shell.execute_reply.started":"2025-09-03T17:43:15.165029Z","shell.execute_reply":"2025-09-03T17:43:15.172849Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plot_gradcam_samples(densenet_vindr, val_ds, \"conv5_block32_concat\", n_samples=6, threshold=0.5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-03T17:43:38.169689Z","iopub.execute_input":"2025-09-03T17:43:38.169987Z","iopub.status.idle":"2025-09-03T17:46:54.866202Z","shell.execute_reply.started":"2025-09-03T17:43:38.169968Z","shell.execute_reply":"2025-09-03T17:46:54.865188Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_gradcam_misclassified(model, val_ds, last_conv_layer_name, max_samples=10, threshold=0.5):\n    \"\"\"\n    Plot Grad-CAM for only misclassified samples from val_ds.\n    \n    Args:\n        model: trained Keras model (binary output with sigmoid).\n        val_ds: tf.data.Dataset yielding (image, label).\n        last_conv_layer_name: str, name of last conv layer for Grad-CAM.\n        max_samples: int, maximum number of misclassified samples to show.\n        threshold: float, decision threshold for binary class (default 0.5).\n    \"\"\"\n    sample_imgs, true_labels, pred_labels, probs = [], [], [], []\n\n    # Go through dataset and collect misclassified samples\n    for img, label in val_ds.unbatch():\n        x_in = np.expand_dims(img.numpy(), axis=0)\n\n        # Predict\n        prob = model(x_in, training=False).numpy().squeeze()\n        prob = float(np.array(prob).squeeze())  # robust squeeze\n        pred_label = int(prob >= threshold)\n        true_label = int(label.numpy())\n\n        if pred_label != true_label:  # misclassified\n            sample_imgs.append(img.numpy())\n            true_labels.append(true_label)\n            pred_labels.append(pred_label)\n            probs.append(prob)\n\n        if len(sample_imgs) >= max_samples:\n            break\n\n    if len(sample_imgs) == 0:\n        print(\"No misclassified samples found in the given dataset.\")\n        return\n\n    rows, cols = len(sample_imgs), 3\n    plt.figure(figsize=(5 * cols, 4 * rows))\n\n    for idx, (img, t, p, prob) in enumerate(zip(sample_imgs, true_labels, pred_labels, probs)):\n        x_in = np.expand_dims(img, axis=0)\n        heatmap = make_gradcam_heatmap(x_in, model, last_conv_layer_name)\n        img_pil = Image.fromarray((img * 255).astype(np.uint8))\n        overlayed_img = overlay_heatmap(img_pil, heatmap)\n\n        # Original\n        ax = plt.subplot(rows, cols, idx * 3 + 1)\n        ax.imshow(img)\n        ax.set_title(f\"Original\\nTrue={t}\")\n        ax.axis(\"off\")\n\n        # Heatmap\n        ax = plt.subplot(rows, cols, idx * 3 + 2)\n        ax.imshow(heatmap, cmap=\"jet\")\n        ax.set_title(\"Grad-CAM Heatmap\")\n        ax.axis(\"off\")\n\n        # Overlay with misclassification info\n        ax = plt.subplot(rows, cols, idx * 3 + 3)\n        ax.imshow(overlayed_img)\n        ax.set_title(f\"Overlay\\nTrue={t} | Pred={p} (p={prob:.2f})\")\n        ax.axis(\"off\")\n\n    plt.tight_layout()\n    plt.show()\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-03T17:53:14.129541Z","iopub.execute_input":"2025-09-03T17:53:14.130129Z","iopub.status.idle":"2025-09-03T17:53:14.141341Z","shell.execute_reply.started":"2025-09-03T17:53:14.130104Z","shell.execute_reply":"2025-09-03T17:53:14.140357Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Show 5 misclassified samples with Grad-CAM\nplot_gradcam_misclassified(densenet_vindr, val_ds, \"conv5_block32_concat\", max_samples=5, threshold=0.5)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-03T17:53:14.768765Z","iopub.execute_input":"2025-09-03T17:53:14.769261Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}