Spaces:
Sleeping
Sleeping
File size: 4,180 Bytes
93138dd cd67ded 93138dd cd67ded 93138dd cd67ded 93138dd cd67ded 93138dd 5a0afb5 93138dd cd67ded 93138dd cd67ded 93138dd cd67ded 93138dd cd67ded 93138dd cd67ded 93138dd cd67ded adbe4c4 93138dd | 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 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 | # ================================
# 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) |