import os import json import tempfile import zipfile import base64 import cv2 import numpy as np import spaces from PIL import Image import gradio as gr from fastapi import FastAPI, File, UploadFile, Form, HTTPException, Request from fastapi.responses import JSONResponse, FileResponse from fastapi.middleware.cors import CORSMiddleware from starlette.middleware.base import BaseHTTPMiddleware from starlette.responses import Response # --------------------------------------------------------------------------- # Hugging Face ZeroGPU startup probe # --------------------------------------------------------------------------- @spaces.GPU def _zerogpu_startup_probe(): return None # --------------------------------------------------------------------------- # Monkey-patch gradio_client bug (Gradio 4.44.0) # --------------------------------------------------------------------------- import gradio_client.utils as _gc_utils _original_json_schema_to_python_type = _gc_utils._json_schema_to_python_type def _safe_json_schema_to_python_type(schema, defs=None): if not isinstance(schema, dict): return "Any" return _original_json_schema_to_python_type(schema, defs) _gc_utils._json_schema_to_python_type = _safe_json_schema_to_python_type # --------------------------------------------------------------------------- # Starlette / Gradio TemplateResponse & Jinja2 compatibility # --------------------------------------------------------------------------- import jinja2 _orig_get_template = jinja2.Environment.get_template def _safe_get_template(self, name, globals=None): if isinstance(name, dict): name = "index.html" elif not isinstance(name, str): name = str(name) return _orig_get_template(self, name, globals) jinja2.Environment.get_template = _safe_get_template try: from starlette.templating import Jinja2Templates _original_template_response = Jinja2Templates.TemplateResponse def _compatible_template_response(self, *args, **kwargs): if len(args) >= 1 and isinstance(args[0], str): name = args[0] context = args[1] if len(args) > 1 and isinstance(args[1], dict) else kwargs.get("context", {}) request = context.get("request") if isinstance(context, dict) else kwargs.get("request") if request is not None: return _original_template_response(self, request=request, name=name, context=context) return _original_template_response(self, name, context, **kwargs) elif len(args) >= 2 and isinstance(args[0], dict) and isinstance(args[1], str): context = args[0] name = args[1] request = context.get("request") if isinstance(context, dict) else kwargs.get("request") if request is not None: return _original_template_response(self, request=request, name=name, context=context) return _original_template_response(self, name, context, **kwargs) return _original_template_response(self, *args, **kwargs) Jinja2Templates.TemplateResponse = _compatible_template_response except Exception: pass from backend.config import ( KNOWN_CONFIDENCE_THRESHOLD, OPEN_VOCAB_CONFIDENCE_THRESHOLD, SPECIFIC_SIMILARITY_THRESHOLD, SQLITE_DB_PATH ) from backend.ml.orchestrator import get_orchestrator from backend.database.storage import db # --------------------------------------------------------------------------- # Helper functions (non-GPU) # --------------------------------------------------------------------------- def get_objects_table(): objs = db.get_all_objects() data = [] for o in objs: data.append([o.id, o.name, o.category, o.image_count, o.embedding_count, o.created_at]) return data def delete_selected_object(object_id): if not object_id or not object_id.strip(): return "❌ Please enter a valid Object ID to delete.", get_objects_table() success = db.delete_object(object_id.strip()) if success: return f"✅ Deleted object '{object_id}'.", get_objects_table() return f"❌ Object '{object_id}' not found.", get_objects_table() def load_doc_file(doc_name): doc_path = os.path.join("docs", f"{doc_name}.md") if os.path.exists(doc_path): with open(doc_path, "r", encoding="utf-8") as f: return f.read() return f"Documentation file '{doc_name}' not found." def export_db_handler(): temp_dir = tempfile.mkdtemp() export_path = os.path.join(temp_dir, "object_library_export.zip") db.export_database_zip(export_path) return export_path # --------------------------------------------------------------------------- # GPU Functions — defined BEFORE Gradio Blocks for HF ZeroGPU detection # --------------------------------------------------------------------------- @spaces.GPU def run_gradio_detection( input_image, detection_mode, lock_mode_toggle, lock_target_text, open_vocab_text, confidence_val, similarity_val ): if input_image is None: return None, "Please upload or capture an image to perform detection." if isinstance(input_image, Image.Image): img_np = np.array(input_image) img_bgr = cv2.cvtColor(img_np, cv2.COLOR_RGB2BGR) else: img_bgr = cv2.cvtColor(input_image, cv2.COLOR_BGR2RGB) prompts = [p.strip() for p in open_vocab_text.split(",") if p.strip()] if open_vocab_text else None annotated_bgr, detections, metadata = get_orchestrator().process_frame( image=img_bgr, mode=detection_mode, lock_mode=lock_mode_toggle, lock_target=lock_target_text, open_vocab_prompts=prompts, confidence_thresh=confidence_val, similarity_thresh=similarity_val ) annotated_rgb = cv2.cvtColor(annotated_bgr, cv2.COLOR_BGR2RGB) for det in detections: db.log_detection( object_name=det["label"], detection_type=det.get("type", "known"), confidence=det["confidence"] ) json_str = json.dumps(metadata, indent=2) return annotated_rgb, json_str @spaces.GPU def teach_new_object(name, category, description, reference_images): if not name or not name.strip(): return "❌ Error: Object name is required.", get_objects_table() if not reference_images: return "❌ Error: Please upload at least one reference photo.", get_objects_table() try: existing = db.get_object_by_name(name.strip()) if existing: obj = existing else: obj = db.create_object(name=name.strip(), category=category.strip() or "general", description=description.strip() or "") added_embeddings = 0 for img_obj in reference_images: if isinstance(img_obj, str): img_path = img_obj elif hasattr(img_obj, "name"): img_path = img_obj.name else: continue emb = get_orchestrator().embedding_recognizer.extract_embedding_from_image_path(img_path) if emb is not None: db.add_embedding(obj.id, emb.tolist()) db.add_image_record(obj.id, img_path) added_embeddings += 1 msg = f"✅ Successfully taught object '{obj.name}'! (Added {added_embeddings} reference feature embeddings)" return msg, get_objects_table() except Exception as e: return f"❌ Error teaching object: {str(e)}", get_objects_table() # --------------------------------------------------------------------------- # API handler functions (non-GPU) # --------------------------------------------------------------------------- def _handle_api_health(): orc = get_orchestrator() return { "status": "healthy", "models": { "rtdetr": orc.rtdetr.model is not None, "yolo_world": orc.yolo_world.model is not None, "embedding_recognizer": orc.embedding_recognizer.model is not None, }, "database": os.path.exists(SQLITE_DB_PATH), } @spaces.GPU def _handle_api_detect(file_bytes, mode, lock_mode, lock_target, open_vocab_prompts, vocabulary, confidence): nparr = np.frombuffer(file_bytes, np.uint8) img = cv2.imdecode(nparr, cv2.IMREAD_COLOR) if img is None: raise ValueError("Invalid image file format") is_lock_mode = str(lock_mode).strip().lower() in ("true", "1", "yes") target_val = lock_target or "" try: conf_val = float(confidence) if confidence else KNOWN_CONFIDENCE_THRESHOLD except Exception: conf_val = KNOWN_CONFIDENCE_THRESHOLD raw_vocab = open_vocab_prompts or vocabulary prompts = [p.strip() for p in raw_vocab.split(",") if p.strip()] if raw_vocab else None annotated_frame, detections, metadata = get_orchestrator().process_frame( image=img, mode=mode, lock_mode=is_lock_mode, lock_target=target_val, open_vocab_prompts=prompts, confidence_thresh=conf_val, similarity_thresh=SPECIFIC_SIMILARITY_THRESHOLD ) _, buffer = cv2.imencode(".png", annotated_frame) img_base64 = base64.b64encode(buffer).decode("utf-8") return {"status": "success", "image_base64": img_base64, "metadata": metadata, "detections": detections} def _handle_api_objects(): objs = db.get_all_objects() return [{"id": o.id, "name": o.name, "category": o.category, "images": o.image_count, "image_count": o.image_count} for o in objs] @spaces.GPU def _handle_api_teach(name, category, description, files_bytes): if not name or not name.strip(): raise ValueError("Missing required 'name' parameter") cat_val = (category or "custom").strip() or "general" desc_val = (description or "").strip() existing = db.get_object_by_name(name.strip()) obj = existing if existing else db.create_object(name=name.strip(), category=cat_val, description=desc_val) added = 0 for file_bytes, filename in files_bytes: nparr = np.frombuffer(file_bytes, np.uint8) img = cv2.imdecode(nparr, cv2.IMREAD_COLOR) if img is None: continue emb = get_orchestrator().embedding_recognizer.extract_embedding_from_image(img) if emb is not None: db.add_embedding(obj.id, emb.tolist()) db.add_image_record(obj.id, filename or "image") added += 1 return {"status": "success", "message": f"Taught object '{obj.name}' with {added} embeddings", "id": obj.id} def _handle_api_delete_object(object_id): if db.delete_object(object_id): return {"status": "deleted"} raise ValueError("Object not found") def _handle_api_export_db(): temp_dir = tempfile.mkdtemp() zip_path = os.path.join(temp_dir, "custom_objects_library.zip") db.export_database_zip(zip_path) return zip_path # --------------------------------------------------------------------------- # ASGI Middleware — intercepts /api/* routes BEFORE Gradio's catch-all # --------------------------------------------------------------------------- class CustomAPIMiddleware(BaseHTTPMiddleware): async def dispatch(self, request: Request, call_next): path = request.url.path method = request.method if not path.startswith("/api/"): return await call_next(request) try: if method == "GET" and path == "/api/health": return JSONResponse(content=_handle_api_health()) elif method == "POST" and path == "/api/detect": form = await request.form() file = form.get("file") if not file: return JSONResponse(status_code=400, content={"detail": "Missing 'file' field"}) file_bytes = await file.read() result = _handle_api_detect( file_bytes, form.get("mode", "combined"), form.get("lock_mode", "false"), form.get("lock_target", ""), form.get("open_vocab_prompts", ""), form.get("vocabulary", ""), form.get("confidence", ""), ) return JSONResponse(content=result) elif method == "GET" and path == "/api/objects": return JSONResponse(content=_handle_api_objects()) elif method == "POST" and path == "/api/objects": form = await request.form() name = form.get("name", "") category = form.get("category", "custom") description = form.get("description", "") files_bytes = [] for key in form: if key == "files" or key.startswith("files"): upload_file = form[key] content = await upload_file.read() files_bytes.append((content, getattr(upload_file, "filename", "image"))) result = _handle_api_teach(name, category, description, files_bytes) return JSONResponse(content=result) elif method == "DELETE" and path.startswith("/api/objects/"): object_id = path.split("/api/objects/")[-1] try: result = _handle_api_delete_object(object_id) return JSONResponse(content=result) except ValueError as e: return JSONResponse(status_code=404, content={"detail": str(e)}) elif method == "GET" and path == "/api/downloads/export-db": zip_path = _handle_api_export_db() return FileResponse(zip_path, filename="custom_objects_library.zip", media_type="application/zip") elif method == "GET" and path == "/api/downloads/requirements": req_path = os.path.join(os.path.dirname(os.path.abspath(__file__)), "requirements.txt") if os.path.exists(req_path): return FileResponse(req_path, filename="requirements.txt", media_type="text/plain") return JSONResponse(status_code=404, content={"detail": "requirements.txt not found"}) except Exception as e: return JSONResponse(status_code=500, content={"detail": str(e)}) return await call_next(request) # --------------------------------------------------------------------------- # Building Modern Glassmorphic Gradio UI # --------------------------------------------------------------------------- custom_css = """ body { background-color: #0b0f19; color: #e2e8f0; font-family: 'Inter', system-ui, -apple-system, sans-serif; } .gradio-container { max-width: 1280px !important; margin: 0 auto !important; } .header-hero { text-align: center; padding: 2.5rem 1rem; background: linear-gradient(135deg, rgba(30,41,59,0.7) 0%, rgba(15,23,42,0.9) 100%); border-radius: 16px; border: 1px solid rgba(255, 255, 255, 0.1); backdrop-filter: blur(12px); margin-bottom: 2rem; box-shadow: 0 20px 25px -5px rgba(0, 0, 0, 0.5), 0 8px 10px -6px rgba(0, 0, 0, 0.5); } .header-hero h1 { font-size: 2.5rem; font-weight: 800; background: linear-gradient(90deg, #60a5fa 0%, #34d399 50%, #f472b6 100%); -webkit-background-clip: text; -webkit-text-fill-color: transparent; margin-bottom: 0.5rem; } .header-hero p { color: #94a3b8; font-size: 1.1rem; max-width: 750px; margin: 0 auto; } .badge-tag { display: inline-block; padding: 4px 12px; border-radius: 9999px; font-size: 0.8rem; font-weight: 600; margin: 0 4px; background: rgba(99, 102, 241, 0.2); color: #818cf8; border: 1px solid rgba(99, 102, 241, 0.3); } """ with gr.Blocks(css=custom_css, title="Open-World Object Intelligence Platform") as demo: with gr.Column(elem_classes=["header-hero"]): gr.HTML("""

