ML-005 vorbereiten: Registry, API-Routen und kompatibler Predictor

This commit is contained in:
2026-06-11 12:04:52 +02:00
parent 79e883f77d
commit fad517e56a
5 changed files with 198 additions and 5 deletions

View File

@@ -1,21 +1,31 @@
from __future__ import annotations
import logging
from functools import lru_cache
from typing import Sequence
from app.ml.feature_store import FeatureVector
from app.ml.training import TrainingPipeline, TrainedArtifact
from app.ml.registry.model_registry import ModelRegistry
from app.ml.training import TrainedArtifact, TrainingPipeline
logger = logging.getLogger(__name__)
class Predictor:
def __init__(self, pipeline: TrainingPipeline) -> None:
def __init__(
self,
pipeline: TrainingPipeline | None = None,
registry: ModelRegistry | None = None,
) -> None:
if isinstance(pipeline, ModelRegistry) and registry is None:
registry = pipeline
pipeline = None
if pipeline is None and registry is None:
raise ValueError("Predictor erfordert TrainingPipeline oder ModelRegistry.")
self._pipeline = pipeline
self._registry = registry
def predict(self, artifact_id: str, entity: FeatureVector) -> str:
artifact = self._pipeline.export(artifact_id)
artifact = self._get_artifact(artifact_id)
if entity.sensor_id not in artifact.supported_sensors:
raise ValueError(
f"Sensor '{entity.sensor_id}' wird vom Modell '{artifact_id}' nicht unterstützt."
@@ -30,4 +40,11 @@ class Predictor:
artifacts = list(pipeline._artifacts)
if not artifacts:
raise ValueError("Kein trainiertes Modell gefunden.")
return pipeline.export(artifacts[-1])
return pipeline.export(artifacts[-1])
def _get_artifact(self, artifact_id: str) -> TrainedArtifact:
if self._registry is not None:
return self._registry.load_artifact(artifact_id)
if self._pipeline is not None:
return self._pipeline.export(artifact_id)
raise RuntimeError("Predictor nicht initialisiert.")