Optuna 如何用 Plotly 可视化函数分析调参结果
2026/9/15 18:32:54 网站建设 项目流程

Optuna 如何用 Plotly 可视化函数分析调参结果

【免费下载链接】optunaA hyperparameter optimization framework项目地址: https://gitcode.com/GitHub_Trending/op/optuna

跑完一轮study.optimize()之后,调参结果还留在Study对象里:哪些 trial 有效、哪些参数相关、最优值何时出现,光看日志无法回答。Optuna 的optuna.visualization模块提供一组基于 Plotly 的可视化函数,把优化历史、参数关系、参数重要性等内容直接画成交互式图表,用于分析调参结果。本文给出从安装依赖到逐图分析再到修改图表的完整路径,示例沿用仓库教程 tutorial/10_key_features/005_visualization.py。

前提只有一个:安装 Plotly。Optuna 本身支持 Python 3.9 及以上(见 docs/source/installation.rst)。

安装依赖并确认 Plotly 可用

$ pip install optuna $ pip install plotly

如果在 Jupyter Notebook 中运行,教程还要求安装nbformat$ pip install nbformat)。

安装后确认版本满足要求。optuna.visualization依赖 plotly 4.0.0 或更高版本,可用模块自带的is_available()判断(实现在 optuna/visualization/_utils.py):

from optuna.visualization import is_available is_available() # True 表示 plotly 已安装且版本可用

返回False时,按该函数文档提示执行$ pip install -U plotly>=4.0.0升级。

跑一个调参任务得到 Study

以教程中的 FashionMNIST 分类任务为例:目标函数通过trial.suggest_*采样超参数(层数n_layers、每层单元数n_units_l{i}、学习率lr),每个 epoch 用trial.report(val_accuracy, epoch)上报中间值,配合MedianPruner做剪枝。最后两个 epoch 上报的验证准确率作为返回值:

import optuna import torch import torch.nn as nn import torch.nn.functional as F import torchvision SEED = 13 torch.manual_seed(SEED) DEVICE = torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu") DIR = ".." # 教程相对 tutorial/ 目录的写法,指向数据集存放位置,可按实际路径替换 BATCHSIZE = 128 N_TRAIN_EXAMPLES = BATCHSIZE * 30 N_VALID_EXAMPLES = BATCHSIZE * 10 def define_model(trial): n_layers = trial.suggest_int("n_layers", 1, 2) layers = [] in_features = 28 * 28 for i in range(n_layers): out_features = trial.suggest_int(f"n_units_l{i}", 64, 512) layers.append(nn.Linear(in_features, out_features)) layers.append(nn.ReLU()) in_features = out_features layers.append(nn.Linear(in_features, 10)) layers.append(nn.LogSoftmax(dim=1)) return nn.Sequential(*layers) def train_model(model, optimizer, train_loader): model.train() for batch_idx, (data, target) in enumerate(train_loader): data, target = data.view(-1, 28 * 28).to(DEVICE), target.to(DEVICE) optimizer.zero_grad() F.nll_loss(model(data), target).backward() optimizer.step() def eval_model(model, valid_loader): model.eval() correct = 0 with torch.no_grad(): for batch_idx, (data, target) in enumerate(valid_loader): data, target = data.view(-1, 28 * 28).to(DEVICE), target.to(DEVICE) pred = model(data).argmax(dim=1, keepdim=True) correct += pred.eq(target.view_as(pred)).sum().item() return correct / N_VALID_EXAMPLES def objective(trial): train_dataset = torchvision.datasets.FashionMNIST( DIR, train=True, download=True, transform=torchvision.transforms.ToTensor() ) train_loader = torch.utils.data.DataLoader( torch.utils.data.Subset(train_dataset, list(range(N_TRAIN_EXAMPLES))), batch_size=BATCHSIZE, shuffle=True, ) val_dataset = torchvision.datasets.FashionMNIST( DIR, train=False, transform=torchvision.transforms.ToTensor() ) val_loader = torch.utils.data.DataLoader( torch.utils.data.Subset(val_dataset, list(range(N_VALID_EXAMPLES))), batch_size=BATCHSIZE, shuffle=True, ) model = define_model(trial).to(DEVICE) optimizer = torch.optim.Adam( model.parameters(), trial.suggest_float("lr", 1e-5, 1e-1, log=True) ) for epoch in range(10): train_model(model, optimizer, train_loader) val_accuracy = eval_model(model, val_loader) trial.report(val_accuracy, epoch) if trial.should_prune(): raise optuna.exceptions.TrialPruned() return val_accuracy study = optuna.create_study( direction="maximize", sampler=optuna.samplers.TPESampler(seed=SEED), pruner=optuna.pruners.MedianPruner(), ) study.optimize(objective, n_trials=30, timeout=300)

