multilingual-absa / src /absa /training /mlflow_utils.py
Aryan Mishra
Add CI, typed ORM models, and packaging cleanup
a00fee9
Raw
History Blame Contribute Delete
2.57 kB
from pathlib import Path
from typing import Any, Dict, Optional, Union
import mlflow # type: ignore
# Default configuration
MLFLOW_TRACKING_URI = "sqlite:///mlflow/mlflow.db"
EXPERIMENT_NAME = "multilingual-absa"
def setup_mlflow():
"""Initializes MLflow tracking URI and experiment."""
# Ensure the directory exists
Path("mlflow").mkdir(parents=True, exist_ok=True)
mlflow.set_tracking_uri(MLFLOW_TRACKING_URI)
mlflow.set_experiment(EXPERIMENT_NAME)
def log_training_run(
params: Dict[str, Any],
metrics: Dict[str, float],
model_path: Optional[Union[str, Path]] = None,
run_name: Optional[str] = None,
) -> str:
"""
Logs parameters, metrics, and optionally a model artifact to MLflow.
Args:
params: Dictionary of hyperparameters or configuration.
metrics: Dictionary of evaluation metrics.
model_path: Path to the saved model directory or file.
run_name: Optional name for the run.
Returns:
The ID of the created MLflow run.
"""
setup_mlflow()
with mlflow.start_run(run_name=run_name) as run: # type: ignore[attr-defined]
mlflow.log_params(params) # type: ignore[attr-defined]
mlflow.log_metrics(metrics) # type: ignore[attr-defined]
if model_path:
model_path_obj = Path(model_path)
if model_path_obj.exists():
mlflow.log_artifact(str(model_path_obj), artifact_path="model") # type: ignore[attr-defined]
else:
print(f"Warning: Model path {model_path} does not exist. Artifact not logged.")
return run.info.run_id # type: ignore[no-any-return]
def get_best_run(metric: str = "eval_macro_f1", ascending: bool = False) -> Optional[Any]: # type: ignore
"""
Retrieves the best run from the experiment based on a specific metric.
Args:
metric: The metric to sort by.
ascending: True if a lower metric is better (e.g., loss), False for higher is better (e.g., F1).
Returns:
The MLflow Run object for the best run, or None if no runs exist.
"""
setup_mlflow()
experiment = mlflow.get_experiment_by_name(EXPERIMENT_NAME) # type: ignore[attr-defined]
if not experiment:
return None
runs = mlflow.search_runs( # type: ignore[attr-defined]
experiment_ids=[experiment.experiment_id],
order_by=[f"metrics.{metric} {'ASC' if ascending else 'DESC'}"],
max_results=1,
output_format="list",
)
if not runs:
return None
return runs[0]