EngReem85 commited on
Commit
59d6247
·
verified ·
1 Parent(s): 8b1d6f5

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +60 -58
app.py CHANGED
@@ -1,20 +1,20 @@
1
- import gradio as gr
 
2
  import numpy as np
3
  import tensorflow as tf
4
- import cv2
5
  from tensorflow.keras import datasets, layers, models
6
- import os
7
- import spaces
8
 
9
  # ============================================
10
- # 1. تحميل النموذج
11
  # ============================================
12
  model_path = 'mnist_cnn_model.keras'
 
13
  if os.path.exists(model_path):
14
  model = tf.keras.models.load_model(model_path)
15
- print("✅ تم تحميل النموذج المحفوظ.")
16
  else:
17
- print("⚠️ بناء نموذج جديد...")
18
  model = models.Sequential([
19
  layers.Conv2D(32, (3, 3), activation='relu', input_shape=(28, 28, 1)),
20
  layers.MaxPooling2D((2, 2)),
@@ -25,45 +25,51 @@ else:
25
  layers.Dropout(0.5),
26
  layers.Dense(10, activation='softmax')
27
  ])
28
- model.compile(optimizer='adam',
29
- loss='sparse_categorical_crossentropy',
30
- metrics=['accuracy'])
 
 
31
 
 
32
  (x_train, y_train), _ = datasets.mnist.load_data()
33
  x_train = x_train.reshape((60000, 28, 28, 1)).astype('float32') / 255.0
 
34
  model.fit(x_train, y_train, epochs=5, validation_split=0.1, verbose=1)
35
  model.save(model_path)
36
- print("✅ تم بناء النموذج وتدريبه وحفظه.")
37
 
38
  # ============================================
39
- # 2. تحميل بيانات الاختبار
40
  # ============================================
41
  (_, _), (test_images, test_labels) = datasets.mnist.load_data()
42
  test_images = test_images.reshape((10000, 28, 28, 1)).astype('float32') / 255.0
43
 
44
  # ============================================
45
- # 3. دالة معالجة الصورة
46
  # ============================================
47
  def process_image_for_mnist(image_input):
48
- """
49
- تعالج الصورة المدخلة من Sketchpad لتناسب تنسيق MNIST
50
- """
51
  if image_input is None:
52
- raise ValueError("لم يتم إدخال أي صورة!")
53
 
54
- # التعامل مع مدخلات Sketchpad (تكون Dictionary في الإصدارات الحديثة)
55
  if isinstance(image_input, dict):
56
  image = image_input.get('composite', None)
57
  if image is None:
58
- image = image_input.get('background', list(image_input.values())[0])
 
 
59
  else:
60
  image = image_input
61
 
62
- # تحويل الصورة إلى RGB إذا كانت تحتوي على قناة Alpha (RGBA)
63
- if image.shape[-1] == 4:
 
 
 
64
  image = cv2.cvtColor(image, cv2.COLOR_RGBA2RGB)
65
 
66
- # تحويل إلى تدرج رمادي Gray Scale
67
  if len(image.shape) == 3:
68
  gray = cv2.cvtColor(image, cv2.COLOR_RGB2GRAY)
69
  else:
@@ -72,79 +78,75 @@ def process_image_for_mnist(image_input):
72
  # تغيير الحجم إلى 28x28
73
  resized = cv2.resize(gray, (28, 28), interpolation=cv2.INTER_AREA)
74
 
75
- # عتبة أوتوماتيكية لجعل الصورة ثنائية (أسود وأبيض)
76
  _, binary = cv2.threshold(resized, 0, 255, cv2.THRESH_BINARY + cv2.THRESH_OTSU)
77
 
78
- # جعل الخلفية سوداء (0) والرسم أبيض (255) كما في MNIST
79
- white_pixels = np.sum(binary == 255)
80
- black_pixels = np.sum(binary == 0)
81
- if white_pixels > black_pixels:
82
  binary = 255 - binary
83
 
