You can not select more than 25 topics
Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
72 lines
2.6 KiB
72 lines
2.6 KiB
2 years ago
|
# Ultralytics YOLO 🚀, AGPL-3.0 license
|
||
2 years ago
|
|
||
|
import os
|
||
|
import re
|
||
|
from pathlib import Path
|
||
|
|
||
2 years ago
|
from ultralytics.utils import LOGGER, TESTS_RUNNING, colorstr
|
||
2 years ago
|
|
||
|
try:
|
||
|
import mlflow
|
||
|
|
||
|
assert not TESTS_RUNNING # do not log pytest
|
||
|
assert hasattr(mlflow, '__version__') # verify package is not directory
|
||
|
except (ImportError, AssertionError):
|
||
|
mlflow = None
|
||
|
|
||
|
|
||
|
def on_pretrain_routine_end(trainer):
|
||
2 years ago
|
"""Logs training parameters to MLflow."""
|
||
2 years ago
|
global mlflow, run, run_id, experiment_name
|
||
|
|
||
|
if os.environ.get('MLFLOW_TRACKING_URI') is None:
|
||
|
mlflow = None
|
||
|
|
||
|
if mlflow:
|
||
|
mlflow_location = os.environ['MLFLOW_TRACKING_URI'] # "http://192.168.xxx.xxx:5000"
|
||
|
mlflow.set_tracking_uri(mlflow_location)
|
||
|
|
||
2 years ago
|
experiment_name = os.environ.get('MLFLOW_EXPERIMENT_NAME') or trainer.args.project or '/Shared/YOLOv8'
|
||
2 years ago
|
run_name = os.environ.get('MLFLOW_RUN') or trainer.args.name
|
||
2 years ago
|
experiment = mlflow.get_experiment_by_name(experiment_name)
|
||
|
if experiment is None:
|
||
|
mlflow.create_experiment(experiment_name)
|
||
|
mlflow.set_experiment(experiment_name)
|
||
|
|
||
|
prefix = colorstr('MLFlow: ')
|
||
|
try:
|
||
2 years ago
|
run, active_run = mlflow, mlflow.active_run()
|
||
|
if not active_run:
|
||
2 years ago
|
active_run = mlflow.start_run(experiment_id=experiment.experiment_id, run_name=run_name)
|
||
2 years ago
|
run_id = active_run.info.run_id
|
||
|
LOGGER.info(f'{prefix}Using run_id({run_id}) at {mlflow_location}')
|
||
|
run.log_params(vars(trainer.model.args))
|
||
2 years ago
|
except Exception as err:
|
||
|
LOGGER.error(f'{prefix}Failing init - {repr(err)}')
|
||
|
LOGGER.warning(f'{prefix}Continuing without Mlflow')
|
||
|
|
||
|
|
||
|
def on_fit_epoch_end(trainer):
|
||
2 years ago
|
"""Logs training metrics to Mlflow."""
|
||
2 years ago
|
if mlflow:
|
||
|
metrics_dict = {f"{re.sub('[()]', '', k)}": float(v) for k, v in trainer.metrics.items()}
|
||
|
run.log_metrics(metrics=metrics_dict, step=trainer.epoch)
|
||
|
|
||
|
|
||
|
def on_train_end(trainer):
|
||
2 years ago
|
"""Called at end of train loop to log model artifact info."""
|
||
2 years ago
|
if mlflow:
|
||
|
root_dir = Path(__file__).resolve().parents[3]
|
||
2 years ago
|
run.log_artifact(trainer.last)
|
||
2 years ago
|
run.log_artifact(trainer.best)
|
||
|
run.pyfunc.log_model(artifact_path=experiment_name,
|
||
|
code_path=[str(root_dir)],
|
||
|
artifacts={'model_path': str(trainer.save_dir)},
|
||
|
python_model=run.pyfunc.PythonModel())
|
||
|
|
||
|
|
||
|
callbacks = {
|
||
|
'on_pretrain_routine_end': on_pretrain_routine_end,
|
||
|
'on_fit_epoch_end': on_fit_epoch_end,
|
||
|
'on_train_end': on_train_end} if mlflow else {}
|