80 lines
2.4 KiB
Python
80 lines
2.4 KiB
Python
from __future__ import annotations
|
|
|
|
from fastapi import FastAPI
|
|
|
|
from .engines import VisionInferenceEngine
|
|
from .schemas import (
|
|
HealthResponse,
|
|
InferenceRequest,
|
|
InferenceResponse,
|
|
ModelTestInferenceRequest,
|
|
ModelTestInferenceResponse,
|
|
)
|
|
|
|
app = FastAPI(title="Rail UAV Vision Inference Service", version="0.1.0")
|
|
engine = VisionInferenceEngine()
|
|
|
|
|
|
@app.get("/health", response_model=HealthResponse)
|
|
def health() -> HealthResponse:
|
|
return HealthResponse(status="ok")
|
|
|
|
|
|
@app.get("/api/v1/runtime")
|
|
def runtime_status() -> dict:
|
|
return engine.runtime_status()
|
|
|
|
|
|
@app.post("/api/v1/runtime/unload")
|
|
def unload_models(model_group: str | None = None) -> dict:
|
|
return {"unloaded": engine.unload(model_group), "runtime": engine.runtime_status()}
|
|
|
|
|
|
@app.post("/api/v1/runtime/load")
|
|
def load_model(model_group: str, model_version: str | None = None) -> dict:
|
|
return {"model": engine.load(model_group, model_version), "runtime": engine.runtime_status()}
|
|
|
|
|
|
@app.post("/api/v1/inference/detect", response_model=InferenceResponse)
|
|
def detect(request: InferenceRequest) -> InferenceResponse:
|
|
selected_group = request.model_group
|
|
selected_version = request.model_version or "v1.0.0"
|
|
results = (
|
|
engine.test_infer(
|
|
request.resource,
|
|
request.scenes,
|
|
selected_group,
|
|
selected_version,
|
|
request.parameters,
|
|
)
|
|
if selected_group
|
|
else engine.infer(request.resource, request.scenes)
|
|
)
|
|
if request.deployment_id:
|
|
for result in results:
|
|
result.attributes["deployment_id"] = request.deployment_id
|
|
return InferenceResponse(
|
|
job_id=request.job_id,
|
|
model={"name": selected_group or "rail-multi-scene-inference", "version": selected_version},
|
|
results=results,
|
|
)
|
|
|
|
|
|
@app.post("/api/v1/inference/test", response_model=ModelTestInferenceResponse)
|
|
def test_inference(request: ModelTestInferenceRequest) -> ModelTestInferenceResponse:
|
|
results = engine.test_infer(
|
|
request.resource,
|
|
request.scenes,
|
|
request.model_group,
|
|
request.model_version,
|
|
request.parameters,
|
|
)
|
|
if request.deployment_id:
|
|
for result in results:
|
|
result.attributes["deployment_id"] = request.deployment_id
|
|
return ModelTestInferenceResponse(
|
|
model={"name": request.model_group, "version": request.model_version},
|
|
parameter_snapshot=request.parameters,
|
|
results=results,
|
|
)
|