这是 AI 实操第 3 天。前两篇做了 PyTorch 训练和 ONNX 导出,手上有好几个模型文件了。哪个准确率最高?当时用的什么参数?过一周还记不记得?这就是 MLflow 要解决的问题。这篇记录我实际操作 MLflow 的过程和踩的坑,读者水平跟我差不多——懂点 Python,刚学 PyTorch。 本文首发地址 https://h89.cn/archives/669.html 项目地址 https://gitee.com/chenjim/cockpit-ai-from-zero

为什么需要实验追踪

先说说痛点。调参的时候经常要在命令行反复跑:学习率设 0.001 还是 0.01?batch size 用 64 还是 128?靠脑子记结果,很快就不行了。

比如跑完一轮准确率 98.9%,但当时用的什么学习率、什么 batch size,过两天就忘。想对比一下学习率 0.001 和 0.01 的效果,得翻终端历史记录,对着屏幕比数字,麻烦得很。模型文件也是乱糟糟的,.pth 文件名随便起,根本分不清哪个是哪个。要是两个人一起调,还得把参数和结果整理成表格发来发去。

MLflow 就是来解决这些问题的。它是一个开源实验管理工具,Databricks 公司出的。今天只用到它的实验追踪和模型注册功能,但已经能解决上面大部分问题了。

MLflow 的四个核心概念

实验(Experiment)

实验是一个项目级别的容器,可以理解成"一次调参活动"。比如你准备针对 MNIST 调几组学习率,就可以创建一个名为 "mnist-classification" 的实验。这个实验下面会包含多次运行的记录。

运行(Run)

每次跑训练脚本就是一次运行。每次运行会记录三样东西:参数(params)、指标(metrics)、制品(artifacts)。参数是你设定的超参数,指标是训练过程中产生的数值(如 loss、accuracy),制品是产出的文件(如模型权重、可视化图片)。

Model Registry(模型注册表)

当实验跑完,找到效果最好的模型,可以把它注册到 Model Registry。注册后可以给模型标注阶段:Staging(待验证)、Production(已上线)、Archived(已归档)。这样模型的生命周期就清晰了。

MLflow Tracking Server

这是 MLflow 的后端服务,负责接收训练脚本发来的记录,存入数据库并提供 Web 界面。数据默认存在当前目录的 mlruns/ 文件夹里,跑完 mlflow ui 就能看到。

下面这张图展示了整个流程:

MLflow 实验追踪工作流程

安装和基础用法

项目统一用虚拟环境管理依赖,激活后安装:

source .venv/bin/activate
pip install mlflow

(本项目 requirements.txt 已包含 mlflow,直接 pip install -r requirements.txt 也可)

安装后启动 Web 界面:

cd day03-H-mlflow
MLFLOW_SERVER_DISABLE_SECURITY_MIDDLEWARE=true mlflow ui --host 0.0.0.0 --port 17892

浏览器打开 http://localhost:17892 就能看到 MLflow 界面。刚开始是空的,需要先跑训练脚本把数据写进去。 MLflow 界面

给训练脚本加上 MLflow 追踪

接下来实操。我拿之前的 MNIST 训练脚本改的,加了 MLflow 追踪。核心改动不多,先交代总体结构。

脚本做的事情:先创建或切到一个实验,然后循环跑 3 组不同超参数的训练——每次训练算一次独立的运行,最后找出准确率最高的模型注册到 Model Registry。

下面分段看看关键代码。完整代码在 day03-H-mlflow/train_with_mlflow.py,可以直接运行。

创建实验

experiment_id = mlflow.set_experiment("mnist-classification").experiment_id

这行代码会创建名为 "mnist-classification" 的实验。如果之前已经创建过,它会自动切换到已有实验,不会重复创建。返回的 experiment_id 是一个整数,内部用来标识这个实验。

记录一次运行

每跑一组参数,就开启一次运行:

with mlflow.start_run(experiment_id=experiment_id, run_name=f"lr={lr}_bs={batch_size}"):
    mlflow.log_param("learning_rate", lr)
    mlflow.log_param("batch_size", batch_size)
    mlflow.log_param("epochs", epochs)
    # ... 训练代码 ...
    mlflow.log_metric("train_loss", avg_loss, step=epoch)
    mlflow.log_metric("val_accuracy", acc, step=epoch)
    # ... 模型保存 ...
    mlflow.pytorch.log_model(model, name="mnist_cnn_model",
                             serialization_format="pickle", input_example=example_input)

start_run 可以用 with 语句,也可以手动调用 start 和 end。在 with 块内做的所有记录都会归入本次运行。

记录的内容分三类:

  • 参数(params):调用 log_param,记录训练前设定的值,比如学习率、batch size。这些值在同一个实验的不同运行之间用来做对比。
  • 指标(metrics):调用 log_metric,记录训练过程中产生的数值,比如 loss 和 accuracy。可以多次调用同一个指标名,用 step 区分不同 epoch。
  • 制品(artifacts):调用 log_model(或其他 log_artifact 相关方法),保存模型文件。

注册最佳模型

3 组实验跑完后,用 search_runs 按准确率降序排序,把最好的那条记录注册到 Model Registry:

best_run = mlflow.search_runs(
    experiment_ids=[experiment_id],
    order_by=["metrics.val_accuracy DESC"],
    max_results=1
).iloc[0]
best_run_id = best_run.run_id
model_uri = f"runs:/{best_run_id}/mnist_cnn_model"
mlflow.register_model(model_uri, "MNISTCNN")

注册成功后,在 MLflow UI 的 Models 标签页就能看到这个模型,可以手动修改它的阶段(Staging / Production / Archived)。

