walidchaib's picture
Update app.py
5a0afb5 verified
Raw
History Blame Contribute Delete
4.18 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
import cv2
from PIL import Image
# ================================
# 3. CHARGEMENT DU MODÈLE (détection)
# ================================
# MODEL_PATH = "detection_model.h5" # Note: le modèle sauvegardé est .h5 (HDF5) best_model.keras
MODEL_PATH = "best_model.keras" # Note: le modèle sauvegardé est .h5 (HDF5)
model = tf.keras.models.load_model(MODEL_PATH, compile=False)
# Recompiler pour éviter les warnings (non nécessaire pour l'inférence)
model.compile(optimizer='adam', loss={'class_output':'sparse_categorical_crossentropy', 'bbox_output':'mse'})
IMG_SIZE = 224
# ================================
# 4. FONCTION DE PRÉDICTION
# ================================
def preprocess_image(img):
"""Redimensionne et normalise l'image pour le modèle."""
img = img.resize((IMG_SIZE, IMG_SIZE))
img_array = np.array(img, dtype=np.float32) / 255.0
img_array = np.expand_dims(img_array, axis=0)
return img_array
def predict_detection(img):
"""
Prend une image PIL, renvoie l'image avec boîte englobante dessinée et le texte de classification.
"""
# Prétraitement
processed = preprocess_image(img)
# Prédiction
pred_cls, pred_bbox = model.predict(processed, verbose=0)
cls = np.argmax(pred_cls[0]) # 0 = crop, 1 = weed
bbox = pred_bbox[0] # [xmin, ymin, xmax, ymax] en coordonnées normalisées (0-1)
# Convertir en coordonnées pixel sur l'image originale (taille originale)
# L'utilisateur a téléchargé une image qui peut ne pas être carrée, nous devons adapter.
# Nous travaillons sur l'image redimensionnée pour l'affichage, mais la boîte doit être
# redimensionnée proportionnellement. On va redessiner sur une copie de l'image redimensionnée
# pour l'affichage.
img_disp = img.resize((IMG_SIZE, IMG_SIZE)) # même taille que celle utilisée par le modèle
w, h = img_disp.size
# Coordonnées absolues
xmin = int(bbox[0] * w)
ymin = int(bbox[1] * h)
xmax = int(bbox[2] * w)
ymax = int(bbox[3] * h)
# Convertir PIL en array OpenCV (BGR) pour dessiner
img_cv = cv2.cvtColor(np.array(img_disp), cv2.COLOR_RGB2BGR)
# Dessiner le rectangle
cv2.rectangle(img_cv, (xmin, ymin), (xmax, ymax), (0, 255, 0), 2)
# Ajouter le texte
label = "Crop" if cls == 0 else "Weed"
cv2.putText(img_cv, label, (xmin, ymin-5), cv2.FONT_HERSHEY_SIMPLEX, 0.6, (0, 255, 0), 2)
# Reconvertir en PIL pour affichage
img_result = Image.fromarray(cv2.cvtColor(img_cv, cv2.COLOR_BGR2RGB))
return img_result
# ================================
# 5. INTERFACE GRADIO
# ================================
iface = gr.Interface(
fn=predict_detection,
inputs=gr.Image(type="pil", label="Chargez une image de champ"),
outputs=gr.Image(type="pil", label="Résultat avec détection"),
title="Détection Culture / Mauvaise Herbe",
description="Ce modèle (MobileNetV2) détecte les cultures et les mauvaises herbes et renvoie la boîte englobante de l'objet principal ainsi que sa classe. "
)
if __name__ == "__main__":
iface.launch(share=True)