Spaces:
Running on Zero
Running on Zero
Download app.py from muhammadpriv001/Object-Intelligence-Backend: direct link, hf CLI and curl.
- Browser
- Download file 25.3 kB
-
https://huggingface.co/spaces/muhammadpriv001/Object-Intelligence-Backend/resolve/main/app.py
- Command line
-
hf download hf://spaces/muhammadpriv001/Object-Intelligence-Backend/app.py
-
curl -L -o app.py https://huggingface.co/spaces/muhammadpriv001/Object-Intelligence-Backend/resolve/main/app.py
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 | |
| # --------------------------------------------------------------------------- | |
| 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 | |
| # --------------------------------------------------------------------------- | |
| 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 | |
| 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), | |
| } | |
| 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] | |
| 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 | |
| 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) | |
| async def _startup(): | |
| _zerogpu_startup_probe() | |
| return fastapi_app | |
| _GradioApp.create_app = _patched_create_app | |
| # --------------------------------------------------------------------------- | |
| # Launch | |
| # --------------------------------------------------------------------------- | |
| if __name__ == "__main__": | |
| _zerogpu_startup_probe() | |
| demo.launch() | |