feat(ml): enregistrer le modèle dans le MLflow Model Registry
This commit is contained in:
@@ -181,7 +181,11 @@ def _log_to_mlflow(
|
|||||||
)
|
)
|
||||||
mlflow.log_metrics({f"model_{cle}": valeur for cle, valeur in model_metrics.items()})
|
mlflow.log_metrics({f"model_{cle}": valeur for cle, valeur in model_metrics.items()})
|
||||||
mlflow.log_metrics({f"baseline_{cle}": valeur for cle, valeur in baseline_metrics.items()})
|
mlflow.log_metrics({f"baseline_{cle}": valeur for cle, valeur in baseline_metrics.items()})
|
||||||
mlflow.lightgbm.log_model(booster, name="model")
|
mlflow.lightgbm.log_model(
|
||||||
|
booster,
|
||||||
|
name="model",
|
||||||
|
registered_model_name="consumption-forecast-lightgbm",
|
||||||
|
)
|
||||||
mlflow.log_artifact(str(model_output))
|
mlflow.log_artifact(str(model_output))
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ from pathlib import Path
|
|||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
|
import pytest
|
||||||
|
|
||||||
from enervision_ml.features import TARGET_COLUMN, build_features, feature_columns
|
from enervision_ml.features import TARGET_COLUMN, build_features, feature_columns
|
||||||
from enervision_ml.train import chronological_split, prepare_dataset, train
|
from enervision_ml.train import chronological_split, prepare_dataset, train
|
||||||
@@ -74,3 +75,18 @@ def test_train_runs_end_to_end_on_synthetic_data_and_beats_a_dummy_baseline(
|
|||||||
assert model_metrics["n_observations"] > 0
|
assert model_metrics["n_observations"] > 0
|
||||||
assert model_metrics["mae"] >= 0
|
assert model_metrics["mae"] >= 0
|
||||||
assert baseline_metrics["n_observations"] == model_metrics["n_observations"]
|
assert baseline_metrics["n_observations"] == model_metrics["n_observations"]
|
||||||
|
assert model_metrics["mae"] < baseline_metrics["mae"]
|
||||||
|
|
||||||
|
def test_train_raises_when_the_validation_window_is_empty(tmp_path: Path) -> None:
|
||||||
|
depart = datetime(2026, 1, 1, tzinfo=UTC)
|
||||||
|
frame = make_frame("site-a", heures=50, depart=depart) # trop court pour un lag de 168h
|
||||||
|
csv_path = tmp_path / "trop_court.csv"
|
||||||
|
frame.to_csv(csv_path, index=False)
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match="Fenetre d'entrainement ou de validation vide"):
|
||||||
|
train(
|
||||||
|
csv_path=csv_path,
|
||||||
|
model_output=tmp_path / "model.txt",
|
||||||
|
test_fraction=0.2,
|
||||||
|
tracking_uri=f"sqlite:///{tmp_path / 'mlflow.db'}",
|
||||||
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user