feat(ml,backend): implemente le service de scoring et GET /predictions (#37)
This commit is contained in:
+51
-5
@@ -57,6 +57,46 @@ validation. La coupure est **chronologique**, jamais un tirage aleatoire de lign
|
||||
aleatoire laisserait des lignes de validation "voir" des lignes d'entrainement via leurs
|
||||
lags/moyennes glissantes, une fuite qui masquerait un surapprentissage.
|
||||
|
||||
## Scoring
|
||||
|
||||
```bash
|
||||
uv run python -m enervision_ml.score --csv data/all_sites_combined.csv
|
||||
# ou, une fois la base peuplee et ML_DATABASE_URL positionnee :
|
||||
uv run python -m enervision_ml.score
|
||||
```
|
||||
|
||||
Calcule, pour chaque site (ou un seul avec `--site-id`), la consommation prevue de l'heure suivant
|
||||
sa derniere lecture connue, et ecrit une ligne dans `prediction`. Etapes, cf. `ML-START.md`
|
||||
section 2 :
|
||||
|
||||
1. Lit une fenetre recente de `reading`+`site` (21 jours par defaut, une marge au-dessus des 168h
|
||||
necessaires au lag hebdomadaire) plutot que tout l'historique -- le meme piege que celui deja
|
||||
corrige sur `GET /readings` (fenetre non plafonnee sur une hypertable).
|
||||
2. Ajoute une ligne "future" par site (l'heure suivante) et calcule ses features avec
|
||||
`enervision_ml.features.build_features`, **exactement** la meme fonction qu'a l'entrainement.
|
||||
3. Si le lag de 168h est absent (moins d'une semaine d'historique pour ce site) : ecrit
|
||||
`status="insufficient_data"` directement, sans jamais appeler LightGBM.
|
||||
4. Sinon : appelle `booster.predict(...)` et ecrit `status="available"` avec la valeur predite.
|
||||
|
||||
`--model` pointe vers le fichier entraine (`models/lightgbm-consumption.txt` par defaut).
|
||||
`model_reference` en base est le hache SHA-256 (tronque) du fichier modele, pas son nom de
|
||||
fichier : `train.py` reecrit toujours le meme chemin a chaque entrainement, donc le nom seul ne
|
||||
distinguerait pas deux versions du modele.
|
||||
|
||||
En mode `--csv`, rien n'est ecrit en base : c'est un instantane historique fige (l'heure "future"
|
||||
calculee a partir de la fin du CSV n'existe dans aucune base reelle), utile pour valider le
|
||||
pipeline sans base joignable.
|
||||
|
||||
**Limite assumee** : la feature `is_working_hours` de la ligne future est recopiee depuis la
|
||||
derniere lecture reelle, pas recalculee -- il n'existe aucune regle horaire ouvrable dans ce
|
||||
depot (elle vit dans le generateur du jeu de donnees d'origine). L'approximation n'est fausse
|
||||
qu'aux heures de bascule ouverture/fermeture, sur une seule feature parmi une dizaine, pour une
|
||||
prevision a un seul pas.
|
||||
|
||||
`prediction` n'a pas de contrainte d'unicite sur `(site_id, target_at)` : chaque run de scoring
|
||||
insere une nouvelle ligne plutot que d'ecraser la precedente, pour garder une trace de chaque
|
||||
prevision (utile plus tard pour comparer prevision et realise, surveillance de derive #44/#45).
|
||||
|
||||
## Commandes
|
||||
|
||||
```bash
|
||||
@@ -81,8 +121,14 @@ environnement de developpement pour le moment.
|
||||
## Piege a connaitre
|
||||
|
||||
`enervision_ml.features.build_features` est **le seul endroit** qui doit construire les features
|
||||
du modele, a l'entrainement comme au futur scoring (service #37, pas encore construit). Si les
|
||||
deux divergent meme legerement (une fenetre de moyenne glissante calculee differemment, par
|
||||
exemple), le modele recoit en production des features qui ne ressemblent plus a ce qu'il a
|
||||
appris, et ses predictions deviennent silencieusement mauvaises sans qu'aucune erreur ne se
|
||||
declenche. Ne jamais reecrire cette logique ailleurs : importer `enervision_ml.features`.
|
||||
du modele, a l'entrainement comme au scoring (`enervision_ml.score`). Si les deux divergent meme
|
||||
legerement (une fenetre de moyenne glissante calculee differemment, par exemple), le modele
|
||||
recoit en production des features qui ne ressemblent plus a ce qu'il a appris, et ses predictions
|
||||
deviennent silencieusement mauvaises sans qu'aucune erreur ne se declenche. Ne jamais reecrire
|
||||
cette logique ailleurs : importer `enervision_ml.features`.
|
||||
|
||||
## Et cote API ?
|
||||
|
||||
`GET /api/v1/predictions` (backend, `apps/backend`) lit ce que `enervision_ml.score` a ecrit dans
|
||||
`prediction` -- la derniere prevision par site, jamais un recalcul a la volee. FastAPI ne fait
|
||||
jamais tourner LightGBM lui-meme, cf. `ML-START.md` section 3.
|
||||
|
||||
@@ -16,6 +16,7 @@ Deux chemins, qui doivent produire le meme schema de sortie (colonnes `site_id`,
|
||||
colonne est renvoyee a `NaN`, que LightGBM gere nativement comme valeur manquante.
|
||||
"""
|
||||
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
|
||||
import pandas as pd
|
||||
@@ -34,6 +35,14 @@ OUTPUT_COLUMNS = [
|
||||
"capacity_kw",
|
||||
]
|
||||
|
||||
NUMERIC_COLUMNS = [
|
||||
"consumption_kwh",
|
||||
"temperature_celsius",
|
||||
"humidity_percent",
|
||||
"solar_irradiance_wm2",
|
||||
"capacity_kw",
|
||||
]
|
||||
|
||||
_READING_QUERY = text(
|
||||
"""
|
||||
SELECT
|
||||
@@ -53,10 +62,43 @@ _READING_QUERY = text(
|
||||
)
|
||||
|
||||
|
||||
_RECENT_READING_QUERY = text(
|
||||
"""
|
||||
SELECT
|
||||
r.site_id,
|
||||
r.timestamp,
|
||||
r.consumption_kwh,
|
||||
r.temperature_celsius,
|
||||
r.humidity_percent,
|
||||
r.solar_irradiance_wm2,
|
||||
r.is_working_hours,
|
||||
s.site_type,
|
||||
s.capacity_kw
|
||||
FROM reading r
|
||||
JOIN site s ON s.site_id = r.site_id
|
||||
WHERE r.timestamp >= :since
|
||||
ORDER BY r.site_id, r.timestamp
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def load_from_database(connection: Connectable) -> pd.DataFrame:
|
||||
"""Lit l'historique complet `reading` + `site` depuis PostgreSQL."""
|
||||
"""Lit l'historique complet `reading` + `site` depuis PostgreSQL. Entrainement seulement :
|
||||
le scoring n'a besoin que d'une fenetre recente, cf. `load_recent_from_database`.
|
||||
"""
|
||||
frame = pd.read_sql(_READING_QUERY, connection)
|
||||
return frame[OUTPUT_COLUMNS]
|
||||
return _typer(frame[OUTPUT_COLUMNS])
|
||||
|
||||
|
||||
def load_recent_from_database(connection: Connectable, *, since: datetime) -> pd.DataFrame:
|
||||
"""Lit `reading` + `site` depuis `since` seulement, pour le scoring.
|
||||
|
||||
Piege evite : un `SELECT` sans borne sur l'hypertable complete juste pour scorer le prochain
|
||||
pas horaire serait la meme erreur que celle corrigee sur `GET /readings` (fenetre non
|
||||
plafonnee sur une table pouvant porter des annees d'historique).
|
||||
"""
|
||||
frame = pd.read_sql(_RECENT_READING_QUERY, connection, params={"since": since})
|
||||
return _typer(frame[OUTPUT_COLUMNS])
|
||||
|
||||
|
||||
def load_from_csv(csv_path: Path) -> pd.DataFrame:
|
||||
@@ -65,4 +107,20 @@ def load_from_csv(csv_path: Path) -> pd.DataFrame:
|
||||
frame["capacity_kw"] = float("nan")
|
||||
frame["is_working_hours"] = frame["is_working_hours"].astype(bool)
|
||||
|
||||
return frame[OUTPUT_COLUMNS]
|
||||
return _typer(frame[OUTPUT_COLUMNS])
|
||||
|
||||
|
||||
def _typer(frame: pd.DataFrame) -> pd.DataFrame:
|
||||
"""Force le typage numerique attendu par LightGBM.
|
||||
|
||||
Piege reel, pas theorique : `site.capacity_kw` n'est peuple par aucun pipeline d'ingestion
|
||||
aujourd'hui (`historical_import.py` ne pose que `site_type`/`site_name`). Une colonne
|
||||
entierement `NULL` revient de `pd.read_sql` en dtype `object` plutot que `float64`, ce que
|
||||
LightGBM refuse ("pandas dtypes must be int, float or bool"). `pd.to_numeric` corrige aussi
|
||||
n'importe quelle autre colonne mesuree entierement absente sur une fenetre de scoring, pas
|
||||
seulement `capacity_kw`.
|
||||
"""
|
||||
typee = frame.copy()
|
||||
for colonne in NUMERIC_COLUMNS:
|
||||
typee[colonne] = pd.to_numeric(typee[colonne], errors="coerce")
|
||||
return typee
|
||||
|
||||
@@ -0,0 +1,281 @@
|
||||
"""Scoring du modele LightGBM : calcule et enregistre la consommation prevue du prochain pas
|
||||
horaire, par site.
|
||||
|
||||
CLI autonome, sur le meme gabarit que `enervision_ml.train` et
|
||||
`apps/backend/app/etl/historical_import.py`. Cf. `docs/ML-START.md`, section 2.
|
||||
|
||||
uv run python -m enervision_ml.score --csv ../ml/data/all_sites_combined.csv
|
||||
uv run python -m enervision_ml.score # lit ML_DATABASE_URL, ecrit dans `prediction`
|
||||
|
||||
Reutilise `enervision_ml.features.build_features` tel quel (jamais reecrit) : c'est la garantie
|
||||
contre le train/serve skew documentee dans ce module.
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import hashlib
|
||||
from dataclasses import dataclass
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from pathlib import Path
|
||||
from typing import Any, cast
|
||||
|
||||
import lightgbm as lgb
|
||||
import pandas as pd
|
||||
from sqlalchemy import create_engine, text
|
||||
from sqlalchemy.engine import Connection
|
||||
|
||||
from enervision_ml import config
|
||||
from enervision_ml.data import load_from_csv, load_recent_from_database
|
||||
from enervision_ml.features import TARGET_COLUMN, WEATHER_COLUMNS, build_features, feature_columns
|
||||
|
||||
# Marge au-dessus des 168h necessaires au lag hebdomadaire, pour absorber les trous de mesure.
|
||||
LOOKBACK = timedelta(days=21)
|
||||
|
||||
TARGET_METRIC = "consumption_kwh"
|
||||
PERIOD_MINUTES = 60
|
||||
LAG_168H_COLUMN = f"{TARGET_COLUMN}_lag_168h"
|
||||
INSUFFICIENT_DATA_REASON = (
|
||||
"Historique insuffisant : moins de 168h de consumption_kwh disponibles pour ce site."
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ScoredSite:
|
||||
site_id: str
|
||||
target_at: datetime
|
||||
status: str
|
||||
predicted_value: float | None
|
||||
failure_reason: str | None
|
||||
|
||||
|
||||
def model_reference(model_path: Path) -> str:
|
||||
"""Identifiant stable du modele utilise, insensible au fait que `train.py` reecrive
|
||||
toujours le meme nom de fichier a chaque entrainement (pas de versioning par nom, cf.
|
||||
`ml/README.md`)."""
|
||||
empreinte = hashlib.sha256(model_path.read_bytes()).hexdigest()
|
||||
return f"lightgbm-{empreinte[:12]}"
|
||||
|
||||
|
||||
def build_scoring_frame(recent: pd.DataFrame, *, site_id: str | None = None) -> pd.DataFrame:
|
||||
"""Ajoute une ligne future (l'heure suivant la derniere lecture connue) par site, et calcule
|
||||
ses features par `build_features` -- exactement comme a l'entrainement, seule la cible de
|
||||
cette ligne est inconnue.
|
||||
|
||||
Piege assume : `is_working_hours` de la ligne future est copie de la derniere lecture reelle,
|
||||
pas recalcule. Il n'existe aucune regle horaire ouvrable dans ce depot (elle vit dans le
|
||||
generateur du jeu de donnees d'origine, hors de ce code) ; l'approximation n'est fausse
|
||||
qu'aux heures de bascule (ouverture/fermeture), sur une seule feature parmi une dizaine, pour
|
||||
une prevision a un pas seulement.
|
||||
"""
|
||||
travail = recent if site_id is None else recent[recent["site_id"] == site_id]
|
||||
if travail.empty:
|
||||
return build_features(travail)
|
||||
|
||||
dernieres = (
|
||||
travail.sort_values("timestamp").groupby("site_id", as_index=False, sort=False).tail(1)
|
||||
).copy()
|
||||
dernieres["timestamp"] = dernieres["timestamp"] + pd.Timedelta(hours=1)
|
||||
dernieres[TARGET_COLUMN] = float("nan")
|
||||
# Meteo future inconnue (cf. piege documente dans `enervision_ml.features.build_features`) :
|
||||
# laisser `NaN` ici n'a aucun effet sur les features utilisees, qui ne prennent la meteo que
|
||||
# decalee.
|
||||
for colonne in WEATHER_COLUMNS:
|
||||
dernieres[colonne] = float("nan")
|
||||
|
||||
etendu = pd.concat([travail, dernieres], ignore_index=True)
|
||||
features = build_features(etendu)
|
||||
return features.groupby("site_id", as_index=False, sort=False).tail(1).reset_index(drop=True)
|
||||
|
||||
|
||||
def score(booster: lgb.Booster, scoring_frame: pd.DataFrame) -> list[ScoredSite]:
|
||||
resultats: list[ScoredSite] = []
|
||||
|
||||
insuffisants = scoring_frame[scoring_frame[LAG_168H_COLUMN].isna()]
|
||||
for enregistrement in _records(insuffisants):
|
||||
resultats.append(
|
||||
ScoredSite(
|
||||
site_id=enregistrement["site_id"],
|
||||
target_at=enregistrement["timestamp"].to_pydatetime(),
|
||||
status="insufficient_data",
|
||||
predicted_value=None,
|
||||
failure_reason=INSUFFICIENT_DATA_REASON,
|
||||
)
|
||||
)
|
||||
|
||||
suffisants = scoring_frame[scoring_frame[LAG_168H_COLUMN].notna()]
|
||||
if not suffisants.empty:
|
||||
typee = suffisants.copy()
|
||||
typee["site_type"] = typee["site_type"].astype("category")
|
||||
predictions = booster.predict(typee[feature_columns()])
|
||||
for enregistrement, valeur in zip(_records(suffisants), predictions, strict=True):
|
||||
resultats.append(
|
||||
ScoredSite(
|
||||
site_id=enregistrement["site_id"],
|
||||
target_at=enregistrement["timestamp"].to_pydatetime(),
|
||||
status="available",
|
||||
predicted_value=float(valeur),
|
||||
failure_reason=None,
|
||||
)
|
||||
)
|
||||
|
||||
return resultats
|
||||
|
||||
|
||||
def _records(frame: pd.DataFrame) -> list[dict[str, Any]]:
|
||||
return cast(list[dict[str, Any]], frame.to_dict(orient="records"))
|
||||
|
||||
|
||||
_INSERT_PREDICTION = text(
|
||||
"""
|
||||
INSERT INTO prediction (
|
||||
site_id, target_at, target_metric, period_minutes,
|
||||
predicted_value, model_reference, status, failure_reason
|
||||
) VALUES (
|
||||
:site_id, :target_at, :target_metric, :period_minutes,
|
||||
:predicted_value, :model_reference, :status, :failure_reason
|
||||
)
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def write_predictions(
|
||||
connection: Connection, resultats: list[ScoredSite], *, reference: str
|
||||
) -> None:
|
||||
"""Ecrit une ligne par site score. Insertion seule, jamais de mise a jour : `prediction`
|
||||
n'a pas de contrainte d'unicite sur `(site_id, target_at)`, chaque run garde sa propre trace
|
||||
plutot que d'ecraser la precedente -- utile plus tard pour comparer prevision et realise
|
||||
(surveillance de derive, #44/#45)."""
|
||||
if not resultats:
|
||||
return
|
||||
|
||||
lignes = [
|
||||
{
|
||||
"site_id": r.site_id,
|
||||
"target_at": r.target_at,
|
||||
"target_metric": TARGET_METRIC,
|
||||
"period_minutes": PERIOD_MINUTES,
|
||||
"predicted_value": r.predicted_value,
|
||||
"model_reference": reference,
|
||||
"status": r.status,
|
||||
"failure_reason": r.failure_reason,
|
||||
}
|
||||
for r in resultats
|
||||
]
|
||||
connection.execute(_INSERT_PREDICTION, lignes)
|
||||
|
||||
|
||||
def _load_recent(*, csv_path: Path | None, now: datetime | None) -> tuple[pd.DataFrame, datetime]:
|
||||
if csv_path is not None:
|
||||
brute = load_from_csv(csv_path)
|
||||
instant = now or (
|
||||
brute["timestamp"].max().to_pydatetime() if not brute.empty else datetime.now(UTC)
|
||||
)
|
||||
return brute[brute["timestamp"] >= instant - LOOKBACK], instant
|
||||
|
||||
instant = now or datetime.now(UTC)
|
||||
engine = create_engine(config.database_url())
|
||||
try:
|
||||
return load_recent_from_database(engine, since=instant - LOOKBACK), instant
|
||||
finally:
|
||||
engine.dispose()
|
||||
|
||||
|
||||
def run_scoring(
|
||||
*,
|
||||
model_path: Path,
|
||||
csv_path: Path | None = None,
|
||||
site_id: str | None = None,
|
||||
now: datetime | None = None,
|
||||
) -> list[ScoredSite]:
|
||||
"""Score le prochain pas horaire par site et l'ecrit dans `prediction`.
|
||||
|
||||
En mode `--csv`, rien n'est ecrit : c'est un instantane historique fige (l'heure "future"
|
||||
calculee n'existe dans aucune base reelle), utile pour valider le pipeline sans base
|
||||
joignable, cf. `ml/README.md`.
|
||||
"""
|
||||
recent, _instant = _load_recent(csv_path=csv_path, now=now)
|
||||
if site_id is not None:
|
||||
recent = recent[recent["site_id"] == site_id]
|
||||
|
||||
scoring_frame = build_scoring_frame(recent, site_id=site_id)
|
||||
if scoring_frame.empty:
|
||||
return []
|
||||
|
||||
booster = lgb.Booster(model_file=str(model_path))
|
||||
resultats = score(booster, scoring_frame)
|
||||
|
||||
if csv_path is None:
|
||||
reference = model_reference(model_path)
|
||||
engine = create_engine(config.database_url())
|
||||
try:
|
||||
with engine.begin() as connection:
|
||||
write_predictions(connection, resultats, reference=reference)
|
||||
finally:
|
||||
engine.dispose()
|
||||
|
||||
return resultats
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description="Scoring du modele LightGBM EnerVision")
|
||||
|
||||
parser.add_argument(
|
||||
"--model",
|
||||
type=Path,
|
||||
default=Path("models/lightgbm-consumption.txt"),
|
||||
help="Chemin du modele entraine. Defaut : models/lightgbm-consumption.txt.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--csv",
|
||||
type=Path,
|
||||
default=None,
|
||||
help=(
|
||||
"Instantane historique de demarrage/demo, rien n'est ecrit en base. Omis, lit "
|
||||
"ML_DATABASE_URL, se connecte a PostgreSQL et ecrit dans `prediction`."
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--site-id",
|
||||
default=None,
|
||||
help="Ne score que ce site. Omis, tous les sites presents dans la fenetre recente.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--now",
|
||||
type=_parse_instant,
|
||||
default=None,
|
||||
help=(
|
||||
"Instant de reference (ISO 8601), pour tester ou demontrer le scoring cote base sur "
|
||||
"des donnees anciennes (ex. le jeu de donnees historique, qui s'arrete fin 2024). "
|
||||
"Omis, horloge systeme reelle."
|
||||
),
|
||||
)
|
||||
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def _parse_instant(valeur: str) -> datetime:
|
||||
instant = datetime.fromisoformat(valeur)
|
||||
return instant if instant.tzinfo is not None else instant.replace(tzinfo=UTC)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
resultats = run_scoring(
|
||||
model_path=args.model, csv_path=args.csv, site_id=args.site_id, now=args.now
|
||||
)
|
||||
|
||||
if not resultats:
|
||||
print("Aucun site a scorer (aucune lecture recente dans la fenetre).")
|
||||
return
|
||||
|
||||
for r in resultats:
|
||||
if r.status == "available":
|
||||
print(f"{r.site_id} @ {r.target_at} : {r.predicted_value:.2f} kWh")
|
||||
else:
|
||||
print(f"{r.site_id} @ {r.target_at} : {r.status} ({r.failure_reason})")
|
||||
|
||||
if args.csv is not None:
|
||||
print("\nMode --csv : instantane historique, rien ecrit en base.")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
+3
-3
@@ -47,9 +47,9 @@ select = [
|
||||
"S",
|
||||
"PT",
|
||||
]
|
||||
# N806 : `X`/`y` (donnees/cible) est la convention scikit-learn/LightGBM, pas une variable mal
|
||||
# nommee.
|
||||
ignore = ["B008", "N806"]
|
||||
# N806/N803 : `X`/`y` (donnees/cible) est la convention scikit-learn/LightGBM, pas une variable
|
||||
# ou un argument mal nomme.
|
||||
ignore = ["B008", "N806", "N803"]
|
||||
|
||||
[tool.ruff.lint.per-file-ignores]
|
||||
"tests/**/*.py" = ["S101"]
|
||||
|
||||
@@ -0,0 +1,34 @@
|
||||
import pandas as pd
|
||||
|
||||
from enervision_ml.data import NUMERIC_COLUMNS, OUTPUT_COLUMNS, _typer
|
||||
|
||||
|
||||
def make_frame_with_object_dtype_capacity() -> pd.DataFrame:
|
||||
# Reproduit ce que `pd.read_sql` renvoie pour une colonne entierement `NULL` en base :
|
||||
# dtype `object` rempli de `None`, pas `float64` rempli de `NaN`.
|
||||
frame = pd.DataFrame(
|
||||
{colonne: [1.0, 2.0] for colonne in OUTPUT_COLUMNS if colonne not in NUMERIC_COLUMNS}
|
||||
)
|
||||
for colonne in NUMERIC_COLUMNS:
|
||||
frame[colonne] = pd.Series([None, None], dtype="object")
|
||||
return frame
|
||||
|
||||
|
||||
def test_typer_coerces_an_all_null_object_column_to_float() -> None:
|
||||
frame = make_frame_with_object_dtype_capacity()
|
||||
|
||||
typee = _typer(frame)
|
||||
|
||||
for colonne in NUMERIC_COLUMNS:
|
||||
assert typee[colonne].dtype == "float64"
|
||||
assert typee[colonne].isna().all()
|
||||
|
||||
|
||||
def test_typer_preserves_real_numeric_values() -> None:
|
||||
frame = make_frame_with_object_dtype_capacity()
|
||||
frame["capacity_kw"] = pd.Series([100.0, None], dtype="object")
|
||||
|
||||
typee = _typer(frame)
|
||||
|
||||
assert typee["capacity_kw"].tolist()[0] == 100.0
|
||||
assert pd.isna(typee["capacity_kw"].tolist()[1])
|
||||
@@ -0,0 +1,230 @@
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import pandas as pd
|
||||
import pytest
|
||||
|
||||
from enervision_ml.features import TARGET_COLUMN
|
||||
from enervision_ml.score import (
|
||||
LAG_168H_COLUMN,
|
||||
ScoredSite,
|
||||
build_scoring_frame,
|
||||
model_reference,
|
||||
run_scoring,
|
||||
score,
|
||||
write_predictions,
|
||||
)
|
||||
|
||||
|
||||
def make_recent(
|
||||
site_id: str, *, heures: int, depart: datetime, valeur: float = 10.0
|
||||
) -> pd.DataFrame:
|
||||
instants = [depart + timedelta(hours=h) for h in range(heures)]
|
||||
return pd.DataFrame(
|
||||
{
|
||||
"site_id": site_id,
|
||||
"timestamp": instants,
|
||||
TARGET_COLUMN: [valeur + h for h in range(heures)],
|
||||
"temperature_celsius": 15.0,
|
||||
"humidity_percent": 50.0,
|
||||
"solar_irradiance_wm2": 0.0,
|
||||
"is_working_hours": True,
|
||||
"site_type": "office",
|
||||
"capacity_kw": 100.0,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
class FakeBooster:
|
||||
def __init__(self, valeur: float = 42.0) -> None:
|
||||
self.valeur = valeur
|
||||
self.appels: list[int] = []
|
||||
|
||||
def predict(self, X: Any) -> list[float]:
|
||||
self.appels.append(len(X))
|
||||
return [self.valeur] * len(X)
|
||||
|
||||
|
||||
class FakeConnection:
|
||||
def __init__(self) -> None:
|
||||
self.appels: list[tuple[Any, Any]] = []
|
||||
|
||||
def execute(self, statement: Any, parameters: Any = None) -> None:
|
||||
self.appels.append((statement, parameters))
|
||||
|
||||
|
||||
def test_build_scoring_frame_adds_one_row_per_site_one_hour_after_the_last_reading() -> None:
|
||||
depart = datetime(2026, 1, 1, tzinfo=UTC)
|
||||
recent = pd.concat(
|
||||
[
|
||||
make_recent("site-a", heures=200, depart=depart),
|
||||
make_recent("site-b", heures=200, depart=depart),
|
||||
],
|
||||
ignore_index=True,
|
||||
)
|
||||
|
||||
scoring_frame = build_scoring_frame(recent)
|
||||
|
||||
assert set(scoring_frame["site_id"]) == {"site-a", "site-b"}
|
||||
derniere_lecture = depart + timedelta(hours=199)
|
||||
assert (scoring_frame["timestamp"] == derniere_lecture + timedelta(hours=1)).all()
|
||||
|
||||
|
||||
def test_build_scoring_frame_computes_lags_from_real_history() -> None:
|
||||
depart = datetime(2026, 1, 1, tzinfo=UTC)
|
||||
recent = make_recent("site-a", heures=200, depart=depart)
|
||||
|
||||
scoring_frame = build_scoring_frame(recent)
|
||||
|
||||
ligne = scoring_frame.iloc[0]
|
||||
# La cible future n'existe pas : le lag d'1h doit valoir la toute derniere valeur reelle.
|
||||
assert ligne[f"{TARGET_COLUMN}_lag_1h"] == recent[TARGET_COLUMN].iloc[-1]
|
||||
|
||||
|
||||
def test_build_scoring_frame_flags_insufficient_history_under_168_hours() -> None:
|
||||
depart = datetime(2026, 1, 1, tzinfo=UTC)
|
||||
recent = make_recent("site-a", heures=100, depart=depart)
|
||||
|
||||
scoring_frame = build_scoring_frame(recent)
|
||||
|
||||
assert pd.isna(scoring_frame.iloc[0][LAG_168H_COLUMN])
|
||||
|
||||
|
||||
def test_build_scoring_frame_accepts_a_full_week_of_history() -> None:
|
||||
depart = datetime(2026, 1, 1, tzinfo=UTC)
|
||||
recent = make_recent("site-a", heures=169, depart=depart)
|
||||
|
||||
scoring_frame = build_scoring_frame(recent)
|
||||
|
||||
assert not pd.isna(scoring_frame.iloc[0][LAG_168H_COLUMN])
|
||||
|
||||
|
||||
def test_build_scoring_frame_filters_to_a_single_site() -> None:
|
||||
depart = datetime(2026, 1, 1, tzinfo=UTC)
|
||||
recent = pd.concat(
|
||||
[
|
||||
make_recent("site-a", heures=200, depart=depart),
|
||||
make_recent("site-b", heures=200, depart=depart),
|
||||
],
|
||||
ignore_index=True,
|
||||
)
|
||||
|
||||
scoring_frame = build_scoring_frame(recent, site_id="site-a")
|
||||
|
||||
assert scoring_frame["site_id"].tolist() == ["site-a"]
|
||||
|
||||
|
||||
def test_build_scoring_frame_returns_empty_when_there_is_no_recent_reading() -> None:
|
||||
recent = make_recent("site-a", heures=0, depart=datetime(2026, 1, 1, tzinfo=UTC))
|
||||
|
||||
scoring_frame = build_scoring_frame(recent)
|
||||
|
||||
assert scoring_frame.empty
|
||||
|
||||
|
||||
def test_score_marks_insufficient_history_without_calling_the_model() -> None:
|
||||
depart = datetime(2026, 1, 1, tzinfo=UTC)
|
||||
scoring_frame = build_scoring_frame(make_recent("site-a", heures=100, depart=depart))
|
||||
booster = FakeBooster()
|
||||
|
||||
resultats = score(booster, scoring_frame) # type: ignore[arg-type]
|
||||
|
||||
assert resultats == [
|
||||
ScoredSite(
|
||||
site_id="site-a",
|
||||
target_at=resultats[0].target_at,
|
||||
status="insufficient_data",
|
||||
predicted_value=None,
|
||||
failure_reason=resultats[0].failure_reason,
|
||||
)
|
||||
]
|
||||
assert booster.appels == []
|
||||
|
||||
|
||||
def test_score_predicts_when_history_is_sufficient() -> None:
|
||||
depart = datetime(2026, 1, 1, tzinfo=UTC)
|
||||
scoring_frame = build_scoring_frame(make_recent("site-a", heures=200, depart=depart))
|
||||
booster = FakeBooster(valeur=99.5)
|
||||
|
||||
resultats = score(booster, scoring_frame) # type: ignore[arg-type]
|
||||
|
||||
assert len(resultats) == 1
|
||||
assert resultats[0].status == "available"
|
||||
assert resultats[0].predicted_value == 99.5
|
||||
assert resultats[0].failure_reason is None
|
||||
assert booster.appels == [1]
|
||||
|
||||
|
||||
def test_write_predictions_does_nothing_when_there_is_nothing_to_write() -> None:
|
||||
connection = FakeConnection()
|
||||
|
||||
write_predictions(connection, [], reference="lightgbm-test") # type: ignore[arg-type]
|
||||
|
||||
assert connection.appels == []
|
||||
|
||||
|
||||
def test_write_predictions_sends_one_row_per_result() -> None:
|
||||
connection = FakeConnection()
|
||||
resultats = [
|
||||
ScoredSite("site-a", datetime(2026, 1, 1, tzinfo=UTC), "available", 42.0, None),
|
||||
ScoredSite(
|
||||
"site-b",
|
||||
datetime(2026, 1, 1, tzinfo=UTC),
|
||||
"insufficient_data",
|
||||
None,
|
||||
"pas assez d'historique",
|
||||
),
|
||||
]
|
||||
|
||||
write_predictions(connection, resultats, reference="lightgbm-test") # type: ignore[arg-type]
|
||||
|
||||
assert len(connection.appels) == 1
|
||||
_, lignes = connection.appels[0]
|
||||
assert len(lignes) == 2
|
||||
assert lignes[0]["model_reference"] == "lightgbm-test"
|
||||
assert lignes[0]["target_metric"] == "consumption_kwh"
|
||||
assert lignes[0]["period_minutes"] == 60
|
||||
|
||||
|
||||
def test_model_reference_is_stable_for_the_same_file_content(tmp_path: Path) -> None:
|
||||
model_path = tmp_path / "model.txt"
|
||||
model_path.write_bytes(b"contenu-du-modele")
|
||||
|
||||
assert model_reference(model_path) == model_reference(model_path)
|
||||
|
||||
|
||||
def test_model_reference_changes_with_the_file_content(tmp_path: Path) -> None:
|
||||
premier = tmp_path / "model-a.txt"
|
||||
premier.write_bytes(b"version-1")
|
||||
second = tmp_path / "model-b.txt"
|
||||
second.write_bytes(b"version-2")
|
||||
|
||||
assert model_reference(premier) != model_reference(second)
|
||||
|
||||
|
||||
def test_run_scoring_in_csv_mode_scores_without_touching_a_database(tmp_path: Path) -> None:
|
||||
depart = datetime(2026, 1, 1, tzinfo=UTC)
|
||||
frame = pd.concat(
|
||||
[
|
||||
make_recent("site-a", heures=400, depart=depart),
|
||||
make_recent("site-b", heures=400, depart=depart),
|
||||
],
|
||||
ignore_index=True,
|
||||
)
|
||||
csv_path = tmp_path / "recent.csv"
|
||||
frame.to_csv(csv_path, index=False)
|
||||
|
||||
model_path = tmp_path / "model.txt"
|
||||
model_path.write_bytes(b"peu importe le contenu pour ce test")
|
||||
|
||||
with pytest.MonkeyPatch.context() as monkeypatch:
|
||||
monkeypatch.setattr(
|
||||
"enervision_ml.score.lgb.Booster", lambda model_file: FakeBooster(valeur=7.0)
|
||||
)
|
||||
|
||||
resultats = run_scoring(model_path=model_path, csv_path=csv_path)
|
||||
|
||||
assert {r.site_id for r in resultats} == {"site-a", "site-b"}
|
||||
assert all(r.status == "available" for r in resultats)
|
||||
assert all(r.predicted_value == 7.0 for r in resultats)
|
||||
Reference in New Issue
Block a user