Skip to content
Merged
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
15 changes: 9 additions & 6 deletions datamint/mlflow/tracking/default_experiment.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,5 @@
import os
import sys
import os
from mlflow.tracking.default_experiment.abstract_context import (
DefaultExperimentProvider,
)
Expand All @@ -17,15 +16,19 @@ def in_context(self): # type: ignore[override]
@override
def get_experiment_id(self): # type: ignore[override]
from mlflow.tracking.client import MlflowClient

from datamint.mlflow.tracking.fluent import get_active_project_name

if DatamintExperimentProvider._experiment_id is not None:
return DatamintExperimentProvider._experiment_id
# Get the filename of the main source file
source_code_filename = os.path.basename(sys.argv[0])

# Prefer the active project (set via `datamint.mlflow.set_project()`) as the
# experiment name. Fall back to the main source file's name.
experiment_name = get_active_project_name() or os.path.basename(sys.argv[0])

mlflowclient = MlflowClient()
exp = mlflowclient.get_experiment_by_name(source_code_filename)
exp = mlflowclient.get_experiment_by_name(experiment_name)
if exp is None:
experiment_id = mlflowclient.create_experiment(source_code_filename)
experiment_id = mlflowclient.create_experiment(experiment_name)
else:
experiment_id = exp.experiment_id
DatamintExperimentProvider._experiment_id = experiment_id
Expand Down
43 changes: 32 additions & 11 deletions datamint/mlflow/tracking/fluent.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,13 +15,14 @@
_LOGGER = logging.getLogger(__name__)

_ACTIVE_PROJECT_ID: str | None = None
_ACTIVE_PROJECT_NAME: str | None = None


def get_active_project_id() -> str | None:
"""
Get the active project ID from the environment variable or the global variable.
"""
global _ACTIVE_PROJECT_ID
global _ACTIVE_PROJECT_ID, _ACTIVE_PROJECT_NAME

if _ACTIVE_PROJECT_ID is not None:
return _ACTIVE_PROJECT_ID
Expand All @@ -35,11 +36,22 @@ def get_active_project_id() -> str | None:
project = _find_project_by_name(project_name)
if project is not None:
_ACTIVE_PROJECT_ID = project['id']
_ACTIVE_PROJECT_NAME = project_name
return _ACTIVE_PROJECT_ID

return None


def get_active_project_name() -> str | None:
"""
Get the active project's name, if one was set via `set_project()` or resolved
from `DATAMINT_PROJECT_NAME`/`DATAMINT_PROJECT_ID`.
"""
if _ACTIVE_PROJECT_NAME is None:
get_active_project_id()
return _ACTIVE_PROJECT_NAME


def _find_project_by_name(project_name: str):
dt_client = Api(check_connection=False)
project = dt_client.projects.get_by_name(project_name)
Expand Down Expand Up @@ -68,38 +80,47 @@ def _get_project_by_name_or_id(project_name_or_id: str) -> 'Project':
def set_project(project: 'Project | str'):
"""
Set the active project for the current session.


The project's name also becomes the default MLflow experiment name for
any run started without an explicit experiment (see `DatamintExperimentProvider`).

Args:
project: The Project instance or project name/ID to set as active.
"""
global _ACTIVE_PROJECT_ID
global _ACTIVE_PROJECT_ID, _ACTIVE_PROJECT_NAME

# Ensure MLflow is properly configured before proceeding
ensure_mlflow_configured()

with _PROJECT_LOCK:
if isinstance(project, str):
project_id = None
project = _get_project_by_name_or_id(project)
project_id = project.id
else:
# It's a Project entity
project_id = project.id

_ACTIVE_PROJECT_ID = project_id
_ACTIVE_PROJECT_ID = project.id
_ACTIVE_PROJECT_NAME = project.name
_reset_default_experiment_cache()

# Set 'DATAMINT_PROJECT_ID' environment variable
# so that subprocess can inherit it.
os.environ[EnvVars.DATAMINT_PROJECT_ID.value] = project_id
os.environ[EnvVars.DATAMINT_PROJECT_ID.value] = project.id

return project


def _reset_active_project():
"""Clear the active project, restoring the pre-``set_project()`` state. """
global _ACTIVE_PROJECT_ID
global _ACTIVE_PROJECT_ID, _ACTIVE_PROJECT_NAME

with _PROJECT_LOCK:
_ACTIVE_PROJECT_ID = None
_ACTIVE_PROJECT_NAME = None
_reset_default_experiment_cache()

os.environ.pop(EnvVars.DATAMINT_PROJECT_ID.value, None)


def _reset_default_experiment_cache():
"""Invalidate the cached default-experiment id so a project switch takes effect
on the next run started without an explicit experiment."""
from datamint.mlflow.tracking.default_experiment import DatamintExperimentProvider
DatamintExperimentProvider._experiment_id = None
Loading