Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
22 changes: 13 additions & 9 deletions mmengine/visualization/vis_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@
import warnings
from abc import ABCMeta, abstractmethod
from collections.abc import MutableMapping
from pathlib import Path
from typing import Any, Callable, List, Optional, Sequence, Union

import cv2
Expand Down Expand Up @@ -662,16 +663,18 @@ class MLflowVisBackend(BaseVisBackend):
Defaults to None.
params (dict, optional): The params to be added to the experiment.
Defaults to None.
tracking_uri (str, optional): The tracking uri. Defaults to None.
tracking_uri (str, optional): The tracking uri. If None, a local
SQLite database is created in ``save_dir``. Defaults to None.
artifact_suffix (Tuple[str] or str, optional): The artifact suffix.
Defaults to ('.json', '.log', '.py', 'yaml').
tracked_config_keys (dict, optional): The top level keys of config that
will be added to the experiment. If it is None, which means all
the config will be added. Defaults to None.
`New in version 0.7.4.`
artifact_location (str, optional): The location to store run artifacts.
If None, the server picks an appropriate default.
Defaults to None.
If None, local artifacts are stored in ``save_dir/artifacts`` when
using the default SQLite database. Otherwise, the tracking server
picks an appropriate default. Defaults to None.
`New in version 0.10.4.`
"""

Expand Down Expand Up @@ -718,22 +721,23 @@ def _init_env(self):
if handler.stream is None or handler.stream.closed:
handler.stream = open(handler.baseFilename, 'a')

artifact_location = self._artifact_location
if self._tracking_uri is not None:
logger.warning(
'Please make sure that the mlflow server is running.')
self._mlflow.set_tracking_uri(self._tracking_uri)
else:
if os.name == 'nt':
file_url = f'file:\\{os.path.abspath(self._save_dir)}'
else:
file_url = f'file://{os.path.abspath(self._save_dir)}'
self._mlflow.set_tracking_uri(file_url)
save_dir = Path(self._save_dir).resolve()
database_path = (save_dir / 'mlflow.db').as_posix()
self._mlflow.set_tracking_uri(f'sqlite:///{database_path}')
if artifact_location is None:
artifact_location = (save_dir / 'artifacts').as_uri()

self._exp_name = self._exp_name or 'Default'

if self._mlflow.get_experiment_by_name(self._exp_name) is None:
self._mlflow.create_experiment(
self._exp_name, artifact_location=self._artifact_location)
self._exp_name, artifact_location=artifact_location)

self._mlflow.set_experiment(self._exp_name)

Expand Down
9 changes: 9 additions & 0 deletions tests/test_visualizer/test_vis_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -274,6 +274,12 @@ def test_define_metric_cfg(self):

class TestMLflowVisBackend:

@pytest.fixture(autouse=True)
def use_null_pool(self, monkeypatch):
# Release SQLite handles so the test directory can be removed
# on Windows.
monkeypatch.setenv('MLFLOW_SQLALCHEMYSTORE_POOLCLASS', 'NullPool')

def test_init(self):
MLflowVisBackend('temp_dir')
VISBACKENDS.build(dict(type='MLflowVisBackend', save_dir='temp_dir'))
Expand Down Expand Up @@ -318,6 +324,9 @@ def test_close(self):
cfg = Config(dict(work_dir='temp_dir'))
mlflow_vis_backend = MLflowVisBackend('temp_dir')
mlflow_vis_backend._init_env()
assert mlflow_vis_backend._mlflow.get_tracking_uri().startswith(
'sqlite:///')
assert os.path.isfile(os.path.join('temp_dir', 'mlflow.db'))
mlflow_vis_backend.add_config(cfg)
mlflow_vis_backend.close()
shutil.rmtree('temp_dir')
Expand Down