MLflow Tracking 主要跟踪四类东西:
模型超参、配置项(如 learning_rate)
训练/验证过程中的数值(如 accuracy、loss)
模型文件、图片、日志、配置文件
pip install mlflowimport mlflow
import mlflow.sklearn
from sklearn.ensemble import RandomForestClassifier
from sklearn.datasets import load_iris
from sklearn.model_selection import train_test_split
mlflow.set_experiment("iris_classification")
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)
with mlflow.start_run():
n_estimators = 100
max_depth = 5
mlflow.log_param("n_estimators", n_estimators)
mlflow.log_param("max_depth", max_depth)
model = RandomForestClassifier(n_estimators=n_estimators, max_depth=max_depth)
model.fit(X_train, y_train)
acc = model.score(X_test, y_test)
mlflow.log_metric("accuracy", acc)
mlflow.sklearn.log_model(model, "model")mlruns/。mlflow ui浏览器打开:
http://localhost:5000你可以看到:
for epoch in range(10):
loss = train_one_epoch()
mlflow.log_metric("loss", loss, step=epoch)import matplotlib.pyplot as plt
plt.plot([1,2,3], [1,4,9])
plt.savefig("curve.png")
mlflow.log_artifact("curve.png")mlflow.set_tag("model_type", "random_forest")
mlflow.set_tag("author", "alice")with mlflow.start_run(run_name="parent"):
for lr in [0.01, 0.1]:
with mlflow.start_run(run_name=f"lr={lr}", nested=True):
mlflow.log_param("lr", lr)mlflow server \
--host 0.0.0.0 \
--port 5000 \
--backend-store-uri sqlite:///mlflow.db \
--default-artifact-root s3://my-mlflow-bucketmlflow.set_tracking_uri("http://server-ip:5000")
mlflow.set_experiment("team_exp")✅ 每个实验一个 set_experiment
✅ 参数用 log_param,不要写死在代码里
✅ 模型一定要 log_model
✅ 对比实验时多看 metric + tag
| 场景 | 工具 |
|---|---|
| 本地实验 | MLflow + sklearn |
| 深度学习 | MLflow + PyTorch / TensorFlow |
| 自动调参 | MLflow + Optuna / Ray Tune |
| 模型部署 | MLflow Model Registry |
如果你愿意,我可以:
你现在是个人用还是团队用?