Spaces:
Sleeping
Sleeping
File size: 2,695 Bytes
72e2b6e | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 | from fastapi import APIRouter, HTTPException, Request
from api.schemas import (
BatchRequest,
BatchResponse,
PredictRequest,
PredictResponse,
TopIntent,
)
from src.monitoring.drift import append_current_data
router = APIRouter()
def _to_predict_response(result, ab_variant: str | None = None) -> PredictResponse:
return PredictResponse(
intent=result.intent,
confidence=result.confidence,
top5=[TopIntent(**item) for item in result.top5],
latency_ms=result.latency_ms,
is_oos=result.is_oos,
model_used=result.model_used,
ab_variant=ab_variant,
)
@router.post("/predict", response_model=PredictResponse)
def predict(payload: PredictRequest, request: Request) -> PredictResponse:
predictors = request.app.state.predictors
ab_router = request.app.state.ab_router
if payload.model_type:
if payload.model_type not in predictors:
raise HTTPException(status_code=400, detail=f"unknown model_type: {payload.model_type}")
result = predictors[payload.model_type].predict(payload.text)
response = _to_predict_response(result)
else:
result_dict = ab_router.predict(payload.text)
response = PredictResponse(
intent=result_dict["intent"],
confidence=result_dict["confidence"],
top5=[TopIntent(**item) for item in result_dict["top5"]],
latency_ms=result_dict["latency_ms"],
is_oos=result_dict["is_oos"],
model_used=result_dict["model_used"],
ab_variant=result_dict["ab_variant"],
)
append_current_data(
texts=[payload.text],
predictions=[response.intent],
confidences=[response.confidence],
)
return response
@router.post("/predict/batch", response_model=BatchResponse)
def predict_batch(payload: BatchRequest, request: Request) -> BatchResponse:
predictors = request.app.state.predictors
model_type = payload.model_type or "transformer"
if model_type not in predictors:
raise HTTPException(status_code=400, detail=f"unknown model_type: {model_type}")
predictor = predictors[model_type]
results = predictor.predict_batch(payload.texts)
responses = [_to_predict_response(r) for r in results]
total_latency = sum(r.latency_ms for r in responses)
append_current_data(
texts=payload.texts,
predictions=[r.intent for r in responses],
confidences=[r.confidence for r in responses],
)
return BatchResponse(predictions=responses, total_latency_ms=round(total_latency, 3))
|