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
quality / test (3.11) (pull_request) Has been cancelled
quality / test (3.13) (pull_request) 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

1
.gitignore vendored
View File

@@ -4,6 +4,7 @@
/.vscode /.vscode
__pycache__/ __pycache__/
*.pyc *.pyc
*.egg-info/
.mypy_cache/ .mypy_cache/
.pytest_cache/ .pytest_cache/
.ruff_cache/ .ruff_cache/

View File

@@ -8,3 +8,4 @@
- Persistente, validierte und gegen Path Traversal gehärtete Model Registry - Persistente, validierte und gegen Path Traversal gehärtete Model Registry
- Reproduzierbares Packaging, CI-Gates und gehärteter non-root Container - Reproduzierbares Packaging, CI-Gates und gehärteter non-root Container
- Definierte API-Fehler und korrigierte Evaluationsmetriken - Definierte API-Fehler und korrigierte Evaluationsmetriken
- Scheduler-tauglicher Retraining-Service mit API und atomischem Registry-Update

View File

@@ -43,6 +43,7 @@ uvicorn app.main:app --reload
- `http://127.0.0.1:8000/docs/` - OpenAPI-Dokumentation - `http://127.0.0.1:8000/docs/` - OpenAPI-Dokumentation
- `http://127.0.0.1:8000/v1/entities` - Home-Assistant-Entities - `http://127.0.0.1:8000/v1/entities` - Home-Assistant-Entities
- `http://127.0.0.1:8000/ml/health` - Registry-/Serving-Health - `http://127.0.0.1:8000/ml/health` - Registry-/Serving-Health
- `POST http://127.0.0.1:8000/ml/retrain` - Modell-Metadaten aktualisieren
Ohne vollständige HA-Konfiguration liefert `/v1/entities` bewusst `503`. Ohne vollständige HA-Konfiguration liefert `/v1/entities` bewusst `503`.

View File

@@ -1,5 +1,14 @@
"""Machine-Learning-Grundbausteine für SillyHome Next.""" """Machine-Learning-Grundbausteine für SillyHome Next."""
__all__ = ["FeatureStore", "FeatureVector", "TrainedArtifact", "TrainingPipeline"] __all__ = [
"FeatureStore",
"FeatureVector",
"RetrainingResult",
"RetrainingService",
"TrainedArtifact",
"TrainingPipeline",
"retrain_model",
]
from app.ml.feature_store import FeatureStore, FeatureVector from app.ml.feature_store import FeatureStore, FeatureVector
from app.ml.retraining import RetrainingResult, RetrainingService, retrain_model
from app.ml.training import TrainedArtifact, TrainingPipeline from app.ml.training import TrainedArtifact, TrainingPipeline

View File

@@ -5,6 +5,7 @@ import logging
import os import os
from pathlib import Path from pathlib import Path
import re import re
from threading import RLock
from collections.abc import Iterable from collections.abc import Iterable
from app.ml.training import TrainedArtifact from app.ml.training import TrainedArtifact
@@ -19,22 +20,31 @@ class ModelRegistry:
self._root = Path(root).resolve() self._root = Path(root).resolve()
self._root.mkdir(parents=True, exist_ok=True) self._root.mkdir(parents=True, exist_ok=True)
self._artifacts: dict[str, TrainedArtifact] = {} self._artifacts: dict[str, TrainedArtifact] = {}
self._lock = RLock()
self._load_existing() self._load_existing()
def register(self, artifact: TrainedArtifact) -> TrainedArtifact: def register(self, artifact: TrainedArtifact) -> TrainedArtifact:
registered, _ = self.register_with_status(artifact)
return registered
def register_with_status(self, artifact: TrainedArtifact) -> tuple[TrainedArtifact, bool]:
self._validate_artifact_id(artifact.artifact_id) self._validate_artifact_id(artifact.artifact_id)
self._persist(artifact) with self._lock:
self._artifacts[artifact.artifact_id] = artifact replaced = artifact.artifact_id in self._artifacts
return artifact self._persist(artifact)
self._artifacts[artifact.artifact_id] = artifact
return artifact, replaced
def load_artifact(self, artifact_id: str) -> TrainedArtifact: def load_artifact(self, artifact_id: str) -> TrainedArtifact:
self._validate_artifact_id(artifact_id) self._validate_artifact_id(artifact_id)
if artifact_id not in self._artifacts: with self._lock:
raise KeyError(f"Artifact '{artifact_id}' nicht registriert.") if artifact_id not in self._artifacts:
return self._artifacts[artifact_id] raise KeyError(f"Artifact '{artifact_id}' nicht registriert.")
return self._artifacts[artifact_id]
def list_models(self) -> Iterable[TrainedArtifact]: def list_models(self) -> Iterable[TrainedArtifact]:
return list(self._artifacts.values()) with self._lock:
return [self._artifacts[key] for key in sorted(self._artifacts)]
def _load_existing(self) -> None: def _load_existing(self) -> None:
for source in sorted(self._root.glob("*.json")): for source in sorted(self._root.glob("*.json")):

