walidchaib's picture
Update app.py
47c1670 verified
Raw
History Blame Contribute Delete
2.96 kB
# ================================
# 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)