Team Ai
Apppublic

jojo9956/Breast_Cancer_Detection_using_Computer_Vision

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
streamlit_1.py244 linesDownload Raw Back to root
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