TP3 Partie 1 : log_model + Model Registry + promotion par alias

This commit is contained in:
Johan LEROY
2026-07-22 11:22:50 +02:00
parent e6a05fbbc7
commit 9a52395c91
6 changed files with 127 additions and 0 deletions

View File

@@ -11,5 +11,12 @@ MLFLOW_EXPERIMENT_NAME=tp02_electricity_consumption
REQUESTS_CA_BUNDLE=/etc/ssl/certs/ca-certificates.crt
AWS_CA_BUNDLE=/etc/ssl/certs/ca-certificates.crt
# --- S3 Garage : artefacts MLflow (TP03) ---
# Necessaire pour log_model (upload de l'artefact) ET pour le chargement du modele
# par l'API (download depuis s3://mlflow-artifacts). Cle S3 "mlflow" (RWO sur le bucket).
AWS_ACCESS_KEY_ID=GKxxxxxxxxxxxxxxxxxxxxxxxx
AWS_SECRET_ACCESS_KEY=change-me
MLFLOW_S3_ENDPOINT_URL=https://garage.192-168-122-143.nip.io
# --- Import du package lab ---
PYTHONPATH=/home/user/tp

View File

@@ -70,3 +70,18 @@ MODELLING_FEATURES: dict[ModellingStrategy, list] = {
# Valeurs d'alpha demandees par l'enonce (Partie 3)
RIDGE_ALPHAS = [1, 1e3, 1e9]
# --- TP03 : Model Registry + service de prediction ---
# Nom sous lequel les modeles sont enregistres dans le MLflow Model Registry.
# On garde un nom stable pour retrouver le modele et empiler ses versions.
REGISTERED_MODEL_NAME = "electricity-consumption"
# Alias pointant vers la version promue (chargee par l'API). On identifie le modele
# a servir par son alias (mobile) plutot que par un numero de version (fige).
MODEL_ALIAS = "champion"
# Strategie de features du modele expose par l'API : "full" (les 6 features).
# L'ordre des colonnes servies doit correspondre a celui de l'entrainement.
SERVING_STRATEGY = ModellingStrategy.FULL
SERVING_FEATURES = MODELLING_FEATURES[SERVING_STRATEGY]

View File

@@ -3,6 +3,7 @@ import logging
import mlflow
import pandas as pd
import typer
from mlflow.models import infer_signature
from sklearn import linear_model
from sklearn import metrics
@@ -17,6 +18,11 @@ logger = logging.getLogger(__name__)
@app.command()
def main(
strategy: constants.ModellingStrategy,
register: bool = typer.Option(
False,
"--register/--no-register",
help="Enregistrer le modele dans le Model Registry (cree une nouvelle version).",
),
):
training_file_path = constants.DATASET_DIR / "train.parquet"
validation_file_path = constants.DATASET_DIR / "validation.parquet"
@@ -72,6 +78,24 @@ def main(
):
mlflow.log_metric(f"coef_{feature_name}", float(coefficient))
# Sauvegarde de l'artefact du modele (poids + signature + environnement
# d'execution : requirements.txt, conda.yaml, MLmodel). --register empile
# une nouvelle version dans le Model Registry pour les meilleures experiences.
signature = infer_signature(X_train, train_predictions)
mlflow.sklearn.log_model(
sk_model=model,
name="model",
signature=signature,
input_example=X_train.iloc[:5],
registered_model_name=(
constants.REGISTERED_MODEL_NAME if register else None
),
)
if register:
logger.info(
f"Model registered as '{constants.REGISTERED_MODEL_NAME}' (nouvelle version)"
)
if __name__ == "__main__":
app()

View File

@@ -3,6 +3,7 @@ import logging
import mlflow
import pandas as pd
import typer
from mlflow.models import infer_signature
from sklearn import linear_model
from sklearn import metrics
@@ -17,6 +18,11 @@ logger = logging.getLogger(__name__)
@app.command()
def main(
strategy: constants.ModellingStrategy = constants.ModellingStrategy.MIXED,
register: bool = typer.Option(
False,
"--register/--no-register",
help="Enregistrer chaque modele (par alpha) dans le Model Registry.",
),
):
training_file_path = constants.DATASET_DIR / "train.parquet"
validation_file_path = constants.DATASET_DIR / "validation.parquet"
@@ -74,6 +80,18 @@ def main(
):
mlflow.log_metric(f"coef_{feature_name}", float(coefficient))
# Sauvegarde de l'artefact du modele (voir lab/modeling/cli.py).
signature = infer_signature(X_train, train_predictions)
mlflow.sklearn.log_model(
sk_model=model,
name="model",
signature=signature,
input_example=X_train.iloc[:5],
registered_model_name=(
constants.REGISTERED_MODEL_NAME if register else None
),
)
if __name__ == "__main__":
app()

0
lab/registry/__init__.py Normal file
View File

63
lab/registry/cli.py Normal file
View File

@@ -0,0 +1,63 @@
"""Gestion du MLflow Model Registry : consulter les versions et promouvoir via alias.
La promotion = pointer un alias (mobile) vers une version precise (figee) du modele.
L'API charge ensuite le modele par `models:/<nom>@<alias>` sans connaitre le numero.
"""
import logging
import typer
from mlflow import MlflowClient
from .. import constants
app = typer.Typer(help="MLflow Model Registry (versions + promotion par alias).")
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
@app.command()
def versions(
name: str = constants.REGISTERED_MODEL_NAME,
):
"""Lister les versions du modele enregistre, avec leurs alias."""
client = MlflowClient()
results = client.search_model_versions(f"name = '{name}'")
if not results:
logger.warning(f"Aucune version pour le modele '{name}'.")
raise typer.Exit(code=1)
# Les alias sont portes par le modele enregistre (dict alias -> version),
# pas par les objets renvoyes par search_model_versions.
alias_by_version: dict[str, list[str]] = {}
for alias, version in client.get_registered_model(name).aliases.items():
alias_by_version.setdefault(str(version), []).append(alias)
for mv in sorted(results, key=lambda v: int(v.version)):
aliases = ", ".join(alias_by_version.get(str(mv.version), [])) or "-"
run = client.get_run(mv.run_id) if mv.run_id else None
strategy = run.data.params.get("strategy", "?") if run else "?"
val_rmse = run.data.metrics.get("validation_rmse") if run else None
rmse_txt = f"{val_rmse:.3f}" if val_rmse is not None else "?"
typer.echo(
f"v{mv.version:<3} | alias: {aliases:<12} | strategy={strategy:<12}"
f" | validation_rmse={rmse_txt} | run={mv.run_id}"
)
@app.command()
def promote(
version: int = typer.Option(..., help="Numero de version a promouvoir."),
alias: str = typer.Option(constants.MODEL_ALIAS, help="Alias a (re)pointer."),
name: str = constants.REGISTERED_MODEL_NAME,
):
"""Promouvoir une version : (re)pointer l'alias vers cette version."""
client = MlflowClient()
client.set_registered_model_alias(name=name, alias=alias, version=str(version))
mv = client.get_model_version_by_alias(name=name, alias=alias)
logger.info(f"Alias '{alias}' -> {name} v{mv.version} (source: {mv.source})")
if __name__ == "__main__":
app()