43
app/ml/retraining.py Normal file
View File

@@ -0,0 +1,43 @@
from __future__ import annotations
from collections.abc import Iterable
from dataclasses import dataclass
from app.ml.feature_store import FeatureStore, FeatureVector
from app.ml.registry.model_registry import ModelRegistry
from app.ml.training import TrainedArtifact, TrainingPipeline
@dataclass(frozen=True)
class RetrainingResult:
artifact: TrainedArtifact
replaced: bool
class RetrainingService:
"""Runs one retraining cycle without owning scheduling or background threads."""
def __init__(self, registry: ModelRegistry) -> None:
self._registry = registry
def retrain(
self,
artifact_id: str,
vectors: Iterable[FeatureVector],
) -> RetrainingResult:
store = FeatureStore()
store.add_batch(vectors)
pipeline = TrainingPipeline(store)
artifact = pipeline.run(artifact_id)
_, replaced = self._registry.register_with_status(artifact)
return RetrainingResult(artifact=artifact, replaced=replaced)
def retrain_model(
registry: ModelRegistry,
artifact_id: str,
vectors: Iterable[FeatureVector],
) -> RetrainingResult:
"""Scheduler-compatible entry point for exactly one retraining run."""
return RetrainingService(registry).retrain(artifact_id, vectors)

View File

@@ -10,6 +10,7 @@ from pydantic import BaseModel, Field
from app.ml.feature_store import FeatureVector from app.ml.feature_store import FeatureVector
from app.ml.predictor import Predictor from app.ml.predictor import Predictor
from app.ml.registry.model_registry import ModelRegistry from app.ml.registry.model_registry import ModelRegistry
from app.ml.retraining import retrain_model
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -45,6 +46,23 @@ class ModelsResponse(BaseModel):
models: list[str] 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) @router.get("/health", response_model=HealthResponse, status_code=200)
def health() -> HealthResponse: def health() -> HealthResponse:
return HealthResponse(status="ok") return HealthResponse(status="ok")
@@ -57,6 +75,31 @@ def list_models(request: Request) -> ModelsResponse:
return ModelsResponse(models=models) 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) @router.post("/predict", response_model=PredictResponse, status_code=200)
def predict(payload: PredictRequest, request: Request) -> PredictResponse: def predict(payload: PredictRequest, request: Request) -> PredictResponse:
registry = _require_registry(request) registry = _require_registry(request)

View File

