You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何通过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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.11 15:41:03