CoolFace
Apppublic

KamdiO/Diabetic_Retinopathy_Class

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
0likes
app.py128 linesDownload Raw Back to root
1import numpy as np
2import pickle
3import streamlit as st
4from PIL import Image
5import os
6from fastai.vision.all import *
7
8# set page config
9st.set_page_config(
10    page_title="Diabetic Retinopathy Model Comparison",
11    layout="wide"
12)
13
14# define function to load models
15@st.cache_resource
16def load_models():
17    models = {}
18    model_files = {
19        "Model 1": "C:\Users\Kamdi\Desktop\diabeticretin_proj\models\dr_model_resnet1.pkl",
20        "Model 2":"C:\Users\Kamdi\Desktop\diabeticretin_proj\models\dr_model_resnet1.pkl",
21        "Model 3":"C:\Users\Kamdi\Desktop\diabeticretin_proj\models\dr_model_resnet1.pkl",
22        "Model 4":"C:\Users\Kamdi\Desktop\diabeticretin_proj\models\dr_model_resnet1.pkl",
23        "Model 5":"C:\Users\Kamdi\Desktop\diabeticretin_proj\models\dr_model_resnet1.pkl",
24    }
25
26    for model_name, model_path in model_files.items():
27        try:
28            models[model_name] = load_learner(model_path)
29        except Exception as e:
30            st.error(f"Error loading {model_name}: {str(e)}")
31    return models
32
33# define function to load test images
34@st.cache_data
35def load_test_images():
36    images = {}
37    image_files = {
38        "Test Image 1": "C:\Users\Kamdi\Desktop\diabeticretin_proj\test_images\image1.jpeg",
39        "Test Image 2": "C:\Users\Kamdi\Desktop\diabeticretin_proj\test_images\image2.jpeg",
40        "Test Image 3": "C:\Users\Kamdi\Desktop\diabeticretin_proj\test_images\image3.jpeg",
41        "Test Image 4": "C:\Users\Kamdi\Desktop\diabeticretin_proj\test_images\image4.jpeg",
42        "Test Image 5": "C:\Users\Kamdi\Desktop\diabeticretin_proj\test_images\image5.jpeg",
43    }
44
45    for image_name, image_path in image_files.items():
46        try:
47            images[image_name] = Image.open(image_path)
48        except Exception as e:
49            st.error(f"Error loading {image_name}: {str(e)}")
50    return images
51
52# define a function to preprocess images
53def preprocess_image(image):
54    # Convert PIL Image to fastai format
55    img = PILImage.create(np.array(image))
56    return img
57
58def get_prediction(model, image):
59    processed_image = preprocess_image(image)
60    # Get prediction and probability
61    pred, pred_idx, probs = model.predict(processed_image)
62    return pred_idx
63
64def main():
65    st.title("Diabetic Retinopathy Model Comparison")
66    
67    # Load models and images
68    try:
69        models = load_models()
70        images = load_test_images()
71    except Exception as e:
72        st.error(f"Error loading models or images: {str(e)}")
73        return
74
75    # Create two columns
76    col1, col2 = st.columns(2)
77
78    with col1:
79        st.subheader("Select Test Image")
80        # Image selection
81        selected_image_name = st.selectbox(
82            "Choose a test image:",
83            list(images.keys())
84        )
85        
86        # Display selected image
87        if selected_image_name:
88            st.image(images[selected_image_name], caption=selected_image_name, use_column_width=True)
89
90    with col2:
91        st.subheader("Model Selection and Results")
92        # Model selection
93        selected_models = st.multiselect(
94            "Select models to compare:",
95            list(models.keys())
96        )
97
98        if st.button("Run Analysis"):
99            if selected_models and selected_image_name:
100                # Create results table
101                results = []
102                for model_name in selected_models:
103                    try:
104                        prediction = get_prediction(models[model_name], images[selected_image_name])
105                        results.append({
106                            "Model": model_name,
107                            "Severity Score (0-4)": int(prediction)
108                        })
109                    except Exception as e:
110                        st.error(f"Error with {model_name}: {str(e)}")
111                
112                # Display results
113                st.table(results)
114                
115                # Display interpretation guide
116                st.markdown("""
117                **Severity Score Interpretation:**
118                - 0: No DR
119                - 1: Mild DR
120                - 2: Moderate DR
121                - 3: Severe DR
122                - 4: Proliferative DR
123                """)
124            else:
125                st.warning("Please select at least one model and an image.")
126
127if __name__ == "__main__":
128    main()