muhammadpriv001's picture
fix: add @spaces.GPU to API handlers so middleware calls get GPU context
4583f36
Raw History Blame Contribute Delete
25.3 kB
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("""
<h1>Open-World Object Intelligence Platform</h1>
<p>Unified Multi-Pipeline Vision System powered by <strong>RT-DETR</strong> (Known Objects), <strong>YOLO-World</strong> (Open Vocabulary), and <strong>Visual Embeddings</strong> (Specific Taught Identity).</p>
<div style="margin-top: 1rem;">
<span class="badge-tag">🎯 RT-DETR Known Detector</span>
<span class="badge-tag">🌐 YOLO-World Open Vocab</span>
<span class="badge-tag">🧠 Specific Object Recognition</span>
<span class="badge-tag">πŸ”’ Target Lock Mode</span>
</div>
""")
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()