MLflow 实验跟踪(MLflow Tracking)是 MLflow 平台的核心功能之一,用于记录、组织和管理机器学习模型开发过程中的所有关键信息,帮助数据科学家和工程师高效追踪实验、对比结果、复现模型,并避免“实验混乱”的问题。
在机器学习开发中,你可能会遇到这些场景:
MLflow 实验跟踪就是为了解决这些问题——它像一个“实验笔记本”,自动或手动记录实验的输入、过程和输出,让整个开发流程可追溯、可对比、可复现。
MLflow 定义了几个核心概念,覆盖实验的全生命周期:
| 概念 | 说明 |
|---|---|
| 实验(Experiment) | 一组相关实验的集合(比如“信用卡欺诈检测模型实验”“图像分类模型优化实验”)。每个实验有唯一 ID 和名称,方便归类管理。 |
| 运行(Run) | 一次具体的实验执行(比如一次模型训练)。每个 Run 会记录所有关键信息,是实验跟踪的基本单元。 |
| 参数(Parameters) | 实验的输入配置,比如超参数(学习率 lr=0.01、批量大小 batch_size=32)、数据路径、模型类型等。 |
| 指标(Metrics) | 实验的量化结果,比如准确率 accuracy=0.95、损失值 loss=0.12、AUC 等。支持随时间记录(比如每个 epoch 的损失变化)。 |
| artifact(工件) | 实验输出的文件,比如模型文件(.pkl、.pt)、日志、可视化图表(如损失曲线图)、数据样本等。MLflow 会持久化存储这些文件。 |
| 标签(Tags) | 给 Run 添加的描述性信息,比如 version=v1.2、creator=张三、status=best,方便筛选和搜索。 |
| 源代码版本 | 自动记录实验对应的 Git Commit Hash(如果当前目录是 Git 仓库),确保代码可追溯。 |
| 环境信息 | 可选记录运行环境(如 Python 版本、依赖库版本 requirements.txt),辅助复现。 |
MLflow 提供了 Python API(最常用)、CLI 和 UI 三种交互方式,核心流程如下:
pip install mlflow在训练代码中,通过 mlflow 模块记录实验信息:
import mlflow
import mlflow.sklearn
from sklearn.ensemble import RandomForestClassifier
from sklearn.datasets import load_iris
from sklearn.model_selection import train_test_split
from sklearn.metrics import accuracy_score
# 1. 加载数据
iris = load_iris()
X, y = iris.data, iris.target
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2)
# 2. 设置实验名称(如果不存在会自动创建)
mlflow.set_experiment("iris-classification-experiment")
# 3. 开始一次实验运行(Run)
with mlflow.start_run(run_name="rf-baseline"):
# 记录参数
params = {
"n_estimators": 100,
"max_depth": 5,
"random_state": 42
}
mlflow.log_params(params) # 批量记录参数
# 训练模型
model = RandomForestClassifier(**params)
model.fit(X_train, y_train)
# 记录指标
y_pred = model.predict(X_test)
acc = accuracy_score(y_test, y_pred)
mlflow.log_metric("accuracy", acc) # 记录单个指标
# 也可以记录多个指标:mlflow.log_metrics({"accuracy": acc, "f1": f1_score})
# 记录模型(作为 artifact)
mlflow.sklearn.log_model(model, "random-forest-model")
# 添加标签
mlflow.set_tag("model_type", "RandomForest")
mlflow.set_tag("dataset", "iris")运行代码后,在终端启动 MLflow UI:
mlflow ui然后在浏览器访问 http://localhost:5000,你会看到:
log_* 方法。例如:mlflow.sklearn.autolog() # 自动记录 sklearn 模型的参数、指标、模型mlruns 目录。也可以配置远程跟踪服务器(如用 MySQL 存储元数据,S3/OSS 存储工件),方便团队协作:mlflow.set_tracking_uri("http://remote-mlflow-server:5000") # 设置远程服务器地址metrics.accuracy > 0.9 找准确率超过 90% 的实验。MLflow 实验跟踪是机器学习工程化的基础工具,它将“混乱的实验过程”转化为“结构化的可管理数据”,帮助开发者和团队更高效地迭代模型、降低协作成本,是 MLOps 流程中不可或缺的一环。