如何通过MLFlow API设置默认图表,避免UI重复配置?
解决方案:通过标准化日志+MLFlow API实现默认图表自动生成
当然可以实现,不用每次在UI里重复操作。核心思路是标准化训练代码中的日志逻辑,结合MLFlow的API批量处理Run数据,达成默认图表对比的效果,具体方法如下:
1. 统一日志所有需要的指标和图表
在你的训练代码中,固定好每次都要对比的指标、图表的日志逻辑,让每个Run自动带上这些数据,不用手动在UI添加。
比如,训练完成后,统一执行以下操作:
- 用
mlflow.log_metrics()记录关键指标(如验证集准确率、损失值) - 用
mlflow.log_figure()生成并日志所需图表(如ROC曲线、混淆矩阵、训练/验证损失曲线)
示例代码片段:
import mlflow import matplotlib.pyplot as plt from sklearn.metrics import roc_curve, auc # 训练逻辑... # 日志指标 mlflow.log_metrics({ "val_accuracy": val_acc, "val_loss": val_loss, "train_loss": train_loss }) # 生成并日志ROC曲线 fpr, tpr, _ = roc_curve(val_y, model.predict_proba(val_x)[:,1]) roc_auc = auc(fpr, tpr) plt.figure() plt.plot(fpr, tpr, label=f'ROC curve (area = {roc_auc:.2f})') plt.xlabel('False Positive Rate') plt.ylabel('True Positive Rate') plt.title('Receiver Operating Characteristic') plt.legend(loc="lower right") mlflow.log_figure(plt.gcf(), "roc_curve.png") plt.close()
2. 批量生成对比图表
完成标准化日志后,有两种方式快速获取默认对比视图:
方式一:用MLFlow UI的对比功能
进入新建的Experiment页面,选中所有需要对比的Run,点击顶部的「Compare」按钮。MLFlow会自动展示你日志的所有指标的折线对比图,以及所有日志的图片类图表(每张图表会按Run分别展示,方便对比)。
方式二:用MLFlow API编写自定义对比脚本
如果需要更定制化的对比图表(比如把所有Run的损失曲线画在同一张图里),可以用mlflow.search_runs()获取当前Experiment的所有Run数据,然后自己生成对比图,甚至可以把生成的图再日志到MLFlow中。
示例脚本:
import mlflow import pandas as pd import matplotlib.pyplot as plt # 指定要处理的Experiment ID或名称 experiment_name = "你的验证集Experiment" experiment = mlflow.get_experiment_by_name(experiment_name) runs = mlflow.search_runs(experiment_ids=[experiment.experiment_id]) # 提取所有Run的训练损失数据(假设是按step日志的序列值) plt.figure(figsize=(10,6)) for _, run in runs.iterrows(): plt.plot(run['metrics.train_loss'], label=f"Run {run.run_id[:8]}") plt.xlabel("Training Step") plt.ylabel("Loss") plt.title("Training Loss Comparison Across Runs") plt.legend() mlflow.log_figure(plt.gcf(), "training_loss_comparison.png") plt.close()
3. 进阶:封装通用日志模板
把上述日志逻辑封装成一个可复用的函数,每次新建Experiment启动训练时直接调用,确保所有Run的日志规范完全一致:
def log_default_artifacts(model, val_y, val_x, metrics): # 日志指标 mlflow.log_metrics(metrics) # 生成并日志ROC曲线 fpr, tpr, _ = roc_curve(val_y, model.predict_proba(val_x)[:,1]) roc_auc = auc(fpr, tpr) plt.figure() plt.plot(fpr, tpr, label=f'ROC curve (area = {roc_auc:.2f})') plt.xlabel('False Positive Rate') plt.ylabel('True Positive Rate') plt.title('Receiver Operating Characteristic') plt.legend(loc="lower right") mlflow.log_figure(plt.gcf(), "roc_curve.png") plt.close() # 可添加更多默认图表的生成与日志逻辑... # 训练时调用 log_default_artifacts(trained_model, val_y, val_x, {"val_acc": val_acc, "val_loss": val_loss})
这样每次新建Experiment后,所有Run都会自动包含你需要的图表和指标,无论是用UI对比还是自定义脚本,都不用再手动重复操作。
内容的提问来源于stack exchange,提问作者George Pearse
相关产品推荐
相关产品推荐