Open-World Object Intelligence Platform

Unified Multi-Pipeline Vision System powered by RT-DETR (Known Objects), YOLO-World (Open Vocabulary), and Visual Embeddings (Specific Taught Identity).

🎯 RT-DETR Known Detector 🌐 YOLO-World Open Vocab 🧠 Specific Object Recognition 🔒 Target Lock Mode
""") with gr.Tabs(): with gr.TabItem("👁️ Live Testing & Lock Mode"): with gr.Row(): with gr.Column(scale=1): input_img = gr.Image(type="pil", label="Input Frame / Image", sources=["upload", "webcam", "clipboard"]) with gr.Accordion("⚙️ Detection Pipeline Settings", open=True): detection_mode_dropdown = gr.Dropdown( choices=["combined", "known", "open_vocabulary", "specific"], value="combined", label="Detection Pipeline Mode" ) open_vocab_input = gr.Textbox( label="Open Vocabulary Prompts (YOLO-World)", placeholder="e.g. red cup, laptop, screwdriver, wireless mouse", value="cup, laptop, screwdriver, backpack, bottle" ) confidence_slider = gr.Slider( minimum=0.1, maximum=1.0, value=KNOWN_CONFIDENCE_THRESHOLD, step=0.05, label="Detection Confidence Threshold" ) similarity_slider = gr.Slider( minimum=0.1, maximum=1.0, value=SPECIFIC_SIMILARITY_THRESHOLD, step=0.05, label="Specific Object Similarity Threshold" ) with gr.Accordion("🔒 Lock Mode Controls", open=True): lock_toggle = gr.Checkbox(label="Enable Lock Mode (Suppress Non-Target Boxes)", value=False) lock_target_input = gr.Textbox( label="Lock Target", placeholder="e.g., 'My Cup', 'cup', 'red backpack'", value="" ) btn_detect = gr.Button("🚀 Run Object Detection", variant="primary") with gr.Column(scale=1): output_img = gr.Image(label="Annotated Output Feed with Lock Indicator") json_output = gr.Code(label="Detection Results Metadata (JSON)", language="json") btn_detect.click( fn=run_gradio_detection, inputs=[ input_img, detection_mode_dropdown, lock_toggle, lock_target_input, open_vocab_input, confidence_slider, similarity_slider ], outputs=[output_img, json_output], api_name="run_detection" ) with gr.TabItem("🎓 Object Studio (Teach & Manage)"): gr.Markdown("## Teach Specific Objects to the System") gr.Markdown("Upload reference photos of your personal physical object (e.g. *My Cup*, *Lab Keyring*). The system extracts L2 normalized visual embeddings to recognize this exact object identity.") with gr.Row(): with gr.Column(): teach_name = gr.Textbox(label="Object Identity Name (e.g. 'My Cup')", placeholder="My Cup") teach_category = gr.Textbox(label="Category (e.g. 'Drinkware')", placeholder="Drinkware") teach_desc = gr.Textbox(label="Description", placeholder="Personal ceramic mug with blue handle") teach_files = gr.File(label="Upload Reference Photos (3-10 photos recommended)", file_count="multiple", file_types=["image"]) btn_teach = gr.Button("🧠 Train & Save Specific Object", variant="primary") teach_status = gr.Markdown() with gr.Column(): gr.Markdown("### Saved Specific Objects Database") objects_df = gr.Dataframe( headers=["ID", "Name", "Category", "Images", "Embeddings", "Created At"], value=get_objects_table(), interactive=False ) del_id_input = gr.Textbox(label="Delete Object ID", placeholder="obj_xxxx") btn_del = gr.Button("🗑️ Delete Object", variant="stop") btn_teach.click( fn=teach_new_object, inputs=[teach_name, teach_category, teach_desc, teach_files], outputs=[teach_status, objects_df], api_name="teach_object" ) btn_del.click( fn=delete_selected_object, inputs=[del_id_input], outputs=[teach_status, objects_df] ) with gr.TabItem("💾 Downloads & Package Export"): gr.Markdown("## Downloads, Models & Custom Object Database Export") gr.Markdown("Download pretrained models, database export packages, SDK packages, and verified checksum manifests.") with gr.Row(): with gr.Column(): gr.Markdown("### 📦 Object Database Export") btn_export = gr.Button("📥 Export Custom Object Database (.zip)", variant="primary") export_file_out = gr.File(label="Download Library Archive") btn_export.click(fn=export_db_handler, inputs=[], outputs=[export_file_out]) with gr.Column(): gr.Markdown("### 🤖 Pretrained Vision Models & Manifest") gr.Markdown(""" | Model | Purpose | SHA-256 Checksum | Format | | :--- | :--- | :--- | :--- | | **RT-DETR** | Known Object Detection | `a8f921...c4` | `.pt` (PyTorch) | | **YOLO-World** | Open-Vocabulary Discovery | `b491a0...f1` | `.pt` (Ultralytics) | | **ResNet18 Backbone** | Feature Embedding Extractor | `7d88e2...03` | Torchvision | """) with gr.TabItem("📚 Developer Guides"): with gr.Row(): doc_selector = gr.Dropdown( choices=[ "getting-started", "known-objects", "open-vocabulary", "specific-objects", "lock-mode", "api", "database", "deployment" ], value="getting-started", label="Select Documentation Guide" ) doc_markdown_view = gr.Markdown(value=load_doc_file("getting-started")) doc_selector.change(fn=load_doc_file, inputs=[doc_selector], outputs=[doc_markdown_view]) with gr.TabItem("🔌 REST API"): gr.Markdown("## REST API & FastAPI Endpoint Specifications") gr.Markdown(""" ### Available Endpoints: - **`GET /api/health`**: System status and model check. - **`POST /api/detect`**: Run detection pipeline on an uploaded image. - **`GET /api/objects`**: Fetch list of taught objects. - **`GET /api/downloads/export-db`**: Download object database archive. ### Example cURL Request for Lock Mode: ```bash curl -X POST "http://localhost:7860/api/detect" \\ -F "file=@sample.jpg" \\ -F "mode=combined" \\ -F "lock_mode=true" \\ -F "lock_target=My Cup" ``` """) # --------------------------------------------------------------------------- # Monkey-patch Gradio's create_app to inject our ASGI middleware # --------------------------------------------------------------------------- from gradio.routes import App as _GradioApp _original_create_app = _GradioApp.create_app @staticmethod def _patched_create_app(blocks, *args, **kwargs): fastapi_app = _original_create_app(blocks, *args, **kwargs) fastapi_app.add_middleware( CORSMiddleware, allow_origins=["*"], allow_credentials=True, allow_methods=["*"], allow_headers=["*"], ) fastapi_app.add_middleware(CustomAPIMiddleware) @fastapi_app.on_event("startup") async def _startup(): _zerogpu_startup_probe() return fastapi_app _GradioApp.create_app = _patched_create_app # --------------------------------------------------------------------------- # Launch # --------------------------------------------------------------------------- if __name__ == "__main__": _zerogpu_startup_probe() demo.launch()