跑 3 组实验的结果对比

我用 3 组不同的超参数分别跑了 5 个 epoch:

三组超参数实验准确率对比

从结果可以看出两个有意思的现象:

一是学习率不是越大越好。0.01 的学习率(实验 B)反而比 0.001 差,准确率低了 1 个多百分点。这说明学习率太大时,优化器容易"跳过"最优值,一直在最优点附近震荡。

二是 batch size 的影响没那么大。128(实验 C)虽然每次更新参数用的样本更多,但在 MNIST 这个简单任务上,和 64 的差距很小。

最终准确率最高的是实验 C,98.82%,被注册到了 Model Registry。

在 WebUI 中对比实验

启动 mlflow ui 之后,浏览器打开就能看到刚才跑的结果。左侧选实验,右侧列出所有运行记录,表格里直接能看到每次运行的参数和指标,点列名就能排序。想对比某几组实验,勾选对应的行再点 Compare,上面是参数对比表,下面是指标对比图。

Models 标签页能看到注册表里的模型,点进去可以看版本列表和当前阶段,支持手动切换 Staging / Production / Archived。 Models 标签页

踩坑记录

坑一:Python 文件命名

MLflow 在保存 PyTorch 模型时,会用 pickle 把模型类和代码路径一起存起来。如果你的代码在名字带连字符的目录里(比如 day03-H-mlflow),用 mlflow.pyfunc.load_model 加载模型时会报错。这个问题的根源是 Python 的 import 机制不允许目录名有连字符。

我目前的解决方法分两部分:

保存模型时用 serialization_format="pickle" 并传 input_example

加载模型时直接用训练脚本中的模型类定义配合 state_dict,或者用 MLflow 推荐的 mlflow.pyfunc.load_model 做推理(前提是代码所在的目录可以被正常 import)。

坑二:MLflow 版本差异

我装的 MLflow 是 3.14.0 版本,这个版本对 PyTorch 模型的默认保存格式是 "pt2"(TorchScript 的图格式),需要在 log_model 时额外传 input_exampleserialization_format。老版本的 MLflow 默认用 pickle 格式,不需要传这两个参数。

坑三:小数据集上实验差异不大

MNIST 是个很简单的数据集,准确率普遍在 98% 以上,不同超参数之间的差异很小。如果想看到更明显的对比效果,可以用稍微复杂一点的数据集或者减少 epoch 数。

坑四:局域网访问被安全中间件拦截

MLflow 3.14 自带安全中间件,默认只允许 localhost 访问。如果用 --host 0.0.0.0 允许局域网访问,浏览器打开页面时,AJAX 请求会被 CORS 中间件拦住,数据死活刷不出来。

一开始试了 --cors-allowed-origins '*'MLFLOW_CORS_ALLOWED_ORIGINS=* 都不行,原因是:

  1. 多 worker 模式下 CLI 参数没有正确传给子进程
  2. 环境变量名应该是 MLFLOW_SERVER_CORS_ALLOWED_ORIGINS,不是 MLFLOW_CORS_ALLOWED_ORIGINS
  3. 用通配符 * 会导致浏览器拒绝带 cookie 的请求

开发环境最简单的解决方法是关掉安全中间件:

MLFLOW_SERVER_DISABLE_SECURITY_MIDDLEWARE=true mlflow ui --host 0.0.0.0 --port 17892

总结

MLflow 不是非用不可,但用了之后训练记录确实好查多了。一个人调参省了自己记笔记的时间,几个人一起干活也不用互相发 Excel 了。

这次只用了 Tracking 和 Registry,它还有 Serving、Evaluation、Projects 等组件,以后慢慢试。不过工具归工具,核心还是得养成每次实验都记录的习惯。

下一天预告

到今天为止,PyTorch 训练、ONNX 导出、MLflow 实验追踪三条主线都跑通了。第 4 天开始走上部署路线:用阿里开源的 MNN 在 CPU 上跑端侧推理。MLflow 后续还可以试试这些进阶玩法:在训练脚本里自动导出 ONNX 作为 artifact 存进实验、设置远程 Tracking Server(用 PostgreSQL 做后端)多人协作、用 mlflow.pyfunc.load_model 加载注册好的模型写推理服务。


系列文章

本系列「车载端侧 AI 工程化从零上手」共 10 篇,建议按序阅读:

  1. 零基础用 PyTorch 识别手写数字:MNIST 实战入门
  2. 把 PyTorch 模型变成 ONNX:导出、验证和可视化一次学会
  3. MLflow 实验追踪:从入门到上手
  4. 用 MNN 在 CPU 上跑 AI 推理:阿里端侧推理框架上手记
  5. 小白用 TensorRT 给模型加速:ONNX 转 Engine 踩坑实录
  6. 零基础上手:高通 SNPE 模型转换实战(从 ONNX 到 DLC)
  7. 零基础搞懂 SNPE 模型量化:INT8 精度损失 = 0% 的秘密
  8. Python 工程化重构:从 print 脚本到 pytest 项目
  9. 车载语音指令识别:从 44% 到 96% 的调优之路
  10. 车载语音指令识别:SNPE 转换与真机部署实测记录

本文链接:MLflow 实验追踪:从入门到上手 - https://h89.cn/archives/669.html

版权声明:原创文章 遵循 CC 4.0 BY-SA 版权协议,转载请附上原文链接和本声明。

标签: ONNX, PyTorch, MNIST, MLflow, Model Registry, Tracking Server, 超参数, MNN, 模型注册, 实验追踪

欸谨特公众号
微信扫码关注:欸谨特
Agent · 效率工具 · 实战笔记

添加新评论