@@ -10,6 +10,7 @@ from pydantic import BaseModel, Field
|
||||
from app.ml.feature_store import FeatureVector
|
||||
from app.ml.predictor import Predictor
|
||||
from app.ml.registry.model_registry import ModelRegistry
|
||||
from app.ml.retraining import retrain_model
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -45,6 +46,23 @@ class ModelsResponse(BaseModel):
|
||||
models: list[str]
|
||||
|
||||
|
||||
class TrainingSample(BaseModel):
|
||||
sensor_id: str = Field(min_length=1)
|
||||
values: dict[str, float]
|
||||
label: str | None = None
|
||||
|
||||
|
||||
class RetrainRequest(BaseModel):
|
||||
model_id: str = Field(..., alias="modelId", min_length=1, max_length=128)
|
||||
samples: list[TrainingSample] = Field(min_length=1)
|
||||
|
||||
|
||||
class RetrainResponse(BaseModel):
|
||||
model_id: str
|
||||
supported_sensors: list[str]
|
||||
replaced: bool
|
||||
|
||||
|
||||
@router.get("/health", response_model=HealthResponse, status_code=200)
|
||||
def health() -> HealthResponse:
|
||||
return HealthResponse(status="ok")
|
||||
@@ -57,6 +75,31 @@ def list_models(request: Request) -> ModelsResponse:
|
||||
return ModelsResponse(models=models)
|
||||
|
||||
|
||||
@router.post("/retrain", response_model=RetrainResponse, status_code=200)
|
||||
def retrain(payload: RetrainRequest, request: Request) -> RetrainResponse:
|
||||
registry = _require_registry(request)
|
||||
vectors = [
|
||||
FeatureVector(
|
||||
sensor_id=sample.sensor_id,
|
||||
values=sample.values,
|
||||
label=sample.label,
|
||||
)
|
||||
for sample in payload.samples
|
||||
]
|
||||
try:
|
||||
result = retrain_model(registry, payload.model_id, vectors)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_CONTENT,
|
||||
detail=str(exc),
|
||||
) from exc
|
||||
return RetrainResponse(
|
||||
model_id=result.artifact.artifact_id,
|
||||
supported_sensors=list(result.artifact.supported_sensors),
|
||||
replaced=result.replaced,
|
||||
)
|
||||
|
||||
|
||||
@router.post("/predict", response_model=PredictResponse, status_code=200)
|
||||
def predict(payload: PredictRequest, request: Request) -> PredictResponse:
|
||||
registry = _require_registry(request)
|
||||
|
||||
Reference in New Issue
Block a user