ARViewer / api /register.py
wizardsmagic's picture
Add Wizara Vision API endpoints without changing UI or inference.
3820d5b
Raw
History Blame Contribute Delete
2.5 kB
"""Register Wizara Vision API endpoints on the Gradio Server app."""
from __future__ import annotations
import os
from typing import Any
from .endpoints import handle_detect, handle_ocr, handle_unsupported
def register_wizara_api(app, deps: dict[str, Any]) -> None:
base_dir = os.path.dirname(os.path.abspath(deps["__file__"]))
run_image_gpu = deps["run_image_gpu_api"]
generate_prompt = deps["generate_raw_prompt"]
parse_results = deps["parse_mixed_results"]
shared = {
"base_dir": base_dir,
"run_image_gpu": run_image_gpu,
"generate_prompt": generate_prompt,
"parse_results": parse_results,
}
@app.api(name="detect")
def detect_api(
image: Any = None,
categories: str = "objects",
task_type: str = "Detection",
model_mode: str = "hybrid",
temp: float = 0.7,
top_p: float = 0.9,
top_k: int = 20,
short_size: int | None = None,
advanced_settings: str = "",
) -> dict:
"""Detect objects in an uploaded image and return normalized JSON."""
return handle_detect(
image_file=image,
categories=categories,
task_type=task_type,
model_mode=model_mode,
temp=temp,
top_p=top_p,
top_k=top_k,
short_size=short_size,
advanced_settings=advanced_settings,
**shared,
)
@app.api(name="ocr")
def ocr_api(
image: Any = None,
model_mode: str = "hybrid",
temp: float = 0.7,
top_p: float = 0.9,
top_k: int = 20,
short_size: int | None = None,
advanced_settings: str = "",
) -> dict:
"""Run OCR localization using the existing LocateAnything OCR task."""
return handle_ocr(
image_file=image,
model_mode=model_mode,
temp=temp,
top_p=top_p,
top_k=top_k,
short_size=short_size,
advanced_settings=advanced_settings,
**shared,
)
@app.api(name="caption")
def caption_api() -> dict:
return handle_unsupported("caption")
@app.api(name="segment")
def segment_api() -> dict:
return handle_unsupported("segment")
@app.api(name="count")
def count_api() -> dict:
return handle_unsupported("count")
@app.api(name="classify")
def classify_api() -> dict:
return handle_unsupported("classify")