EngReem85 commited on
Commit
f1946ce
·
verified ·
1 Parent(s): a247f01

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +82 -140
app.py CHANGED
@@ -1,146 +1,88 @@
1
- import os
2
- import cv2
3
  import gradio as gr
4
  import numpy as np
5
  import tensorflow as tf
6
- from scipy import ndimage
7
- from tensorflow.keras import layers, models
8
-
9
- MODEL_PATH = "mnist_model.keras"
10
- import spaces # استيراد مكتبة ZeroGPU
11
-
12
-
13
- # إضافة المزخرف فوق دالة التنبؤ
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
14
  @spaces.GPU
15
-
16
- def get_or_build_model():
17
- # 1. إذا كان النموذج موجوداً، قم بتحميله فوراً
18
- if os.path.exists(MODEL_PATH):
19
- print("✅ تم تحميل النموذج بنجاح من الملف")
20
- return tf.keras.models.load_model(MODEL_PATH)
21
-
22
- # 2. إذا لم يكن موجوداً، يتم بناؤه وتدريبه سريعاً لمرة واحدة
23
- print("⚠️ لم يتم العثور على ملف النموذج. جاري بناء وتدريب نموذج جديد على CPU...")
24
- model = models.Sequential(
25
- [
26
- layers.Input(shape=(28, 28, 1)),
27
- layers.Conv2D(32, (3, 3), activation="relu"),
28
- layers.BatchNormalization(),
29
- layers.MaxPooling2D((2, 2)),
30
- layers.Conv2D(64, (3, 3), activation="relu"),
31
- layers.BatchNormalization(),
32
- layers.MaxPooling2D((2, 2)),
33
- layers.Flatten(),
34
- layers.Dense(64, activation="relu"),
35
- layers.Dropout(0.3),
36
- layers.Dense(10, activation="softmax"),
37
- ]
38
- )
39
-
40
- model.compile(
41
- optimizer="adam",
42
- loss="sparse_categorical_crossentropy",
43
- metrics=["accuracy"],
44
- )
45
-
46
- # تحميل البيانات وتقسيمها يدوياً لتفادي أخطاء Keras 3
47
- (x_train_full, y_train_full), _ = tf.keras.datasets.mnist.load_data()
48
- x_train_full = x_train_full.reshape(-1, 28, 28, 1).astype("float32") / 255.0
49
-
50
- split = int(len(x_train_full) * 0.9)
51
- x_tr, x_va = x_train_full[:split], x_train_full[split:]
52
- y_tr, y_va = y_train_full[:split], y_train_full[split:]
53
-
54
- # تدريب خفيف جداً لـ 3 جولات فقط (سريع ومناسب لـ CPU Space)
55
- model.fit(
56
- x_tr,
57
- y_tr,
58
- epochs=3,
59
- batch_size=128,
60
- validation_data=(x_va, y_va),
61
- verbose=1,
62
- )
63
-
64
- model.save(MODEL_PATH)
65
- print("✅ تم تدريب النموذج وحفظه بنجاح")
66
- return model
67
-
68
-
69
- # تحميل أو بناء النموذج
70
- model = get_or_build_model()
71
-
72
-
73
- # دالة تمركز كتلة الرسم
74
- def center_image(img):
75
- cy, cx = ndimage.center_of_mass(img)
76
- rows, cols = img.shape
77
- shiftx = np.round(cols / 2.0 - cx)
78
- shifty = np.round(rows / 2.0 - cy)
79
- M = np.float32([[1, 0, shiftx], [0, 1, shifty]])
80
- return cv2.warpAffine(img, M, (cols, rows))
81
-
82
-
83
- # دالة التنبؤ
84
- def predict(image):
85
- if image is None:
86
- return None, {"error": "لم يتم الرسم"}
87
-
88
- if isinstance(image, dict):
89
- image = image.get("composite", image.get("background"))
90
-
91
- gray = cv2.cvtColor(image, cv2.COLOR_RGB2GRAY)
92
- _, thresh = cv2.threshold(gray, 200, 255, cv2.THRESH_BINARY_INV)
93
-
94
- coords = cv2.findNonZero(thresh)
95
- if coords is not None:
96
- x, y, w, h = cv2.boundingRect(coords)
97
- cropped = thresh[y : y + h, x : x + w]
98
-
99
- if w > h:
100
- new_w, new_h = 20, int(h * (20 / w))
101
- else:
102
- new_h, new_w = 20, int(w * (20 / h))
103
-
104
- resized = cv2.resize(
105
- cropped, (new_w, new_h), interpolation=cv2.INTER_AREA
106
- )
107
-
108
- pad_v = (28 - new_h) // 2
109
- pad_h = (28 - new_w) // 2
110
- padded = np.pad(
111
- resized,
112
- ((pad_v, 28 - new_h - pad_v), (pad_h, 28 - new_w - pad_h)),
113
- "constant",
114
- )
115
- else:
116
- padded = cv2.resize(thresh, (28, 28))
117
-
118
- final_img = center_image(padded)
119
- input_data = (final_img.astype("float32") / 255.0).reshape(1, 28, 28, 1)
120
-
121
- pred = model.predict(input_data, verbose=0)[0]
122
- return final_img, {str(i): float(pred[i]) for i in range(10)}
123
-
124
-
125
- # واجهة Gradio
126
- with gr.Blocks(title="MNIST Predictor") as demo:
127
- gr.Markdown("# 🧠 التنبؤ بالأرقام المرسومة (MNIST)")
128
 