@@ -12,6 +12,7 @@ Modell-Artefakt- und Vorhersage-Schnittstelle.
- Standard: `http://127.0.0.1:8000/ml` - Standard: `http://127.0.0.1:8000/ml`
- Health: `/health` - Health: `/health`
- Modelle: `/models` - Modelle: `/models`
- Retraining: `/retrain`
- Einzelvorhersage: `/predict` - Einzelvorhersage: `/predict`
- Batchvorhersage: `/batch` - Batchvorhersage: `/batch`
@@ -65,6 +66,35 @@ Einzelne Vorhersage für einen Sensor.
} }
``` ```
### `POST /ml/retrain`
Trainiert die Artefakt-Metadaten aus neuen Sensordaten. Existiert `modelId`
bereits, wird das Artefakt atomisch ersetzt und beim nächsten Prozessstart aus
dem Modellverzeichnis geladen.
**Request**
```json
{
"modelId": "home-model",
"samples": [
{
"sensor_id": "sensor.kitchen",
"values": {"temperature": 21.0},
"label": "occupied"
}
]
}
```
**Antwort**
```json
{
"model_id": "home-model",
"supported_sensors": ["sensor.kitchen"],
"replaced": false
}
```
### `POST /ml/batch` ### `POST /ml/batch`
Batch-Vorhersage für mehrere Sensorwerte. Batch-Vorhersage für mehrere Sensorwerte.
@@ -114,11 +144,15 @@ Batch-Vorhersage für mehrere Sensorwerte.
## Betrieb ## Betrieb
Die produktive App lädt Artefakte aus `SILLYHOME_MODEL_STORE`. Neue Artefakte Die produktive App lädt Artefakte aus `SILLYHOME_MODEL_STORE`. Neue Artefakte
werden derzeit intern über `ModelRegistry.register(...)` registriert. Die werden über `/ml/retrain`, `RetrainingService` oder direkt über
Registry speichert validiertes JSON atomisch und lädt es beim Neustart. `ModelRegistry.register(...)` registriert. Die Registry speichert validiertes
JSON atomisch und lädt es beim Neustart. Die API sollte nur in einem
vertrauenswürdigen Netz oder hinter einem authentifizierenden Reverse Proxy
erreichbar sein.
## Verweise ## Verweise
- `app/ml/predictor.py` - `app/ml/predictor.py`
- `app/ml/retraining.py`
- `app/ml/registry/model_registry.py` - `app/ml/registry/model_registry.py`
- `backend/routes/ml.py` - `backend/routes/ml.py`

View File

@@ -38,6 +38,21 @@ Der Report enthält:
Das trainierte Artefakt kann anschließend über `ModelRegistry.register(artifact)` bereitgestellt werden. Die ML-Serving-API stellt es unter `/ml/predict` und `/ml/batch` zur Verfügung. Das trainierte Artefakt kann anschließend über `ModelRegistry.register(artifact)` bereitgestellt werden. Die ML-Serving-API stellt es unter `/ml/predict` und `/ml/batch` zur Verfügung.
## 5. Retraining ausführen
`RetrainingService.retrain(...)` führt genau einen Trainingslauf aus und ersetzt
ein vorhandenes Artefakt mit derselben ID atomisch in der Registry:
```python
service = RetrainingService(registry)
result = service.retrain("home-model", vectors)
```
Scheduler, Cronjobs oder Home-Assistant-Automationen können alternativ die
zustandslose Funktion `retrain_model(registry, artifact_id, vectors)` aufrufen.
Der Service startet bewusst keinen eigenen Hintergrundprozess. Über
`POST /ml/retrain` kann derselbe Ablauf per API angestoßen werden.
## Hinweise ## Hinweise
- Für reproduzierbare Sensor-Reihenfolgen wird in `TrainingPipeline.run(...)` eine sortierte Sensor-Liste verwendet. - Für reproduzierbare Sensor-Reihenfolgen wird in `TrainingPipeline.run(...)` eine sortierte Sensor-Liste verwendet.
- Fehlende Trainingsdaten lösen `ValueError` aus; nicht registrierte Artefakte lösen `KeyError` aus. - Fehlende Trainingsdaten lösen `ValueError` aus; nicht registrierte Artefakte lösen `KeyError` aus.

View File

