jojo9956/Breast_Cancer_Detection_using_Computer_Vision
0
1import streamlit as st
2import os
3import numpy as np
4import joblib
5import lime
6import lime.lime_image
7import matplotlib.pyplot as plt
8from skimage.io import imread
9from skimage.transform import resize
10from skimage.segmentation import mark_boundaries, slic
11from process_data_non_dl import extract_lbp_features # Import LBP feature extractor
12from tensorflow.keras.models import load_model
13import cv2
14import tensorflow as tf
15import io
16
17################################################################Deep learning part#####################################################################################################################################################
18ml_model_path = 'E:/540_dl/Computer_Vision_DL_Proj/models/'
19
20ml_model = joblib.load()
21
22DL_MODEL_PATH = "E:/540_dl/dl_model"
23
24dl_model = load_model(DL_MODEL_PATH)
25
26def df_preprocess(image, img_size=(128, 128)):
27 img = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)
28 img = cv2.resize(img, img_size)
29 img = img / 255.0
30
31 return np.expand_dims(img, axis=0)
32
33def DL_classify(image, model):
34 img = df_preprocess(image)
35 predictions = model.predict(img)
36 predicted_class = np.argmax(predictions, axis=-1)[0]
37 return "Malignant (Yes)" if predicted_class == 1 else "Benign (No)"
38
39def normalize_map(heatmap):
40 """Normalize a heatmap to the [0,1] range."""
41 heatmap -= heatmap.min()
42 denom = (heatmap.max() - heatmap.min()) + 1e-8
43 heatmap /= denom
44 return heatmap
45
46def overlay_heatmap(img, heatmap, alpha=0.5, cmap='jet'):
47 """
48 Overlay `heatmap` on `img`.
49 """
50 colormap = plt.cm.get_cmap(cmap)
51 colored_heatmap = colormap(heatmap)[..., :3]
52 overlay = (1 - alpha) * img + alpha * colored_heatmap
53 overlay = np.clip(overlay, 0, 1)
54 return overlay
55
56def integrated_gradients(model, x, baseline=None, steps=50, class_idx=0):
57 """
58 Computes Integrated Gradients
59 """
60 x = tf.cast(x, tf.float32)
61 if baseline is None:
62 baseline = tf.zeros_like(x)
63 else:
64 baseline = tf.cast(baseline, tf.float32)
65
66 B, H, W, C = x.shape
67
68 alphas = tf.reshape(tf.linspace(0.0, 1.0, steps+1), [steps+1, 1, 1, 1, 1])
69 x_expanded = tf.expand_dims(x, axis=0)
70 baseline_expanded = tf.expand_dims(baseline, axis=0)
71 interpolated = baseline_expanded + alphas * (x_expanded - baseline_expanded)
72
73 with tf.GradientTape() as tape:
74 interpolated_reshaped = tf.reshape(interpolated, [(steps+1)*B, H, W, C])
75 tape.watch(interpolated_reshaped)
76 preds = model(interpolated_reshaped)
77 preds_for_class = preds[:, class_idx]
78 loss = tf.reduce_sum(preds_for_class)
79
80 grads = tape.gradient(loss, interpolated_reshaped)
81 if grads is None:
82 raise ValueError("Gradient is None. Not differentiable.")
83
84 grads_reshaped = tf.reshape(grads, [steps+1, B, H, W, C])
85 avg_grads = tf.reduce_mean(grads_reshaped[1:], axis=0)
86 ig = (x - baseline) * avg_grads
87 return ig.numpy()
88
89
90def DL_explainability(model, image, class_idx=1):
91 img_batch = df_preprocess(image)
92 ig_map = integrated_gradients(model, img_batch, class_idx=class_idx)
93
94 original_img = img_batch[0]
95 normalized_img = (original_img - original_img.min()) / (original_img.max() - original_img.min() + 1e-8)
96
97 ig_map_single = ig_map[0]
98 ig_map_2d = np.mean(ig_map_single, axis=-1)
99 ig_map_norm = normalize_map(ig_map_2d)
100 overlay_img = overlay_heatmap(normalized_img, ig_map_norm, alpha=0.5, cmap='jet')
101
102 fig, axes = plt.subplots(1, 3, figsize=(15, 5))
103
104 # Original image
105 axes[0].imshow(normalized_img)
106 axes[0].set_title("Original Image")
107 axes[0].axis("off")
108
109 # IG heatmap
110 im = axes[1].imshow(ig_map_norm, cmap="jet")
111 axes[1].set_title("IG Map (2D Mean)")
112 plt.colorbar(im, ax=axes[1])
113 axes[1].axis("off")
114
115 # Overlay images
116 axes[2].imshow(overlay_img)
117 axes[2].set_title("Overlay")
118 axes[2].axis("off")
119
120 # Adjust layout
121 plt.tight_layout()
122
123 # load image
124 img_buffer = io.BytesIO()
125 plt.savefig(img_buffer, format="png")
126 img_buffer.seek(0)
127 plt.close(fig)
128
129 return img_buffer
130
131################################################################Deep learning part#####################################################################################################################################################
132
133
134
135
136################################################################machine learning part#####################################################################################################################################################
137
138MODEL_PATH = "models/RandomForest_breakhis.pkl"
139ml_model = joblib.load(MODEL_PATH)
140
141# ✅ Function to preprocess image and extract LBP features
142def process_image(image):
143 image_resized = resize(image, (128, 128)) # Resize to required size
144 features = extract_lbp_features(image_resized)
145 return np.array(features).reshape(1, -1) # Convert to 2D array
146
147# ✅ Function to predict using model
148def classify_image(image, model):
149 features = process_image(image)
150 prediction = model.predict(features)[0]
151 return "Malignant (Yes)" if prediction == 1 else "Benign (No)"
152
153# ✅ **Better LIME prediction function**
154def model_predict(image_batch):
155 processed_features = []
156 for img in image_batch:
157 img_resized = resize(img, (128, 128))
158 features = extract_lbp_features(img_resized)
159 processed_features.append(features)
160
161 return ml_model.predict_proba(np.array(processed_features)) # ✅ LIME now passes LBP features!
162
163# ✅ **Improved LIME Explanation**
164def generate_lime_explanation(image):
165 explainer = lime.lime_image.LimeImageExplainer()
166
167 # ✅ Use SLIC segmentation with more superpixels
168 segmentation_fn = lambda x: slic(x, n_segments=300, compactness=15, sigma=1)
169
170 explanation = explainer.explain_instance(
171 image,
172 model_predict,
173 top_labels=1,
174 hide_color=0,
175 num_samples=1000,
176 segmentation_fn=segmentation_fn # Use improved segmentation
177 )
178
179 # ✅ **Extract important positive regions**
180 temp_pos, mask_pos = explanation.get_image_and_mask(
181 explanation.top_labels[0], positive_only=True, num_features=15, hide_rest=False
182 )
183
184 # ✅ **Extract both positive & negative regions**
185 temp_neg, mask_neg = explanation.get_image_and_mask(
186 explanation.top_labels[0], positive_only=False, num_features=15, hide_rest=False
187 )
188
189 return temp_pos, mask_pos, temp_neg, mask_neg
190
191################################################################machine learning part#####################################################################################################################################################
192
193st.title("🎯 Machine Learning vs Deep Learning Classification and Explainability for Breast Cancer")
194uploaded_file = st.file_uploader("Upload a Histopathological Image", type=["jpg", "png", "jpeg"])
195
196if uploaded_file is not None:
197
198 image = imread(uploaded_file)
199
200 # Create two columns
201 col1, col2 = st.columns(2)
202
203 # Machine learning
204 with col1:
205 st.header("Machine Learning")
206
207 # ✅ Predict Classification
208 prediction = classify_image(image, ml_model)
209 st.subheader(f"Prediction: **{prediction}**")
210
211 # ✅ Generate LIME explanation
212 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 Only
219 axes[0].imshow(mark_boundaries(temp_pos, mask_pos, color=(1, 1, 0), mode='thick')) # Yellow boundaries
220 axes[0].axis("off")
221 axes[0].set_title("LIME: Important Regions", fontsize=12)
222
223 # ✅ Positive & Negative Regions
224 axes[1].imshow(mark_boundaries(temp_neg, mask_neg, color=(1, 1, 0), mode='thick')) # Yellow boundaries
225 axes[1].axis("off")
226 axes[1].set_title("LIME: Positive & Negative Regions", fontsize=12)
227
228 st.pyplot(fig) # Show the plot in Streamlit
229
230 # Deep learning
231 with col2:
232 st.header("Deep Learning")
233
234 # Predict Classification
235 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 