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)