CoolFace
Apppublic

vsaez/object-detection-app

sourceHugging Faceupdated 1y agoView on Hugging Face
2likes
app.py519 linesDownload Raw Back to root
1import gradio as gr2import torch3from PIL import Image, ImageDraw, ImageFont4from transformers import DetrImageProcessor, DetrForObjectDetection5from pathlib import Path6import transformers7import warnings8import traceback9import datetime10 11warnings.filterwarnings("ignore", message=".*copying from a non-meta parameter.*")12 13# Global variables to cache models14current_model = None15current_processor = None16current_model_name = None17 18# Global debug state19debug_info = {"last_error": "", "step": "", "language": "", "timestamp": ""}20 21# Available models with better selection22available_models = {23    "DETR ResNet-50": "facebook/detr-resnet-50",24    "DETR ResNet-101": "facebook/detr-resnet-101",25    "DETR DC5": "facebook/detr-resnet-50-dc5",26    "DETR ResNet-50 Face Only": "esraakh/detr_fine_tune_face_detection_final"27}28 29 30def load_model(model_key):31    """Load model and processor based on selected model key"""32    global current_model, current_processor, current_model_name, debug_info33 34    model_name = available_models[model_key]35 36    # Only load if it's a different model37    if current_model_name != model_name:38        debug_info["step"] = f"Loading model: {model_name}"39        print(f"Loading model: {model_name}")40        current_processor = DetrImageProcessor.from_pretrained(model_name)41        current_model = DetrForObjectDetection.from_pretrained(model_name)42        current_model_name = model_name43        print(f"Model loaded: {model_name}")44        print(f"Available labels: {list(current_model.config.id2label.values())}")45        debug_info["step"] = f"Model loaded successfully: {model_name}"46 47    return current_model, current_processor48 49 50# Load font51font_path = Path("assets/fonts/arial.ttf")52if not font_path.exists():53    print(f"Font file {font_path} not found. Using default font.")54    font = ImageFont.load_default()55else:56    font = ImageFont.truetype(str(font_path), size=100)57 58# Set up translations for the app59translations = {60    "English": {61        "title": "## Enhanced Object Detection App\nUpload an image to detect objects using various DETR models.",62        "input_label": "Input Image",63        "output_label": "Detected Objects",64        "dropdown_label": "Label Language",65        "dropdown_detection_model_label": "Detection Model",66        "threshold_label": "Detection Threshold",67        "button": "Detect Objects",68        "info_label": "Detection Info",69        "error_label": "Error Messages",70        "debug_label": "Debug Status",71        "debug_button": "Show Debug Status",72        "model_fast": "General Objects (fast)",73        "model_precision": "General Objects (high precision)",74        "model_small": "Small Objects/Details (slow)",75        "model_faces": "Face Detection (people only)"76    },77    "Spanish": {78        "title": "## Aplicación Mejorada de Detección de Objetos\nSube una imagen para detectar objetos usando varios modelos DETR.",79        "input_label": "Imagen de entrada",80        "output_label": "Objetos detectados",81        "dropdown_label": "Idioma de las etiquetas",82        "dropdown_detection_model_label": "Modelo de detección",83        "threshold_label": "Umbral de detección",84        "button": "Detectar objetos",85        "info_label": "Información de detección",86        "error_label": "Mensajes de error",87        "debug_label": "Estado de depuración",88        "debug_button": "Mostrar estado de depuración",89        "model_fast": "Objetos generales (rápido)",90        "model_precision": "Objetos generales (precisión alta)",91        "model_small": "Objetos pequeños/detalles (lento)",92        "model_faces": "Detección de caras (solo personas)"93    },94    "French": {95        "title": "## Application Améliorée de Détection d'Objets\nTéléchargez une image pour détecter des objets avec divers modèles DETR.",96        "input_label": "Image d'entrée",97        "output_label": "Objets détectés",98        "dropdown_label": "Langue des étiquettes",99        "dropdown_detection_model_label": "Modèle de détection",100        "threshold_label": "Seuil de détection",101        "button": "Détecter les objets",102        "info_label": "Information de détection",103        "error_label": "Messages d'erreur",104        "debug_label": "État de débogage",105        "debug_button": "Afficher l'état de débogage",106        "model_fast": "Objets généraux (rapide)",107        "model_precision": "Objets généraux (haute précision)",108        "model_small": "Petits objets/détails (lent)",109        "model_faces": "Détection de visages (personnes uniquement)"110    }111}112 113 114def t(language, key):115    return translations.get(language, translations["English"]).get(key, key)116 117 118def get_translated_model_choices(language):119    """Get model choices translated to the selected language"""120    global debug_info121    debug_info["step"] = f"Translating model choices for {language}"122 123    model_mapping = {124        "DETR ResNet-50": "model_fast",125        "DETR ResNet-101": "model_precision",126        "DETR DC5": "model_small",127        "DETR ResNet-50 Face Only": "model_faces"128    }129 130    translated_choices = []131    for model_key in available_models.keys():132        if model_key in model_mapping:133            translation_key = model_mapping[model_key]134            translated_name = t(language, translation_key)135        else:136            translated_name = model_key137        translated_choices.append(translated_name)138 139    debug_info["step"] = f"Model choices translated: {translated_choices}"140    return translated_choices141 142 143def get_model_key_from_translation(translated_name, language):144    """Get the original model key from translated name"""145    model_mapping = {146        "DETR ResNet-50": "model_fast",147        "DETR ResNet-101": "model_precision",148        "DETR DC5": "model_small",149        "DETR ResNet-50 Face Only": "model_faces"150    }151 152    # Reverse lookup153    for model_key, translation_key in model_mapping.items():154        if t(language, translation_key) == translated_name:155            return model_key156 157    # If not found, try direct match158    if translated_name in available_models:159        return translated_name160 161    # Default fallback162    return "DETR ResNet-50"163 164 165def get_helsinki_model(language_label):166    """Returns the Helsinki-NLP model name for translating from English to the selected language."""167    lang_map = {168        "Spanish": "es",169        "French": "fr",170        "English": "en"171    }172    target = lang_map.get(language_label)173    if not target or target == "en":174        return None175    return f"Helsinki-NLP/opus-mt-en-{target}"176 177 178# Translation cache179translation_cache = {}180 181 182def translate_label(language_label, label):183    """Translates the given label to the target language."""184    # Check cache first185    cache_key = f"{language_label}_{label}"186    if cache_key in translation_cache:187        return translation_cache[cache_key]188 189    model_name = get_helsinki_model(language_label)190    if not model_name:191        return label192 193    try:194        translator = transformers.pipeline("translation", model=model_name)195        result = translator(label, max_length=40)196        translated = result[0]['translation_text']197        # Cache the result198        translation_cache[cache_key] = translated199        return translated200    except Exception as e:201        print(f"Translation error (429 or other): {e}")202        return label  # Return original if translation fails203 204 205def detect_objects(image, language_selector, translated_model_selector, threshold):206    """Enhanced object detection with adjustable threshold and better info"""207    global debug_info208 209    try:210        debug_info["step"] = "Starting object detection"211        debug_info["timestamp"] = str(datetime.datetime.now())212 213        # Get the actual model key from the translated name214        model_selector = get_model_key_from_translation(translated_model_selector, language_selector)215        debug_info["step"] = f"Model key resolved: {model_selector}"216 217        print(f"Processing image. Language: {language_selector}, Model: {model_selector}, Threshold: {threshold}")218 219        # Load the selected model220        debug_info["step"] = "Loading model"221        model, processor = load_model(model_selector)222 223        # Process the image224        debug_info["step"] = "Processing image with model"225        inputs = processor(images=image, return_tensors="pt")226        outputs = model(**inputs)227 228        # Convert model output to usable detection results with custom threshold229        debug_info["step"] = "Post-processing results"230        target_sizes = torch.tensor([image.size[::-1]])231        results = processor.post_process_object_detection(232            outputs, threshold=threshold, target_sizes=target_sizes233        )[0]234 235        # Create a copy of the image for drawing236        debug_info["step"] = "Drawing bounding boxes"237        image_with_boxes = image.copy()238        draw = ImageDraw.Draw(image_with_boxes)239 240        # Detection info241        detection_info = f"Detected {len(results['scores'])} objects with threshold {threshold}\n"242        detection_info += f"Model: {translated_model_selector} ({model_selector})\n\n"243 244        # Colors for different confidence levels245        colors = {246            'high': 'red',  # > 0.8247            'medium': 'orange',  # 0.5-0.8248            'low': 'yellow'  # < 0.5249        }250 251        detected_objects = []252 253        for score, label, box in zip(results["scores"], results["labels"], results["boxes"]):254            confidence = score.item()255            box = [round(x, 2) for x in box.tolist()]256 257            # Choose color based on confidence258            if confidence > 0.8:259                color = colors['high']260            elif confidence > 0.5:261                color = colors['medium']262            else:263                color = colors['low']264 265            # Draw bounding box266            draw.rectangle(box, outline=color, width=3)267 268            # Prepare label text269            label_text = model.config.id2label[label.item()]270            translated_label = translate_label(language_selector, label_text)271            display_text = f"{translated_label}: {round(confidence, 3)}"272 273            # Store detection info274            detected_objects.append({275                'label': label_text,276                'translated': translated_label,277                'confidence': confidence,278                'box': box279            })280 281            # Calculate text position and size282            try:283                text_bbox = draw.textbbox((0, 0), display_text, font=font)284                text_width = text_bbox[2] - text_bbox[0]285                text_height = text_bbox[3] - text_bbox[1]286            except:287                # Fallback for older PIL versions288                text_width, text_height = draw.textsize(display_text, font=font)289 290            # Draw text background291            text_bg = [292                box[0], box[1] - text_height - 4,293                        box[0] + text_width + 4, box[1]294            ]295            draw.rectangle(text_bg, fill="black")296            draw.text((box[0] + 2, box[1] - text_height - 2), display_text, fill="white", font=font)297 298        # Create detailed detection info299        if detected_objects:300            detection_info += "Objects found:\n"301            for obj in sorted(detected_objects, key=lambda x: x['confidence'], reverse=True):302                detection_info += f"- {obj['translated']} ({obj['label']}): {obj['confidence']:.3f}\n"303        else:304            detection_info += "No objects detected. Try lowering the threshold."305 306        debug_info["step"] = "Detection completed successfully"307        debug_info["last_error"] = ""308 309        return image_with_boxes, detection_info, ""310 311    except Exception as e:312        error_message = f"Error in object detection:\n{str(e)}\n\nStack trace:\n{traceback.format_exc()}"313        debug_info["last_error"] = error_message314        debug_info["step"] = f"ERROR in detection: {str(e)}"315        print(error_message)316        return image if image else None, "Detection failed. See error panel below.", error_message317 318 319def update_interface(selected_language):320    global debug_info321 322    debug_info["language"] = selected_language323    debug_info["timestamp"] = str(datetime.datetime.now())324    debug_info["step"] = "Starting language interface update"325 326    try:327        translated_choices = get_translated_model_choices(selected_language)328        default_model = t(selected_language, "model_fast")329 330        updates = [331            gr.update(value=t(selected_language, "title")),332            # gr.update(label=t(selected_language, "dropdown_label")), # <-- ELIMINADA ESTA LÍNEA333            gr.update(334                choices=translated_choices,335                value=default_model,336                label=t(selected_language, "dropdown_detection_model_label")337            ),338            gr.update(label=t(selected_language, "threshold_label")),339            gr.update(label=t(selected_language, "input_label")),340            gr.update(value=t(selected_language, "button")),341            gr.update(label=t(selected_language, "output_label")),342            gr.update(label=t(selected_language, "info_label")),343            gr.update(label=t(selected_language, "error_label"), value="", visible=False),344            gr.update(label=t(selected_language, "debug_label")),345            gr.update(value=t(selected_language, "debug_button"))346        ]347 348        debug_info["step"] = "Interface update completed successfully"349        debug_info["last_error"] = ""350 351        return updates352 353    except Exception as e:354        error_msg = f"ERROR in interface update at step '{debug_info['step']}':\n{str(e)}\n\nTraceback:\n{traceback.format_exc()}"355        debug_info["last_error"] = error_msg356        debug_info["step"] = f"FAILED: {str(e)}"357 358        # Safe fallback359        safe_updates = [gr.update() for _ in range(10)]360        return safe_updates361 362 363def get_debug_status():364    """Get current debug status for display"""365    global debug_info366 367    status = f"""🔍 DEBUG STATUS:368Current Language: {debug_info.get('language', 'N/A')}369Last Timestamp: {debug_info.get('timestamp', 'N/A')}370Current Step: {debug_info.get('step', 'N/A')}371Last Error: {debug_info.get('last_error', 'None')}372 373Available Models: {list(available_models.keys())}374Current Model: {current_model_name or 'None loaded'}375Translation Cache Size: {len(translation_cache)}376"""377    return status378 379 380def safe_detect_objects(image, language_selector, translated_model_selector, threshold):381    """Safe wrapper for object detection with error handling"""382    global debug_info383 384    if image is None:385        debug_info["step"] = "No image provided"386        return None, "Please upload an image first.", ""387 388    try:389        result_image, info, error = detect_objects(image, language_selector, translated_model_selector, threshold)390 391        # Update error panel visibility based on whether there's an error392        error_visible = bool(error.strip())393 394        return (395            result_image,396            info,397            gr.update(value=error, visible=error_visible)398        )399 400    except Exception as e:401        error_message = f"Unexpected error in detection:\n{str(e)}\n\nStack trace:\n{traceback.format_exc()}"402        debug_info["last_error"] = error_message403        debug_info["step"] = f"UNEXPECTED ERROR: {str(e)}"404        print(error_message)405        return (406            image,407            "Detection failed due to unexpected error. See error panel below.",408            gr.update(value=error_message, visible=True)409        )410 411 412def build_app():413    with gr.Blocks(theme=gr.themes.Soft()) as app:414        with gr.Row():415            title = gr.Markdown(t("English", "title"))416 417        with gr.Row():418            with gr.Column(scale=1):419                language_selector = gr.Dropdown(420                    choices=["English", "Spanish", "French"],421                    value="English",422                    label=t("English", "dropdown_label")423                )424            with gr.Column(scale=1):425                model_selector = gr.Dropdown(426                    choices=get_translated_model_choices("English"),427                    value=t("English", "model_fast"),428                    label=t("English", "dropdown_detection_model_label")429                )430            with gr.Column(scale=1):431                threshold_slider = gr.Slider(432                    minimum=0.1,433                    maximum=0.95,434                    value=0.5,435                    step=0.05,436                    label=t("English", "threshold_label")437                )438 439        with gr.Row():440            with gr.Column(scale=1):441                input_image = gr.Image(type="pil", label=t("English", "input_label"))442                button = gr.Button(t("English", "button"), variant="primary")443            with gr.Column(scale=1):444                output_image = gr.Image(label=t("English", "output_label"))445                detection_info = gr.Textbox(446                    label=t("English", "info_label"),447                    lines=10,448                    max_lines=15449                )450 451        # Error panel - only visible when there are errors452        with gr.Row():453            error_panel = gr.Textbox(454                label=t("English", "error_label"),455                lines=8,456                max_lines=20,457                visible=False,458                elem_classes=["error-panel"]459            )460 461        # Debug panel - always visible for debugging in HF462        with gr.Row():463            debug_panel = gr.Textbox(464                label=t("English", "debug_label"),465                lines=10,466                max_lines=20,467                value="Application started - ready for debugging",468                visible=True469            )470 471        with gr.Row():472            debug_button = gr.Button(t("English", "debug_button"), size="sm")473 474        # Connect language change event475        language_selector.change(476            fn=update_interface,477            inputs=language_selector,478            outputs=[479                title,480                # language_selector, # <-- esta línea también debes eliminarla481                model_selector,482                threshold_slider,483                input_image,484                button,485                output_image,486                detection_info,487                error_panel,488                debug_panel,489                debug_button490            ],491            queue=True492        )493 494        # Connect detection button click event495        button.click(496            fn=safe_detect_objects,497            inputs=[input_image, language_selector, model_selector, threshold_slider],498            outputs=[output_image, detection_info, error_panel]499        )500 501        # Connect debug button click event502        debug_button.click(503            fn=get_debug_status,504            outputs=debug_panel505        )506 507    return app508 509 510# Initialize with default model and debug info511debug_info["step"] = "Initializing default model"512debug_info["timestamp"] = str(datetime.datetime.now())513load_model("DETR ResNet-50")514debug_info["step"] = "Application ready"515 516# Launch the application517if __name__ == "__main__":518    app = build_app()519    app.launch()