数据集通过download=True从网络下载,需要可访问网络的环境。优化结束后,下面所有图都只依赖这一个study对象。

分析优化过程:历史与中间值图

from optuna.visualization import plot_intermediate_values from optuna.visualization import plot_optimization_history plot_optimization_history(study) plot_intermediate_values(study)
  • plot_optimization_history(study)画所有 trial 的目标值走势,并叠加逐 trial 的 Best Value 折线。支持传入target(自定义要显示的数值)和target_name(坐标轴与图例名称),还可以传多个 study 对比优化历史(见 optuna/visualization/_optimization_history.py)。
  • plot_intermediate_values(study)画每个 trial 的中间值曲线(学习曲线),数据来自目标函数里的trial.report()调用。如果 study 里没有任何中间值,函数会提示 "You need to set up the pruning feature to utilizeplot_intermediate_values()"——即必须在训练中做剪枝/上报才会出图。

在 Jupyter Notebook 中执行这两行后,交互式图表直接内嵌显示。

分析参数关系:平行坐标、轮廓、切片与排名

from optuna.visualization import plot_contour from optuna.visualization import plot_parallel_coordinate from optuna.visualization import plot_rank from optuna.visualization import plot_slice plot_parallel_coordinate(study) plot_contour(study) plot_slice(study) plot_rank(study)

四个函数都接受params参数选择要画的参数,缺省为全部参数,例如只关注学习率和层数:

plot_parallel_coordinate(study, params=["lr", "n_layers"]) plot_contour(study, params=["lr", "n_layers"]) plot_slice(study, params=["lr", "n_layers"])
  • plot_parallel_coordinate:画高维参数关系,每个参数一条纵轴;缺少某参数的 trial 会连到特殊的None刻度。
  • plot_contour:两两参数的轮廓图,含缺失值的 trial 不绘制。注意当directionminimize或传了target时色标方向会反转。
  • plot_slice:每个参数一个子图的切片散点,按目标值着色,看单个参数的取值分布与好坏 trial 的关系。
  • plot_rank:按目标值排名着色的散点图,要求 plotly 5.0.0 及以上(低于此版本时其余函数仍可用)。

函数文档中统一的target参数说明:单目标 study 缺省画目标值;多目标 study 必须显式传target指定画哪一维。

分析参数重要性与分布:importances 与 EDF

from optuna.visualization import plot_edf from optuna.visualization import plot_param_importances plot_param_importances(study) plot_edf(study)
  • plot_param_importances(study)画各超参数对目标值的重要度条形图,默认使用PedAnovaImportanceEvaluator,也接受evaluatorparams参数。教程里还展示了用它分析 trial 耗时:

    optuna.visualization.plot_param_importances( study, target=lambda t: t.duration.total_seconds(), target_name="duration" )
  • plot_edf(study)画目标值的经验分布函数(EDF),只统计 complete 状态的 trial。文档说明 EDF 可用于分析和改进搜索空间,也可传多个 study 对比。

查看 trial 的时间线

from optuna.visualization import plot_timeline plot_timeline(study)

plot_timeline画各 trial 的执行时间段(lifetime),用于查看 trial 在时间轴上的排布和时长差异。参数n_recent_trials控制只画最近多少个 trial:缺省为全部 trial,指定时必须为正整数,否则抛ValueError(见 optuna/visualization/_timeline.py)。

修改与保存生成的图表

optuna.visualization中每个函数都返回可编辑的plotly.graph_objects.Figure对象,可以用 Plotly API 直接改。教程中的例子是替换plot_intermediate_values生成的图标题和坐标轴标签:

fig = plot_intermediate_values(study) fig.update_layout( title="Hyperparameter optimization for FashionMNIST classification", xaxis_title="Epoch", yaxis_title="Validation Accuracy", )

改完后fig就是普通 Plotly Figure,Jupyter 中显示它即可看到更新后的标题。

边界与替代方案

  • Matplotlib 后端:如果偏好 Matplotlib,教程说明只需把optuna.visualization换成optuna.visualization.matplotlib,函数一一对应;该后端需要$ pip install matplotlib
  • Optuna Dashboard:tutorial 提到,把 study 持久化到 RDB 后端后,可执行$ pip install optuna-dashboard再运行$ optuna-dashboard sqlite:///example-study.db,用交互图表和表格查看优化历史、参数重要度等。
  • 多目标plot_pareto_front属于多目标优化场景,教程要求另行参考多目标教程,不在本文范围内。

各函数完整 API 参考见 docs/source/reference/visualization/index.rst,示例脚本可在 docs/visualization_examples/ 下逐图对照运行。

【免费下载链接】optunaA hyperparameter optimization framework项目地址: https://gitcode.com/GitHub_Trending/op/optuna

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询