84
- # تطبيع القيم بين 0 و 1
85
  normalized = binary.astype('float32') / 255.0
86
  reshaped = normalized.reshape(1, 28, 28, 1)
87
 
88
  return normalized, reshaped
89
 
90
  # ============================================
91
- # 4. دالة التنبؤ
92
  # ============================================
93
- @spaces.GPU
94
  def predict_sketch(image):
95
  try:
96
  if image is None:
97
- return {}, None, "���️ يرجى الرسم أولاً!"
98
 
99
- # معالجة الصورة
100
  normalized, reshaped = process_image_for_mnist(image)
101
 
102
- # التنبؤ بواسطة النموذج
103
- pred = model.predict(reshaped, verbose=0)[0]
104
- predicted_class = int(np.argmax(pred))
105
- confidence = float(np.max(pred))
106
-
107
- # تحضير نص التصحيح Debug Info
108
- top3_idx = np.argsort(pred)[-3:][::-1]
109
- top3_str = "\n".join([f" {i+1}. الرقم {idx}: {pred[idx]:.2%}" for i, idx in enumerate(top3_idx)])
110
-
111
- debug_info = f"""📊 معلومات التوقع:
112
- - الفئة المتوقعة: {predicted_class}
113
- - نسبة الثقة: {confidence:.2%}
114
 
115
- 🔝 أعلى 3 احتمالات:
116
- {top3_str}
 
 
117
 
118
- 🔢 جميع الاحتمالات:
119
- {', '.join([f'{i}: {pred[i]:.1%}' for i in range(10)])}"""
120
 
121
- # تحويل صورة المعالجة لمصفوفة 3D لعرضها في Gradio بشكل مريح
122
  display_img = (normalized * 255).astype(np.uint8)
123
  display_img = cv2.resize(display_img, (140, 140), interpolation=cv2.INTER_NEAREST)
124
 
125
- # إرجاع المخرجات الثلاثة مباشرة بالتفصيل
126
- predictions_dict = {str(i): float(pred[i]) for i in range(10)}
127
- return predictions_dict, display_img, debug_info
 
 
 
 
 
 
 
 
 
128
 
129
  except Exception as e:
130
- return {}, None, f"❌ حدث خطأ أثناء التوقع: {str(e)}"
131
 
132
  # ============================================
133
- # 5. دالة المثال العشوائي
134
  # ============================================
135
  def random_example():
136
  idx = np.random.randint(0, len(test_images))
137
  img = test_images[idx].reshape(28, 28)
138
  img_large = cv2.resize(img, (280, 280), interpolation=cv2.INTER_NEAREST)
139
  img_rgb = np.stack([img_large] * 3, axis=2)
140
- return img_rgb, f"🎲 تم اختيار رقم عشوائي (الرقم الحقيقي: {test_labels[idx]})"
141
 
142
  # ============================================
143
  # 6. بناء واجهة Gradio
144
  # ============================================
145
  with gr.Blocks(title="MNIST Digit Recognizer") as demo:
146
  gr.Markdown("# 🧠 التعرف على الأرقام المكتوبة بخط اليد (MNIST)")
147
- gr.Markdown("ارسم رقماً (من 0 إلى 9) في المربع الأيسر، ثم اضغط على زر **توقع**.")
148
 
149
  with gr.Row():
150
  with gr.Column(scale=1):
@@ -152,13 +154,13 @@ with gr.Blocks(title="MNIST Digit Recognizer") as demo:
152
  with gr.Row():
153
  submit_btn = gr.Button("🔮 توقع", variant="primary")
154
  random_btn = gr.Button("🎲 مثال عشوائي", variant="secondary")
155
- info = gr.Textbox(label="📌 معلومات التصحيح (Debug Info)", interactive=False, lines=10)
156
 
157
  with gr.Column(scale=1):
158
  output = gr.Label(num_top_classes=3, label="📊 احتمالات التوقع")
159
  processed_img = gr.Image(label="🖼️ الصورة بعد المعالجة (ما يراه النموذج)", image_mode="L")
