ML-007: add retraining pipeline and API
Some checks failed
quality / test (3.11) (push) Has been cancelled
quality / test (3.13) (push) Has been cancelled

Closes #13
This commit is contained in:
2026-06-13 19:10:17 +02:00
parent ecd32d4813
commit 840c404c1c
12 changed files with 274 additions and 10 deletions

View File

@@ -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)