hackathon-ENGIN/engin/model.py
Jakub Famulski 2 8e6dcb750d
All checks were successful
ENGIN CI / Build, test and smoke (push) Successful in 2m58s
prod: load versioned model artifact at runtime
2026-08-25 13:06:14 +02:00

90 lines
3.3 KiB
Python

"""Prediction model adapters used through a small dependency-injection boundary."""
from __future__ import annotations
from dataclasses import dataclass
from pathlib import Path
from typing import Protocol
import pandas as pd
from .artifact import ArtifactMetadata, load_model_artifact
from .errors import InferenceError
from .inference import DiagnosticModels, predict_test, validate_submission
@dataclass(frozen=True)
class PredictionBundle:
submission: pd.DataFrame
diagnostics: pd.DataFrame
class PredictionModel(Protocol):
def predict(self, frame: pd.DataFrame) -> PredictionBundle: ...
class SklearnPredictionModel:
"""Runtime adapter around a verified, pre-trained model artifact."""
def __init__(
self, models: DiagnosticModels, metadata: ArtifactMetadata | None = None
) -> None:
self._models = models
self.metadata = metadata
@classmethod
def from_artifact(
cls,
directory: Path,
*,
expected_model_version: str | None = None,
) -> "SklearnPredictionModel":
try:
artifact = load_model_artifact(
directory, expected_model_version=expected_model_version
)
except Exception as exc:
raise InferenceError(
"Nie udało się załadować zweryfikowanego artefaktu modelu.",
hint="Sprawdź manifest, checksumę oraz zgodność wersji zależności.",
) from exc
return cls(artifact.models, artifact.metadata)
def predict(self, frame: pd.DataFrame) -> PredictionBundle:
try:
submission, diagnostics = predict_test(self._models, frame)
validate_submission(submission, frame)
except Exception as exc:
raise InferenceError(
"Model odrzucił pomiary podczas predykcji.",
hint="Zweryfikuj kompletność silników oraz brak nietypowych wartości w widmie.",
) from exc
diagnostics = diagnostics.copy()
diagnostics["model_version"] = (
self.metadata.model_version if self.metadata else "unversioned"
)
diagnostics["model_artifact_sha256"] = (
self.metadata.model_sha256 if self.metadata else "unavailable"
)
return PredictionBundle(submission, diagnostics)
class PrecomputedPredictionModel:
"""Read-only demo fallback; accepts only the exact precomputed key set."""
def __init__(self, submission: pd.DataFrame, diagnostics: pd.DataFrame) -> None:
self._submission = submission.reset_index(drop=True).copy()
self._diagnostics = diagnostics.reset_index(drop=True).copy()
def predict(self, frame: pd.DataFrame) -> PredictionBundle:
keys = ["engine_id", "cylinder"]
if not frame[keys].reset_index(drop=True).equals(self._submission[keys]):
raise InferenceError(
"Tryb awaryjny obsługuje wyłącznie dołączony zestaw demonstracyjny.",
hint="Przywróć artefakt modelu i uruchom ponownie aplikację, aby diagnozować własne pliki.",
)
diagnostics = self._diagnostics.copy()
diagnostics["model_version"] = "demo-precomputed"
diagnostics["model_artifact_sha256"] = "unavailable"
return PredictionBundle(self._submission.copy(), diagnostics)