File size: 2,958 Bytes
47c1670
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
97b8ad5
 
 
 
 
47c1670
 
 
 
 
f497282
47c1670
 
97b8ad5
47c1670
 
 
 
 
 
 
 
 
 
 
 
 
 
97b8ad5
47c1670
 
 
 
 
 
 
97b8ad5
dfecc6d
97b8ad5
dfecc6d
47c1670
 
 
 
 
 
 
 
97b8ad5
47c1670
 
 
 
97b8ad5
47c1670
 
 
 
97b8ad5
 
dfecc6d
47c1670
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
# ================================
# 0. PATCH pour huggingface_hub (contourne l'absence de HfFolder)
# ================================
import huggingface_hub
if not hasattr(huggingface_hub, 'HfFolder'):
    class HfFolder:
        _token = None
        @staticmethod
        def get_token():
            return HfFolder._token
        @staticmethod
        def save_token(token):
            HfFolder._token = token
    huggingface_hub.HfFolder = HfFolder

# ================================
# 1. PATCH pour contourner le bug de Gradio 4.44.0
#    (TypeError: argument of type 'bool' is not iterable)
# ================================
import gradio_client.utils

original_get_type = gradio_client.utils.get_type

def patched_get_type(schema):
    if isinstance(schema, bool):
        return "boolean"
    return original_get_type(schema)

gradio_client.utils.get_type = patched_get_type

# ================================
# 2. IMPORTS STANDARDS
# ================================
import gradio as gr
import tensorflow as tf
import numpy as np
from PIL import Image

# ================================
# 3. CHARGEMENT DU MODÈLE (9 classes)
# ================================
MODEL_PATH = "lemon_model.h5"
model = tf.keras.models.load_model(MODEL_PATH, compile=False)

# Recompilation simple pour éviter les warnings
model.compile(optimizer='adam', loss='categorical_crossentropy', metrics=['accuracy'])

# ================================
# 4. NOMS DES CLASSES (ordre de l'entraînement)
# ================================
class_names = [
    "Anthracnose",
    "Bacterial Blight",
    "Citrus Canker",
    "Curl Virus",
    "Deficiency Leaf",
    "Dry Leaf",
    "Healthy Leaf",
    "Sooty Mould",
    "Spider Mites"
]

# ================================
# 5. FONCTION DE PRÉDICTION
# ================================
def preprocess_image(img):
    """Redimensionne et normalise l'image (division par 255)."""
    img = img.resize((224, 224))
    img_array = np.array(img, dtype=np.float32) / 255.0
    img_array = np.expand_dims(img_array, axis=0)
    return img_array

def predict(img):
    """
    img : PIL Image
    Retourne un dictionnaire {classe: probabilité} pour le composant gr.Label
    """
    processed = preprocess_image(img)
    preds = model.predict(processed, verbose=0)[0]
    results = {class_names[i]: float(preds[i]) for i in range(len(class_names))}
    return results

# ================================
# 6. INTERFACE GRADIO
# ================================
iface = gr.Interface(
    fn=predict,
    inputs=gr.Image(type="pil", label="Téléchargez une photo de feuille de citron"),
    outputs=gr.Label(num_top_classes=3, label="Maladie prédite (top 3)"),
    title="🍋 Classification des maladies des feuilles de citron",
    description="Modèle MobileNet (entraîné sans transfert) pour reconnaître 9 types de pathologies ou états des feuilles de citron."
)

if __name__ == "__main__":
    iface.launch(share=True)