jojo9956/Breast_Cancer_Detection_using_Computer_Vision
0
1import streamlit as st2import os3import numpy as np4import joblib5import lime6import lime.lime_image7import matplotlib.pyplot as plt8from skimage.io import imread9from skimage.transform import resize10from skimage.segmentation import mark_boundaries, slic11from process_data_non_dl import extract_lbp_features # Import LBP feature extractor12from tensorflow.keras.models import load_model13import cv214import tensorflow as tf15import io16 17################################################################Deep learning part#####################################################################################################################################################18 19DL_MODEL_PATH = "models/RandomForest_breakhis.pkl"20 21dl_model = load_model(DL_MODEL_PATH)22 23def df_preprocess(image, img_size=(128, 128)):24 img = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)25 img = cv2.resize(img, img_size)26 img = img / 255.027 28 return np.expand_dims(img, axis=0)29 30def DL_classify(image, model):31 img = df_preprocess(image)32 predictions = model.predict(img)33 predicted_class = np.argmax(predictions, axis=-1)[0]34 return "Malignant (Yes)" if predicted_class == 1 else "Benign (No)"35 36def normalize_map(heatmap):37 """Normalize a heatmap to the [0,1] range."""38 heatmap -= heatmap.min()39 denom = (heatmap.max() - heatmap.min()) + 1e-840 heatmap /= denom41 return heatmap42 43def overlay_heatmap(img, heatmap, alpha=0.5, cmap='jet'):44 """45 Overlay `heatmap` on `img`.46 """47 colormap = plt.cm.get_cmap(cmap)48 colored_heatmap = colormap(heatmap)[..., :3]49 overlay = (1 - alpha) * img + alpha * colored_heatmap50 overlay = np.clip(overlay, 0, 1)51 return overlay52 53def integrated_gradients(model, x, baseline=None, steps=50, class_idx=0):54 """55 Computes Integrated Gradients 56 """57 x = tf.cast(x, tf.float32)58 if baseline is None:59 baseline = tf.zeros_like(x)60 else:61 baseline = tf.cast(baseline, tf.float32)62 63 B, H, W, C = x.shape64 65 alphas = tf.reshape(tf.linspace(0.0, 1.0, steps+1), [steps+1, 1, 1, 1, 1])66 x_expanded = tf.expand_dims(x, axis=0)67 baseline_expanded = tf.expand_dims(baseline, axis=0)68 interpolated = baseline_expanded + alphas * (x_expanded - baseline_expanded)69 70 with tf.GradientTape() as tape:71 interpolated_reshaped = tf.reshape(interpolated, [(steps+1)*B, H, W, C])72 tape.watch(interpolated_reshaped)73 preds = model(interpolated_reshaped)74 preds_for_class = preds[:, class_idx]75 loss = tf.reduce_sum(preds_for_class)76 77 grads = tape.gradient(loss, interpolated_reshaped)78 if grads is None:79 raise ValueError("Gradient is None. Not differentiable.")80 81 grads_reshaped = tf.reshape(grads, [steps+1, B, H, W, C])82 avg_grads = tf.reduce_mean(grads_reshaped[1:], axis=0)83 ig = (x - baseline) * avg_grads84 return ig.numpy()85 86 87def DL_explainability(model, image, class_idx=1):88 img_batch = df_preprocess(image)89 ig_map = integrated_gradients(model, img_batch, class_idx=class_idx)90 91 original_img = img_batch[0]92 normalized_img = (original_img - original_img.min()) / (original_img.max() - original_img.min() + 1e-8)93 94 ig_map_single = ig_map[0]95 ig_map_2d = np.mean(ig_map_single, axis=-1)96 ig_map_norm = normalize_map(ig_map_2d)97 overlay_img = overlay_heatmap(normalized_img, ig_map_norm, alpha=0.5, cmap='jet')98 99 fig, axes = plt.subplots(1, 3, figsize=(15, 5))100 101 # Original image102 axes[0].imshow(normalized_img)103 axes[0].set_title("Original Image")104 axes[0].axis("off")105 106 # IG heatmap107 im = axes[1].imshow(ig_map_norm, cmap="jet")108 axes[1].set_title("IG Map (2D Mean)")109 plt.colorbar(im, ax=axes[1])110 axes[1].axis("off")111 112 # Overlay images113 axes[2].imshow(overlay_img)114 axes[2].set_title("Overlay")115 axes[2].axis("off")116 117 # Adjust layout 118 plt.tight_layout()119 120 # load image121 img_buffer = io.BytesIO()122 plt.savefig(img_buffer, format="png")123 img_buffer.seek(0) 124 plt.close(fig) 125 126 return img_buffer127 128################################################################Deep learning part#####################################################################################################################################################129 130 131 132 133################################################################machine learning part#####################################################################################################################################################134 135MODEL_PATH = "models/RandomForest_breakhis.pkl"136ml_model = joblib.load(MODEL_PATH)137 138# ✅ Function to preprocess image and extract LBP features139def process_image(image):140 image_resized = resize(image, (128, 128)) # Resize to required size141 features = extract_lbp_features(image_resized)142 return np.array(features).reshape(1, -1) # Convert to 2D array143 144# ✅ Function to predict using model145def classify_image(image, model):146 features = process_image(image)147 prediction = model.predict(features)[0]148 return "Malignant (Yes)" if prediction == 1 else "Benign (No)"149 150# ✅ **Better LIME prediction function**151def model_predict(image_batch):152 processed_features = []153 for img in image_batch:154 img_resized = resize(img, (128, 128))155 features = extract_lbp_features(img_resized)156 processed_features.append(features)157 158 return ml_model.predict_proba(np.array(processed_features)) # ✅ LIME now passes LBP features!159 160# ✅ **Improved LIME Explanation**161def generate_lime_explanation(image):162 explainer = lime.lime_image.LimeImageExplainer()163 164 # ✅ Use SLIC segmentation with more superpixels165 segmentation_fn = lambda x: slic(x, n_segments=300, compactness=15, sigma=1) 166 167 explanation = explainer.explain_instance(168 image, 169 model_predict, 170 top_labels=1, 171 hide_color=0, 172 num_samples=1000,173 segmentation_fn=segmentation_fn # Use improved segmentation174 )175 176 # ✅ **Extract important positive regions**177 temp_pos, mask_pos = explanation.get_image_and_mask(178 explanation.top_labels[0], positive_only=True, num_features=15, hide_rest=False179 )180 181 # ✅ **Extract both positive & negative regions**182 temp_neg, mask_neg = explanation.get_image_and_mask(183 explanation.top_labels[0], positive_only=False, num_features=15, hide_rest=False184 )185 186 return temp_pos, mask_pos, temp_neg, mask_neg187 188################################################################machine learning part#####################################################################################################################################################189 190 191 192def main():193 st.title("🎯 Machine Learning vs Deep Learning Classification and Explainability for Breast Cancer")194 uploaded_file = st.file_uploader("Upload a Histopathological Image", type=["jpg", "png", "jpeg"])195 196 if uploaded_file is not None:197 198 image = imread(uploaded_file)199 200 # Create two columns201 col1, col2 = st.columns(2)202 203 # Machine learning204 with col1:205 st.header("Machine Learning")206 207 # ✅ Predict Classification208 prediction = classify_image(image)209 st.subheader(f"Prediction: **{prediction}**")210 211 # ✅ Generate LIME explanation212 with st.spinner("Generating LIME Explainability..."):213 temp_pos, mask_pos, temp_neg, mask_neg = generate_lime_explanation(image)214 215 # ✅ Display LIME Explanation with **Improved Masks**216 fig, axes = plt.subplots(1, 2, figsize=(12, 6))217 218 # ✅ Important Regions Only219 axes[0].imshow(mark_boundaries(temp_pos, mask_pos, color=(1, 1, 0), mode='thick')) # Yellow boundaries220 axes[0].axis("off")221 axes[0].set_title("LIME: Important Regions", fontsize=12)222 223 # ✅ Positive & Negative Regions224 axes[1].imshow(mark_boundaries(temp_neg, mask_neg, color=(1, 1, 0), mode='thick')) # Yellow boundaries225 axes[1].axis("off")226 axes[1].set_title("LIME: Positive & Negative Regions", fontsize=12)227 228 st.pyplot(fig) # Show the plot in Streamlit229 230 # Deep learning231 with col2:232 st.header("Deep Learning")233 234 # Predict Classification235 prediction = DL_classify(image, dl_model)236 st.subheader(f"Prediction: **{prediction}**")237 238 with st.spinner("Generating Integrated Gradients Explainability..."):239 X_image = DL_explainability(dl_model, image)240 241 242 st.image(X_image, caption="Integrated Gradients Visualization", use_column_width=True)243 244 245 246if __name__ == "__main__":247 main()