160
 
161
- # ربط الأحداث والأزرار
162
  submit_btn.click(
163
  fn=predict_sketch,
164
  inputs=sketch,
 
1
+ import os
2
+ import cv2
3
  import numpy as np
4
  import tensorflow as tf
5
+ import gradio as gr
6
  from tensorflow.keras import datasets, layers, models
 
 
7
 
8
  # ============================================
9
+ # 1. تحميل أو بناء النموذج (على CPU)
10
  # ============================================
11
  model_path = 'mnist_cnn_model.keras'
12
+
13
  if os.path.exists(model_path):
14
  model = tf.keras.models.load_model(model_path)
15
+ print("✅ تم تحميل النموذج المحفوظ بنجاح.")
16
  else:
17
+ print("⚠️ لم يتم العثور على النموذج.. جاري البناء والتدريب...")
18
  model = models.Sequential([
19
  layers.Conv2D(32, (3, 3), activation='relu', input_shape=(28, 28, 1)),
20
  layers.MaxPooling2D((2, 2)),
 
25
  layers.Dropout(0.5),
26
  layers.Dense(10, activation='softmax')
27
  ])
28
+ model.compile(
29
+ optimizer='adam',
30
+ loss='sparse_categorical_crossentropy',
31
+ metrics=['accuracy']
32
+ )
33
 
34
+ # تحميل بيانات التدريب
35
  (x_train, y_train), _ = datasets.mnist.load_data()
36
  x_train = x_train.reshape((60000, 28, 28, 1)).astype('float32') / 255.0
37
+
38
  model.fit(x_train, y_train, epochs=5, validation_split=0.1, verbose=1)
39
  model.save(model_path)
40
+ print("✅ تم حفظ النموذج المكتمل.")
41
 
42
  # ============================================
43
+ # 2. تحميل بيانات الاختبار للأمثلة العشوائية
44
  # ============================================
45
  (_, _), (test_images, test_labels) = datasets.mnist.load_data()
46
  test_images = test_images.reshape((10000, 28, 28, 1)).astype('float32') / 255.0
47
 
48
  # ============================================
49
+ # 3. دالة معالجة واستخراج الصورة
50
  # ============================================
51
  def process_image_for_mnist(image_input):
 
 
 
52
  if image_input is None:
53
+ return None, None
54
 
55
+ # التعامل مع مدخلات Sketchpad المتنوعة في Gradio
56
  if isinstance(image_input, dict):
57
  image = image_input.get('composite', None)
58
  if image is None:
59
+ image = image_input.get('background', None)
60
+ if image is None and len(image_input) > 0:
61
+ image = list(image_input.values())[0]
62
  else:
63
  image = image_input
64
 
65
+ if image is None or not isinstance(image, np.ndarray):
66
+ return None, None
67
+
68
+ # تحويل RGBA إلى RGB
69
+ if len(image.shape) == 3 and image.shape[-1] == 4:
70
  image = cv2.cvtColor(image, cv2.COLOR_RGBA2RGB)
71
 
72
+ # تحويل إلى رمادي Gray
73
  if len(image.shape) == 3:
74
  gray = cv2.cvtColor(image, cv2.COLOR_RGB2GRAY)
75
  else:
 
78
  # تغيير الحجم إلى 28x28
79
  resized = cv2.resize(gray, (28, 28), interpolation=cv2.INTER_AREA)
80
 
81
+ # تطبيق العتبة الثنائية (Otsu Thresholding)
82
  _, binary = cv2.threshold(resized, 0, 255, cv2.THRESH_BINARY + cv2.THRESH_OTSU)
83
 
84
+ # عكس الألوان لتصبح الخلفية سوداء (0) والرسم أبيض (255) كبيئة MNIST
85
+ if np.sum(binary == 255) > np.sum(binary == 0):
 
 
86
  binary = 255 - binary
87
 
 
88
  normalized = binary.astype('float32') / 255.0
89
  reshaped = normalized.reshape(1, 28, 28, 1)
90
 
91
  return normalized, reshaped
92
 
93
  # ============================================
94
+ # 4. دالة التنبؤ المباشرة والآمنة
95
  # ============================================
 
96
  def predict_sketch(image):
97
  try:
98
  if image is None:
99
+ return {}, None, "⚠️ يرجى الرسم في المربع أولاً!"
100
 
 
101
  normalized, reshaped = process_image_for_mnist(image)
102
 
103
+ if normalized is None:
104
+ return {}, None, "⚠️ لم يتم التعرف على الصورة، يرجى إعادة الرسم."
 
 
 
 
 
 
 
 
 
 
105
 
106
+ # الاستدعاء المباشر للموديل يمنع تعارض الخيوط (Threads) في TensorFlow
107
+ preds = model(reshaped, training=False).numpy()[0]
108
+ predicted_class = int(np.argmax(preds))
109
+ confidence = float(np.max(preds))
110
 
111
+ # قاموس الاحتمالات الموجه لـ Gradio Label
112
+ probabilities = {str(i): float(preds[i]) for i in range(10)}
113
 
114
+ # تحضير صورة العرض المصغرة
115
  display_img = (normalized * 255).astype(np.uint8)
116
  display_img = cv2.resize(display_img, (140, 140), interpolation=cv2.INTER_NEAREST)
117
 
118
+ # نصوص معلومات التصحيح
119
+ top3_idx = np.argsort(preds)[-3:][::-1]
120
+ top3_str = "\n".join([f" {i+1}. الرقم {idx}: {preds[idx]:.2%}" for i, idx in enumerate(top3_idx)])
121
+
122
+ debug_info = f"""📊 نتائج التحليل:
123
+ - الرقم المتوقع: {predicted_class}
124
+ - نسبة الثقة: {confidence:.2%}
125
+
126
+ 🔝 أعلى 3 احتمالات:
127
+ {top3_str}"""
128
+
129
+ return probabilities, display_img, debug_info
130
 
131
  except Exception as e:
132
+ return {}, None, f"❌ حدث خطأ أثناء التوقّع: {str(e)}"
133
 
134
  # ============================================
135
+ # 5. دالة اختياري: اختيار مثال عشوائي
136
  # ============================================
137
  def random_example():
138
  idx = np.random.randint(0, len(test_images))
139
  img = test_images[idx].reshape(28, 28)
140
  img_large = cv2.resize(img, (280, 280), interpolation=cv2.INTER_NEAREST)
141
  img_rgb = np.stack([img_large] * 3, axis=2)
142
+ return img_rgb, f"🎲 تم جلب رقم عشوائي (الرقم الحقيقي: {test_labels[idx]})"
143
 
144
  # ============================================
145
  # 6. بناء واجهة Gradio
146
  # ============================================
147
  with gr.Blocks(title="MNIST Digit Recognizer") as demo:
148
  gr.Markdown("# 🧠 التعرف على الأرقام المكتوبة بخط اليد (MNIST)")
149
+ gr.Markdown("ارسم رقماً من **0 إلى 9** في المربع الأيسر ثم اضغط على زر **توقع**.")
150
 
151
  with gr.Row():
152
  with gr.Column(scale=1):
 
154
  with gr.Row():
155
  submit_btn = gr.Button("🔮 توقع", variant="primary")
156
  random_btn = gr.Button("🎲 مثال عشوائي", variant="secondary")
157
+ info = gr.Textbox(label="📌 معلومات التصحيح (Debug Info)", interactive=False, lines=8)
158
 
159
  with gr.Column(scale=1):
160
  output = gr.Label(num_top_classes=3, label="📊 احتمالات التوقع")
161
  processed_img = gr.Image(label="🖼️ الصورة بعد المعالجة (ما يراه النموذج)", image_mode="L")
162
 
163
+ # ربط الأحداث
164
  submit_btn.click(
165
  fn=predict_sketch,
166
  inputs=sketch,