TP3 Partie 1 : log_model + Model Registry + promotion par alias
This commit is contained in:
63
lab/registry/cli.py
Normal file
63
lab/registry/cli.py
Normal 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()
|
||||
Reference in New Issue
Block a user