引言
「上周那个效果最好的模型,用的是哪套超参?」——如果这个问题答不上来,说明团队的实验管理已经失控。手写 Excel 记录、文件夹命名 final_v2_真最终版、模型文件散落各处,是机器学习项目最常见的混乱来源。
实验管理(Experiment Tracking) 要解决的就是这三件事:记录(参数、指标、产物)、复现(同样的输入能否得到同样的输出)、对比(哪次实验最好、为什么好)。本文以 MLflow 为主线,串起数据/代码版本化、模型注册表与团队协作规范。
前置:模型上线的完整链路见 https://plumephp.com/ml-model-deployment/;超参搜索与实验管理的配合见 https://plumephp.com/ml-automl-hpo/;流水线组织方式参考 https://plumephp.com/ml-pipelines-feature-selection/。
目录
- 1. 为什么需要实验管理
- 2. MLflow 核心概念与上手
- 3. 参数、指标与产物的记录
- 4. 数据与代码版本化
- 5. 模型注册表与生命周期
- 6. 实验对比与可视化
- 7. 团队协作与实验共享
- 8. 落地规范与常见坑
- 9. 总结
- 延伸阅读
1. 为什么需要实验管理
1.1 没有实验管理的三个后果
| 后果 | 具体表现 |
|---|---|
| 无法复现 | 三个月前的「最佳模型」跑不出来了 |
| 重复劳动 | 两个人做了同一个消融实验 |
| 决策无据 | 说不清为什么选了这个方案 |
1.2 一次实验的完整配方
代码 git commit sha 数据 数据集哈希/快照 ID 环境 依赖版本
超参 学习率/批大小/结构 随机种子 seed 产物 权重/日志/图表/指标
任何一项缺失,实验就不可复现。实验管理工具的作用就是自动把这六项绑在一起。
1.3 工具生态速览
主流选择有四类:MLflow(开源、可自托管、模型注册表完整)、Weights & Biases(SaaS、可视化与协作强)、TensorBoard(本地、只看曲线)、DVC(专注数据与流水线版本化)。本文以 MLflow 为主线,因为它覆盖从跟踪到注册表的完整链路。
2. MLflow 核心概念与上手
2.1 四层概念
Experiment(实验):一个项目/一个任务方向
└── Run(运行) :一次具体的训练
├── Params :输入参数(超参、配置)
├── Metrics :输出指标(可随步数记录)
├── Artifacts:产物文件(模型、图表、日志)
└── Tags :元信息(git sha、作者、数据版本)
2.2 启动与初始化
import mlflow
# 本地默认存到 ./mlruns;生产建议用数据库 + 对象存储
mlflow.set_tracking_uri("http://127.0.0.1:5000")
mlflow.set_experiment("fraud-detection")
# 启动服务:mlflow server --host 0.0.0.0 --port 5000
2.3 最小可运行示例
import mlflow
from sklearn.ensemble import RandomForestClassifier
from sklearn.metrics import f1_score
with mlflow.start_run(run_name="rf-baseline"):
params = {"n_estimators": 200, "max_depth": 8, "random_state": 42}
mlflow.log_params(params)
model = RandomForestClassifier(**params)
model.fit(X_train, y_train)
mlflow.log_metric("f1_macro", f1_score(y_test, model.predict(X_test)))
mlflow.sklearn.log_model(model, "model")
mlflow.set_tag("git_sha", git_sha)
print("run id:", mlflow.active_run().info.run_id)
run_id 是这次实验的唯一身份证,务必记进日志或注释。
2.4 自动记录
不想手写 log_params,可以开启自动记录:
mlflow.sklearn.autolog() # 自动记录超参、指标、模型
mlflow.pytorch.autolog() # PyTorch 训练循环自动记录
mlflow.transformers.autolog() # HuggingFace 微调自动记录
自动记录覆盖 90% 的常规场景,自定义指标仍需手动 log。
3. 参数、指标与产物的记录
3.1 什么该记成参数,什么该记成指标
| 类型 | 特征 | 例子 |
|---|---|---|
| 参数 | 训练前确定、不变 | 学习率、层数、特征版本 |
| 指标 | 训练中/后产生、可比较 | loss、F1、AUC、耗时 |
| 产物 | 文件 | 权重、混淆矩阵图、预测结果 |
| 标签 | 字符串元信息 | git sha、数据快照、负责人 |
3.2 按步记录训练曲线
for epoch in range(epochs):
train_loss = train_one_epoch(model, loader)
val_loss, val_acc = evaluate(model, val_loader)
mlflow.log_metrics({"train_loss": train_loss,
"val_loss": val_loss,
"val_acc": val_acc}, step=epoch)
UI 上会自动画出折线图,一眼看出过拟合的拐点。
3.3 记录产物与图片
import matplotlib.pyplot as plt
fig, ax = plt.subplots(); ax.plot(history["val_loss"])
mlflow.log_figure(fig, "curves/val_loss.png") # 直接 log figure 对象
mlflow.log_artifact("val_loss.png") # 单个文件
mlflow.log_artifacts("reports/", "reports") # 整个目录
3.4 批量记录与结果分析
# 用父 run 包住多个子 run,适合「一次调参扫一批」
with mlflow.start_run(run_name="hpo-sweep") as parent:
for lr in [1e-3, 1e-4, 1e-5]:
with mlflow.start_run(run_name=f"lr-{lr}", nested=True):
mlflow.log_param("lr", lr); mlflow.log_metric("f1", train_and_eval(lr))
# 用 Pandas 拉回结果做分析
df = mlflow.search_runs(experiment_names=["fraud-detection"])
cols = ["params.n_estimators", "params.max_depth",
"metrics.f1_macro", "tags.git_sha"]
print(df[cols].sort_values("metrics.f1_macro", ascending=False).head())
嵌套运行让「一组实验」在 UI 上折叠展示,避免列表被淹没;search_runs 返回标准 DataFrame,可直接做透视表、画散点。
4. 数据与代码版本化
4.1 三层版本化
代码 → git(commit sha 记进 run 的 tag)
数据 → 快照 / 哈希 / DVC
环境 → 依赖锁文件 / 容器镜像
4.2 代码版本:记 sha 而不是分支名
import subprocess
def git_sha():
return subprocess.check_output(["git", "rev-parse", "HEAD"]).decode().strip()
mlflow.set_tag("git_sha", git_sha())
mlflow.set_tag("git_dirty", bool(subprocess.check_output(
["git", "status", "--porcelain"]).strip()))
分支名会移动,sha 不会。git_dirty 标记「代码没提交就跑实验」,这类结果不可复现。
4.3 数据版本:哈希与 DVC
import hashlib
def data_fingerprint(path):
"""对数据内容做哈希,作为数据版本号"""
h = hashlib.sha256()
with open(path, "rb") as f:
for chunk in iter(lambda: f.read(1 << 20), b""):
h.update(chunk)
return h.hexdigest()[:16]
mlflow.set_tag("data_version", data_fingerprint("data/train.parquet"))
大文件交给 DVC,它生成体积很小的 .dvc 指针文件入 git,dvc.lock 记录每一步的输入哈希与输出哈希:
dvc init
dvc remote add -d storage s3://my-bucket/dvc
dvc add data/train.parquet
git add data/train.parquet.dvc data/.gitignore && git commit -m "track data v3"
git checkout <sha> && dvc checkout # 复现历史某次实验
4.4 环境版本
mlflow.log_artifact("requirements.txt")
mlflow.set_tag("python_version", sys.version.split()[0])
mlflow.set_tag("image", "registry.example.com/train:cuda12.1-v3")
容器镜像 tag 比 pip freeze 更可靠——它连 CUDA、系统库都锁定了。
5. 模型注册表与生命周期
5.1 为什么需要注册表
Tracking 记录的是「所有实验」,注册表管理的是「要上线的模型」。前者是流水账,后者是带审批的发布通道。
5.2 注册模型与阶段流转
import mlflow
with mlflow.start_run() as run:
mlflow.sklearn.log_model(model, "model", registered_model_name="fraud-detector")
client = mlflow.MlflowClient()
client.transition_model_version_stage(
name="fraud-detector", version="3", stage="Staging")
# 验证通过后转生产,旧的自动归档
client.transition_model_version_stage(
name="fraud-detector", version="3", stage="Production",
archive_existing_versions=True)
阶段流转路径是 None → Staging → Production → Archived。
5.3 标签、备注与加载
client.set_model_version_tag("fraud-detector", "3", "validation_f1", "0.87")
client.update_model_version(name="fraud-detector", version="3",
description="改用 LightGBM + 新特征 v4,AUC 提升 1.2 个点")
# 线上服务加载 Production 阶段,而不是硬编码某个版本号
model = mlflow.pyfunc.load_model("models:/fraud-detector/Production")
old = mlflow.pyfunc.load_model("models:/fraud-detector/2") # 回滚验证
服务端只认阶段,不认版本号——这样切换版本无需改代码,回滚只是再切一次阶段。
5.4 与 CI/CD 集成
训练完成 → 自动注册为 Staging → 跑验证集与偏置检查
→ 通过则转 Production → 触发部署流水线 → 监控告警
→ 指标劣化则回退到上一 Production 版本
6. 实验对比与可视化
6.1 UI 与代码里的对比
MLflow UI 支持勾选多个 run 做平行坐标图(Parallel Coordinates),直观看出哪组超参把指标推高了。代码里则可以按超参做透视:
runs = df[["params.max_depth", "params.n_estimators", "metrics.f1_macro"]].dropna()
runs["params.max_depth"] = runs["params.max_depth"].astype(int)
print(runs.pivot_table(index="params.max_depth",
columns="params.n_estimators",
values="metrics.f1_macro", aggfunc="max").round(4))
6.2 找最佳 run 与对比纪律
best = df.sort_values("metrics.f1_macro", ascending=False).iloc[0]
print("最佳 run:", best["run_id"])
print("参数:", {k: v for k, v in best.items() if k.startswith("params.")})
让对比有意义的三条纪律:
| 纪律 | 说明 |
|---|---|
| 一次只改一个变量 | 否则说不清是哪个改动起效 |
| 固定随机种子 | 否则指标波动被误认为提升 |
| 记录完整指标 | 只记 F1 无法解释「为什么」 |
7. 团队协作与实验共享
7.1 集中式 Tracking Server
mlflow server \
--backend-store-uri postgresql://user:pwd@db:5432/mlflow \
--default-artifact-root s3://my-bucket/mlflow \
--host 0.0.0.0 --port 5000
元数据(参数指标)进关系库,大文件(模型权重)进对象存储——这是标准生产部署。
7.2 命名与标签规范
实验名 项目-任务(fraud-detection) run 名 模型-关键变量(lgbm-lr1e-3)
标签 owner / team / data_version / git_sha / stage
统一命名让 search_runs 的过滤条件能写出花样:
mlflow.search_runs(filter_string="tags.owner = 'alice' and metrics.f1_macro > 0.8")
7.3 工具选型与复现演练
Weights & Biases 在实时可视化与协作上体验更好:训练曲线实时推送、支持富文本报告、能对比不同人不同机器的 run。它的 SaaS 形态开箱即用,代价是数据要出内网。选型建议:数据敏感需自托管选 MLflow,追求体验快速起步选 W&B,只做本地小项目用 TensorBoard + 文件名规范也够。
复现一次历史实验要走通四步:
git checkout <sha> # 1. 切到记录的代码版本
dvc checkout # 2. 恢复数据版本
mlflow artifacts download -r <run_id> -a config.yaml # 3. 拉取超参配置
docker run --gpus all train:cuda12.1-v3 python train.py --config config.yaml
四步都能走通,才算真正可复现。
8. 落地规范与常见坑
| 现象 | 根因 | 处理 |
|---|---|---|
| run 列表几百条找不到重点 | 无命名规范 | 统一实验名与 run 名,打 owner 标签 |
| 复现结果对不上 | 随机种子未固定 | 固定 seed 并记进参数 |
| 指标忽高忽低 | 数据切分不同 | 记录 data_version,冻结切分 |
| 磁盘被产物撑爆 | 每个 run 都存全量权重 | 只对候选模型存产物,其余存指标 |
| UI 打开极慢 | 元数据存本地文件 | 换 Postgres 后端 |
| 分不清线上版本 | 直接按 run 部署 | 走模型注册表阶段流转 |
8.1 最小可行规范
1. 每次训练必须开 run,禁止裸跑
2. 必记 git_sha、data_version、seed、owner
3. 模型产物只存到注册表候选,其余 run 不存
4. 上线只从 Production 阶段拉取,结论必附 run_id
8.2 常见误区
- 把 MLflow 当模型服务器:它管元数据和产物,不做高并发推理,线上还是 vLLM/Triton/FastAPI;
- 只记指标不记参数:事后无法解释差异来源;
- 用文件夹名当版本号:
v2_final_new这类命名没有任何机器可读性。
9. 总结
9.1 从混乱到有序的路线
阶段一:开 run + 记参数指标(当天可落地)
阶段二:加 git_sha / data_version / seed 标签
阶段三:上集中式 server + 对象存储
阶段四:引入模型注册表与阶段流转
阶段五:接入 CI/CD,自动验证与回滚
9.2 关键决策点
| 问题 | 选择 |
|---|---|
| 数据不能出内网 | MLflow 自托管 |
| 只想快速起步 | W&B SaaS 或本地 MLflow |
| 数据文件巨大 | DVC + 对象存储,不入 git |
| 多人协作 | 集中式 server + 命名规范 |
| 需要审批与回滚 | 模型注册表阶段流转 |
9.3 一句话心法
实验管理的价值不在「记录」,而在「让决策可回溯」——三个月后还能说清「为什么选它」,才算真正做对了。
延伸阅读
- https://plumephp.com/ml-model-deployment/ — 模型上线、版本与回滚的完整链路
- https://plumephp.com/ml-automl-hpo/ — 超参搜索与实验管理的配合
- https://plumephp.com/ml-pipelines-feature-selection/ — Pipeline 与流水线组织
- https://plumephp.com/ml-model-evaluation/ — 评估指标与交叉验证设计
- MLflow 官方文档
继续阅读
探索更多技术文章
浏览归档,发现更多关于系统设计、工具链和工程实践的内容。