@@ -50,3 +50,60 @@ def test_unsupported_sensor_returns_422(tmp_path: Path) -> None:
) )
assert response.status_code == 422 assert response.status_code == 422
def test_retrain_creates_and_replaces_persisted_model(tmp_path: Path) -> None:
from app.ml.registry.model_registry import ModelRegistry
registry = ModelRegistry(tmp_path)
with TestClient(app) as client:
app.state.registry = registry
created = client.post(
"/ml/retrain",
json={
"modelId": "home-model",
"samples": [
{
"sensor_id": "sensor.kitchen",
"values": {"temperature": 21.0},
}
],
},
)
replaced = client.post(
"/ml/retrain",
json={
"modelId": "home-model",
"samples": [
{
"sensor_id": "sensor.bedroom",
"values": {"temperature": 18.0},
}
],
},
)
assert created.status_code == 200
assert created.json() == {
"model_id": "home-model",
"supported_sensors": ["sensor.kitchen"],
"replaced": False,
}
assert replaced.status_code == 200
assert replaced.json() == {
"model_id": "home-model",
"supported_sensors": ["sensor.bedroom"],
"replaced": True,
}
restarted = ModelRegistry(tmp_path)
assert restarted.load_artifact("home-model").supported_sensors == ("sensor.bedroom",)
def test_retrain_rejects_empty_samples() -> None:
with TestClient(app) as client:
response = client.post(
"/ml/retrain",
json={"modelId": "home-model", "samples": []},
)
assert response.status_code == 422

View File

@@ -19,6 +19,17 @@ def test_registry_loads_persisted_artifacts_after_restart(tmp_path: Path) -> Non
assert restarted.load_artifact("model-v1") == artifact assert restarted.load_artifact("model-v1") == artifact
def test_registry_replaces_persisted_artifact_after_restart(tmp_path: Path) -> None:
registry = ModelRegistry(tmp_path)
registry.register(TrainedArtifact("model-v1", ("sensor.kitchen",)))
replacement = TrainedArtifact("model-v1", ("sensor.bedroom",))
registry.register(replacement)
assert registry.load_artifact("model-v1") == replacement
assert ModelRegistry(tmp_path).load_artifact("model-v1") == replacement
@pytest.mark.parametrize("artifact_id", ["../escape", "nested/model", "..", ""]) @pytest.mark.parametrize("artifact_id", ["../escape", "nested/model", "..", ""])
def test_registry_rejects_unsafe_artifact_ids(tmp_path: Path, artifact_id: str) -> None: def test_registry_rejects_unsafe_artifact_ids(tmp_path: Path, artifact_id: str) -> None:
registry = ModelRegistry(tmp_path) registry = ModelRegistry(tmp_path)

View File

@@ -0,0 +1,39 @@
from __future__ import annotations
from pathlib import Path
import pytest
from app.ml.feature_store import FeatureVector
from app.ml.registry.model_registry import ModelRegistry
from app.ml.retraining import RetrainingService, retrain_model
def _vector(sensor_id: str) -> FeatureVector:
return FeatureVector(sensor_id=sensor_id, values={"temperature": 21.0})
def test_retraining_registers_new_artifact(tmp_path: Path) -> None:
registry = ModelRegistry(tmp_path)
result = retrain_model(registry, "home-model", [_vector("sensor.kitchen")])
assert result.replaced is False
assert registry.load_artifact("home-model") == result.artifact
def test_retraining_replaces_existing_artifact(tmp_path: Path) -> None:
registry = ModelRegistry(tmp_path)
service = RetrainingService(registry)
service.retrain("home-model", [_vector("sensor.kitchen")])
result = service.retrain("home-model", [_vector("sensor.bedroom")])
assert result.replaced is True
assert result.artifact.supported_sensors == ("sensor.bedroom",)
assert ModelRegistry(tmp_path).load_artifact("home-model") == result.artifact
def test_retraining_rejects_empty_training_data(tmp_path: Path) -> None:
with pytest.raises(ValueError, match="keine Trainingsdaten"):
retrain_model(ModelRegistry(tmp_path), "home-model", [])