Spaces:
Runtime error
Runtime error
File size: 7,357 Bytes
fe66586 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 | import streamlit as st
import cv2
import numpy as np
import pandas as pd
import joblib
import os
import matplotlib.pyplot as plt
import lime
import lime.lime_tabular
from src import features, config
st.set_page_config(page_title="DermAI Classification", layout="centered")
# Custom styling to widen the center container
st.markdown(
"""
<style>
.block-container {
max-width: 1000px;
padding-top: 2rem;
padding-bottom: 2rem;
}
</style>
""",
unsafe_allow_html=True
)
st.title("🔬 Skin Lesion Classification")
st.markdown("""
This system uses **Classical Machine Vision** techniques (CLAHE, Otsu Thresholding, Morphology)
to classify skin lesions.
""")
try:
model_path = os.path.join(config.MODEL_DIR, 'skin_cancer_model.pkl')
scaler_path = os.path.join(config.MODEL_DIR, 'scaler.pkl')
classes_path = os.path.join(config.MODEL_DIR, 'classes.pkl')
model = joblib.load(model_path)
scaler = joblib.load(scaler_path)
classes = joblib.load(classes_path)
st.success("System Ready: Model Loaded Successfully")
except FileNotFoundError:
st.error("Model files not found. Please run 'train_main.py' first.")
st.stop()
uploaded_file = st.file_uploader("Choose a dermoscopy image...", type=["jpg", "jpeg", "png"])
if uploaded_file is not None:
file_bytes = np.asarray(bytearray(uploaded_file.read()), dtype=np.uint8)
image = cv2.imdecode(file_bytes, 1)
col1, col2 = st.columns(2)
with col1:
st.image(image, channels="BGR", caption="Uploaded Image", use_container_width=True)
with st.spinner('Extracting Handcrafted Features...'):
feat_vector = features.extract_all_features_pipeline(image)
# Reshape for model input
feat_vector_reshaped = feat_vector.reshape(1, -1)
feat_scaled = scaler.transform(feat_vector_reshaped)
# Predict
probs = model.predict_proba(feat_scaled)
pred_idx = np.argmax(probs)
pred_label = classes[pred_idx]
with col2:
st.subheader(f"Prediction: **{pred_label}**")
st.metric("Confidence", f"{probs[0][pred_idx] * 100:.2f}%")
# --- Charts ---
st.subheader("Class Probabilities")
chart_data = pd.DataFrame({"Class": classes, "Probability": probs[0] * 100})
st.bar_chart(chart_data.set_index("Class"))
with st.expander("Abbreviation information"):
df_legend = pd.DataFrame(config.LEGEND_DATA)
st.dataframe(
df_legend,
column_config={
"More Info": st.column_config.LinkColumn(
"More",
help="Click to visit Wikipedia page",
display_text="🔍"
)
},
hide_index=True,
use_container_width=True
)
# --- LIME EXPLANATION (Local XAI) ---
st.divider()
st.subheader("Explainable AI (LIME)")
st.write(f"#### Why was this specific image classified as **{pred_label}**?")
st.write(
"The charts below show which features supported (Green) or contradicted (Red) the decision for **EACH** possible class.")
try:
# 1. Load the training sample (needed to initialize LIME)
train_sample_path = os.path.join(config.MODEL_DIR, 'X_train_sample.npy')
if os.path.exists(train_sample_path):
X_train_sample = np.load(train_sample_path)
feature_names = features.get_feature_names()
# Check for feature mismatch
if X_train_sample.shape[1] != len(feature_names):
st.warning(
f"Feature count mismatch (Model: {X_train_sample.shape[1]}, Code: {len(feature_names)}). Falling back to generic names.")
feature_names = [f"Feature_{i}" for i in range(X_train_sample.shape[1])]
# 2. Initialize Explainer
explainer = lime.lime_tabular.LimeTabularExplainer(
training_data=X_train_sample,
feature_names=feature_names,
class_names=classes,
mode='classification',
verbose=False
)
# 3. Explain this specific instance for ALL classes
# We pass labels=range(len(classes)) to calculate explanations for every class index
exp = explainer.explain_instance(
data_row=feat_scaled[0],
predict_fn=model.predict_proba,
num_features=10,
labels=range(len(classes))
)
# 4. Plot using Tabs
# Create a tab for each class so the user can switch between them
tabs = st.tabs(list(classes))
for i, class_name in enumerate(classes):
with tabs[i]:
st.write(f"**Evidence For/Against: {class_name}**")
# LIME uses the index (i) to retrieve the specific explanation
fig = exp.as_pyplot_figure(label=i)
st.pyplot(fig)
else:
st.warning("LIME initialization data (X_train_sample.npy) not found. Re-run training.")
except Exception as e:
st.error(f"Could not generate explanation: {type(e).__name__}: {e}")
# --- Pipeline Visualization ---
st.divider()
with st.expander("See Internal Logic (Computer Vision Pipeline Steps)", expanded=True):
st.info("Visualizing the exact steps performed by `src.features.py`")
img_resized, img_gray, img_eq, img_blur = features.preprocess_image(image)
mask_raw, mask_clean, mask_connected = features.segment_lesion(img_blur)
mask_final, _, _, _ = features.isolate_largest_component(mask_connected)
_, texture_vis = features.compute_texture_canny(img_gray, mask=mask_final)
img_lesion_only = cv2.bitwise_and(img_resized, img_resized, mask=mask_final)
# Row 1: Preprocessing
st.markdown("### Phase 1: Preprocessing")
c1, c2, c3, c4 = st.columns(4)
c1.image(img_resized, channels="BGR", caption="1. Resize")
c2.image(img_gray, caption="2. Grayscale")
c3.image(img_eq, caption="3. CLAHE (Smart Contrast)")
c4.image(img_blur, caption="4. Blur (Reduce Noise)")
st.divider()
# Row 2: Segmentation
st.markdown("### Phase 2: Segmentation")
c5, c6 = st.columns(2)
c5.image(mask_raw, caption="5. Otsu Threshold")
c6.image(mask_clean, caption="6. Morph Opening")
st.divider()
# Row 3: Connection & Selection
c7, c8 = st.columns(2)
c7.image(mask_connected, caption="7. Morph Dilation")
c8.image(mask_final, caption="8. Final Mask")
st.divider()
# Row 4: Analysis
st.markdown("### Phase 3: Analysis")
c9, c10 = st.columns(2)
c9.image(img_lesion_only, channels="BGR", caption="9. Masked Source")
c10.image(texture_vis, caption="10. Canny Edges (Masked)")
# Histogram
st.write("**11. Lesion Color Histogram**")
fig, ax = plt.subplots(figsize=(10, 3))
colors = ('b', 'g', 'r')
for i, color in enumerate(colors):
hist = cv2.calcHist([img_resized], [i], mask_final, [256], [0, 256])
ax.plot(hist, color=color)
ax.set_xlim([0, 256])
ax.set_title("Color Frequency")
st.pyplot(fig) |