| 概念 | 说明 |
|---|---|
| Experiment | 一组相关实验的集合 |
| Run | 一次训练或实验执行 |
| Parameter | 超参数(如 learning_rate) |
| Metric | 数值指标(如 accuracy、loss) |
| Artifact | 模型文件、图片、日志、CSV 等 |
| Tracking URI | MLflow 记录数据的地址(本地 / 远程) |
pip install mlflow如果你使用 sklearn / PyTorch / TensorFlow,可一并安装:
pip install mlflow scikit-learnmlflow ui访问:
http://127.0.0.1:5000import 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
# 设置实验
mlflow.set_experiment("iris_rf_experiment")
with mlflow.start_run():
# 参数
params = {
"n_estimators": 100,
"max_depth": 5
}
mlflow.log_params(params)
# 数据
X, y = load_iris(return_X_y=True)
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2)
# 模型
model = RandomForestClassifier(**params)
model.fit(X_train, y_train)
# 指标
acc = accuracy_score(y_test, model.predict(X_test))
mlflow.log_metric("accuracy", acc)
# 记录模型
mlflow.sklearn.log_model(model, "model")
print(f"Accuracy: {acc}")mlflow.set_tracking_uri("file:///path/to/mlruns")mlflow.set_tracking_uri("http://mlflow-server:5000")export MLFLOW_TRACKING_URI=http://mlflow-server:5000mlflow.set_experiment("recommendation_system")with mlflow.start_run(run_name="rf_100_trees"):
...mlflow.log_artifact("confusion_matrix.png")
mlflow.log_dict({"features": ["a", "b"]}, "features.json")import mlflow.pytorch
mlflow.pytorch.log_model(model, "model")import mlflow.tensorflow
mlflow.tensorflow.log_model(model, "model")for lr in [0.01, 0.001]:
with mlflow.start_run():
mlflow.log_param("lr", lr)
...训练脚本
↓
MLflow Tracking Server
↓
后端存储(MySQL / PostgreSQL)
↓
制品存储(S3 / MinIO)mlflow server \
--backend-store-uri mysql+pymysql://user:pass@host/mlflow \
--default-artifact-root s3://mlflow-bucket/ \
--host 0.0.0.0 \
--port 5000mlflow.start_run()MLflow 实验跟踪 = 用start_run()包裹训练代码 + 记录参数 / 指标 / 模型 + 统一 Tracking URI
如果你愿意,我可以:
你现在的场景是 个人实验、团队使用,还是生产部署?