AI实验追踪折腾手记
这篇东西主要梳理一件事:从手动记实验笔记到用 MLflow 做自动化追踪,中间踩了哪些坑,哪些问题是真问题,哪些是臆想出来的需求。那时候刚做 NLP 项目,训练脚本跑得慢,一个实验要半天。
为什么要搞实验追踪
先说说我最开始是怎么记实验的。
手动记录的痛苦时期
那时候刚做 NLP 项目,训练脚本跑得慢,一个实验要半天。为了不忘记每次的配置,我在笔记软件里建了个表格:
实验1 lr=0.001 batch=32 bert-base acc=0.78
实验2 lr=0.0001 batch=64 bert-base acc=0.79
实验3 lr=0.001 batch=32 bert-large acc=0.81
看起来还行,对吧?
但问题很快就来了:
- 参数越多,表格越宽。光 LR、batch、模型还不够,还有 dropout、warmup、优化器、随机种子、数据增强方式、训练轮数、early stopping 阈值。一行能写几十个参数。
- 环境变化记不住。PyTorch 1.8 升到 1.9 后,同样的参数效果不一样了,但表格里没有版本记录。
- 模型文件管理混乱。实验跑完了,保存的模型文件叫
model_1.pth、model_2.pth,过两周想不起来哪个对应哪次实验。 - 代码回不去。想复现"实验 3 的效果",但那之后代码改了好几版,git log 告诉你有很多提交,不知道当时是哪个 commit。
这些问题不是小问题。做项目的人都知道,“当时效果最好但记不清怎么来的模型"是最折磨人的。
明确需求:要记什么
从上面的痛苦里,我慢慢梳理出几个核心需求:
- 每次实验的参数必须自动记录,不能靠人手写。人一定会忘,一定会写漏。
- 环境和依赖也要记录。Python 版本、库版本、CUDA 版本,这些都可能影响结果。
- 模型文件要有编号,能跟实验记录关联起来。
- 训练过程中的指标要能看。loss、accuracy 这些曲线,不要只看最终数值,过程很重要。
- 能快速对比几次实验。比如想看"改了 LR 后效果如何”,要能直观看到差异。
有了这些需求,自然就想到:不能靠手动,得自动化。
第一次尝试:自己写记录脚本
最开始我没想用现成工具,觉得"记录几个参数而已,自己写个脚本就够了"。
于是写了个简单的东西:
import json
from datetime import datetime
def log_experiment(params, metrics, model_path):
log_entry = {
"timestamp": datetime.now().isoformat(),
"params": params,
"metrics": metrics,
"model_path": model_path
}
with open("experiments.jsonl", "a") as f:
f.write(json.dumps(log_entry) + "\n")
然后在训练脚本里调用:
params = {"lr": 0.001, "batch_size": 32, "model": "bert-base"}
metrics = {"train_acc": 0.85, "val_acc": 0.78}
model_path = "models/bert_base_001.pt"
log_experiment(params, metrics, model_path)
踩坑一:参数传递不完整
第一个问题是参数传递不完整。
训练脚本里参数到处都是:命令行参数、配置文件、hardcoded 的常数、函数默认参数。把它们全部收集起来传给 log_experiment 很麻烦。
# 到处都是参数
parser.add_argument("--lr", type=float, default=0.001)
config = load_config("config.yaml") # 里面又有 lr
optimizer = Adam(model.parameters(), lr=args.lr or config.get("lr", 0.001))
容易漏,容易不一致。改了参数但忘了更新 log_experiment 的调用,导致记录的参数跟实际用的不一样。
踩坑二:环境信息没记
第二个问题是环境信息完全没记。
同样的代码,在本地机器和服务器上跑,效果可能不一样。Python 版本、PyTorch 版本、CUDA 版本、cuDNN 版本,这些都能影响结果。
但我的脚本完全没有记录这些。
踩坑三:可视化要看手动写
第三个问题是记录有了,但看记录很痛苦。
experiments.jsonl 是一行行 JSON,想看几次实验的对比,得自己写脚本解析、画图。或者用 jq 在命令行里查,但很快就会觉得麻烦。
第二次尝试:用 MLflow
踩了一圈坑后,开始看现成的实验追踪工具。试了几个后,最终选了 MLflow。
选择它的原因很简单:
- 安装简单:
pip install mlflow就能用。 - 跟代码集成容易:几行代码就能记录参数、指标、模型。
- 自带 UI:启动后能直接在浏览器里看实验记录。
- 不绑定框架:PyTorch、TensorFlow、sklearn 都能用。
基础用法
先说最基础的用法。
安装和启动:
pip install mlflow
# 启动 UI
mlflow ui --port 5000
然后在训练脚本里加几行:
import mlflow
# 开始一个实验记录
with mlflow.start_run():
# 记录参数
mlflow.log_param("lr", 0.001)
mlflow.log_param("batch_size", 32)
mlflow.log_param("model", "bert-base")
# 训练过程
for epoch in range(epochs):
train_loss, train_acc = train_one_epoch(...)
val_loss, val_acc = validate(...)
# 记录每个 epoch 的指标
mlflow.log_metric("train_loss", train_loss, step=epoch)
mlflow.log_metric("train_acc", train_acc, step=epoch)
mlflow.log_metric("val_loss", val_loss, step=epoch)
mlflow.log_metric("val_acc", val_acc, step=epoch)
# 记录最终指标
mlflow.log_metric("final_val_acc", val_acc)
# 保存模型
mlflow.pytorch.log_model(model, "model")
就这么简单。每次运行脚本,MLflow 会自动记录这次实验的所有信息,并在 UI 里显示。
进阶用法:自动记录参数
前面的例子还是手动记录参数,容易漏。MLflow 提供了自动记录参数的功能:
from mlflow.tracking import MlflowClient
def log_all_params(params_dict):
for key, value in params_dict.items():
mlflow.log_param(key, value)
# 使用
params = vars(args) # 假设 args 是 argparse 解析出来的
params.update(config) # 加上配置文件里的参数
log_all_params(params)
这样至少不会漏掉命令行参数和配置文件参数。
进阶用法:记录环境和依赖
记录环境信息:
import mlflow
import sys
with mlflow.start_run():
# 记录 Python 版本
mlflow.log_param("python_version", sys.version)
# 记录主要包的版本
import torch
mlflow.log_param("torch_version", torch.__version__)
mlflow.log_param("cuda_version", torch.version.cuda)
# 记录 git commit
import subprocess
commit = subprocess.check_output(["git", "rev-parse", "HEAD"]).decode().strip()
mlflow.log_param("git_commit", commit)
这样就能知道每次实验的环境情况。
进阶用法:记录自定义指标
除了 loss、accuracy 这些常见指标,有时候想记录一些自定义的东西。
比如我想记录"训练过程中参数梯度的平均值":
def log_gradient_stats(model, step):
total_norm = 0
for p in model.parameters():
if p.grad is not None:
param_norm = p.grad.data.norm(2)
total_norm += param_norm.item() ** 2
total_norm = total_norm ** 0.5
mlflow.log_metric("gradient_norm", total_norm, step=step)
这个指标在调试梯度爆炸/消失问题时很有用。
MLflow 实践中的踩坑记录
用了一段时间 MLflow 后,发现也有一些坑需要注意。
踩坑一:实验名称管理混乱
最开始没给实验起名字,每次 mlflow.start_run() 都会创建一个新实验,UI 里很快就会出现一堆"Default"实验,分不清哪个项目。
解决办法是显式设置实验名称:
mlflow.set_experiment("sentiment-analysis-bert")
with mlflow.start_run(run_name="lr-0.001-batch-32"):
# 训练代码
这样相关实验会归类到同一个 experiment 下。
踩坑二:指标太多导致 UI 卡顿
有一次每个 epoch 记录了二十多个指标,跑了 50 个 epoch 后,UI 打开非常慢,渲染图表要等半天。
解决办法是精简指标,只记录关键信息:
# 不推荐:记录太多
mlflow.log_metric("layer1_weight_mean", layer1_weights.mean(), step=epoch)
mlflow.log_metric("layer1_weight_std", layer1_weights.std(), step=epoch)
mlflow.log_metric("layer2_weight_mean", layer2_weights.mean(), step=epoch)
# ... 几十个这样的指标
# 推荐:只记录关键指标
mlflow.log_metric("train_loss", train_loss, step=epoch)
mlflow.log_metric("val_loss", val_loss, step=epoch)
mlflow.log_metric("val_acc", val_acc, step=epoch)
# 确实需要详细指标时,可以用 artifact 保存
if epoch % 10 == 0: # 每 10 个 epoch 保存一次详细统计
stats = compute_detailed_stats(model)
with open(f"stats_epoch_{epoch}.json", "w") as f:
json.dump(stats, f)
mlflow.log_artifact(f"stats_epoch_{epoch}.json")
踩坑三:模型文件太大
MLflow 会记录模型文件,如果模型很大,很快就会占满磁盘。
解决办法:
- 只保存最佳模型:
best_val_acc = 0
for epoch in range(epochs):
val_acc = validate(...)
if val_acc > best_val_acc:
best_val_acc = val_acc
mlflow.pytorch.log_model(model, "best_model")
- 用 checkpoint 代替完整模型:
# 只保存参数,不保存完整模型
torch.save(model.state_dict(), "checkpoint.pth")
mlflow.log_artifact("checkpoint.pth")
- 定期清理:
# 删除旧的实验
mlflow runs delete --experiment-id <experiment-id> <run-id>
踩坑四:远程服务器配置
在服务器上跑实验,想在本地看 MLflow UI,需要配置远程 tracking。
配置方法:
import mlflow
# 设置远程 tracking server
mlflow.set_tracking_uri("http://your-server:5000")
# 然后正常使用
with mlflow.start_run():
# 训练代码
或者在环境变量里设置:
export MLFLOW_TRACKING_URI=http://your-server:5000
要注意防火墙和安全设置,不要把 tracking server 暴露在公网。
整体架构
折腾了一圈后,我现在的实验追踪架构大概是这样的:
- 训练脚本通过 MLflow SDK 记录参数、指标、模型
- 记录的数据保存在本地文件系统或远程服务器
- 通过 MLflow UI 查看和分析实验记录
实际效果
用了 MLflow 一段时间后,效果还是很明显的:
- 再也不会"记不清怎么来的"。每次实验都有完整记录,参数、环境、指标、模型都能查到。
- 快速对比实验。在 UI 里选几个 run,就能直接对比它们的指标曲线。
- 复现更容易。知道每次实验的环境和代码版本,复现成功率大幅提升。
- 节省时间。不用花时间整理笔记,专注在调参和改进模型上。
但也有一些边界:
- 不是万能的。MLflow 只记录你让它记录的东西,如果代码本身有问题或数据有问题,它不会自动发现。
- 学习成本。虽然基础用法简单,但进阶功能还是需要一点学习时间。
- 资源占用。tracking server 和 UI 会占一些内存和磁盘,要注意清理。
结语
实验追踪这件事,跟写测试、写文档一样:不做的时候觉得麻烦,做了之后觉得省事。
从手动记笔记到 MLflow 自动化追踪,这个过程让我明白一个道理:凡是重复的、容易出错的、应该记录但总忘记录的事情,都应该交给机器做。
现在回头看,最开始那个手写的 log_experiment 脚本,确实解决了"记录"这个问题,但没有解决"容易漏"、“难对比”、“环境信息缺失"这些更深的问题。
MLflow 也不是完美工具,但它把实验追踪这件事做到了"够用"的程度:安装简单、使用简单、功能够用、不绑定框架。
如果你的实验数量还不多,手动记记可能还够用。但如果开始觉得"记不清哪个模型是怎么来的”,不妨试试 MLflow 之类的工具。
毕竟,机器应该帮我们省时间,而不是让我们花更多时间在"记笔记"这件事上。
版权声明: 本文首发于 指尖魔法屋-AI实验追踪折腾手记(https://blog.thinkmoon.cn/post/284-ai-experiment-tracking-manual-automated-practice/) 转载或引用必须申明原指尖魔法屋来源及源地址!
评论
使用 GitHub 账号登录后即可留言,支持 Markdown。