KamdiO/Diabetic_Retinopathy_Class
0
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()