129
  with gr.Row():
130
- with gr.Column():
131
- input_pad = gr.Sketchpad(label="ارسم رقماً هنا")
132
- btn_predict = gr.Button("تنبؤ", variant="primary")
133
-
134
- with gr.Column():
135
- out_label = gr.Label(num_top_classes=3, label="النتيجة")
136
- out_img = gr.Image(label="الصورة المعالجة (28x28)", image_mode="L")
137
-
138
- btn_predict.click(
139
- fn=predict,
140
- inputs=[input_pad],
141
- outputs=[out_img, out_label],
142
- )
143
-
144
- if __name__ == "__main__":
145
- demo = gr.Interface(fn=predict_digit, inputs="sketchpad", outputs="label")
146
- demo.launch()
 
 
 
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 # <--- 1. استيراد مكتبة spaces
8
+
9
+ # ============================================
10
+ # 1. تحميل النموذج في النطاق العام (مرة واحدة فقط)
11
+ # ============================================
12
+ model_path = 'mnist_cnn_model.keras'
13
+
14
+ if os.path.exists(model_path):
15
+ model = tf.keras.models.load_model(model_path)
16
+ print("✅ تم تحميل النموذج المحفوظ.")
17
+ else:
18
+ print("⚠️ لم يتم العثور على النموذج، جارٍ البناء والتدريب...")
19
+ model = models.Sequential([
20
+ layers.Conv2D(32, (3, 3), activation='relu', input_shape=(28, 28, 1)),
21
+ layers.MaxPooling2D((2, 2)),
22
+ layers.Conv2D(64, (3, 3), activation='relu'),
23
+ layers.MaxPooling2D((2, 2)),
24
+ layers.Flatten(),
25
+ layers.Dense(64, activation='relu'),
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
34
+ model.fit(x_train, y_train, epochs=3, 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
43
+
44
+ # ============================================
45
+ # 3. دوال التنبؤ والمثال العشوائي
46
+ # ============================================
47
+ # --- الدالة التي تستخدم GPU يتم تزيينها بـ @spaces.GPU ---
48
  @spaces.GPU
49
+ def predict_sketch(image):
50
+ try:
51
+ gray = cv2.cvtColor(image, cv2.COLOR_RGB2GRAY)
52
+ resized = cv2.resize(gray, (28, 28))
53
+ inverted = 255 - resized
54
+ normalized = inverted.astype('float32') / 255.0
55
+ reshaped = normalized.reshape(1, 28, 28, 1)
56
+
57
+ pred = model.predict(reshaped, verbose=0)[0]
58
+ return {str(i): float(pred[i]) for i in range(10)}
59
+ except Exception as e:
60
+ return {"خطأ": str(e)}
61
+
62
+ # --- هذه الدالة لا تحتاج GPU، لذا لا تزينها ---
63
+ def random_example():
64
+ idx = np.random.randint(0, len(test_images))
65
+ img = test_images[idx].reshape(28, 28)
66
+ img_rgb = np.stack([img] * 3, axis=2)
67
+ return img_rgb, f"الرقم الحقيقي: {test_labels[idx]}"
68
+
69
+ # ============================================
70
+ # 4. بناء واجهة Gradio
71
+ # ============================================
72
+ with gr.Blocks(title="MNIST Recognizer") as demo:
73
+ gr.Markdown("# 🧠 التعرف على الأرقام المكتوبة بخط اليد")
74
+ gr.Markdown("ارسم رقماً (0-9) في المربع، أو اضغط على زر **مثال عشوائي**.")
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
75
 
76
  with gr.Row():
77
+ with gr.Column(scale=1):
78
+ sketch = gr.Sketchpad(label="✏️ ارسم هنا")
79
+ with gr.Row():
80
+ submit_btn = gr.Button("🔮 توقع", variant="primary")
81
+ random_btn = gr.Button("🎲 مثال عشوائي", variant="secondary")
82
+ info = gr.Textbox(label="📌 معلومات", interactive=False)
83
+
84
+ with gr.Column(scale=1):
85
+ output = gr.Label(num_top_classes=3, label="📊 الاحتمالات")
86
+
87
+ submit_btn.click(fn=predict_sketch, inputs=sketch, outputs=output)
88
+ random_btn.click(fn=random_example, inputs=[], outputs=[sketch, info])