Team Ai
Apppublic

jojo9956/Breast_Cancer_Detection_using_Computer_Vision

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