diff --git a/optuna_dashboard/app.py b/optuna_dashboard/app.py index e1af9ca9..542b7c0f 100644 --- a/optuna_dashboard/app.py +++ b/optuna_dashboard/app.py @@ -8,10 +8,11 @@ import traceback from typing import Union, Dict, List, Optional, TypeVar, Callable, Any, cast from bottle import Bottle, BaseResponse, redirect, request, response, static_file +import optuna from optuna.exceptions import DuplicatedStudyError from optuna.storages import BaseStorage -from optuna.trial import FrozenTrial -from optuna.study import StudyDirection, StudySummary +from optuna.trial import FrozenTrial, TrialState +from optuna.study import StudyDirection, StudySummary, Study from . import serializer @@ -94,6 +95,13 @@ def get_trials( return trials +def get_distribution_name(param_name: str, study: Study) -> str: + for trial in study.trials: + if param_name in trial.distributions: + return trial.distributions[param_name].__class__.__name__ + assert False, "Must not reach here." + + def create_app(storage: BaseStorage) -> Bottle: app = Bottle() @@ -185,6 +193,38 @@ def create_app(storage: BaseStorage) -> Bottle: trials = get_trials(storage, study_id) return serializer.serialize_study_detail(summary, trials) + @app.get("/api/studies//param_importances") + @handle_json_api_exception + def get_param_importances(study_id: int) -> BottleViewReturn: + # TODO(chenghuzi): add support for selecting params and targets via query parameters. + response.content_type = "application/json" + study_name = storage.get_study_name_from_id(study_id) + study = Study(study_name=study_name, storage=storage) + + trials = [trial for trial in study.trials if trial.state == TrialState.COMPLETE] + if len(trials) == 0: + return "" + evaluator = None + params = None + target = None + importances = optuna.importance.get_param_importances( + study, evaluator=evaluator, params=params, target=target + ) + if target is None: + target_name = "Objective Value" + + return { + "target_name": target_name, + "param_importances": [ + { + "name": name, + "importance": importance, + "distribution": get_distribution_name(name, study), + } + for name, importance in importances.items() + ], + } + @app.get("/static/") def send_static(filename: str) -> BottleViewReturn: return static_file(filename, root=STATIC_DIR) diff --git a/optuna_dashboard/static/apiClient.ts b/optuna_dashboard/static/apiClient.ts index cf9348f9..b3f888c4 100644 --- a/optuna_dashboard/static/apiClient.ts +++ b/optuna_dashboard/static/apiClient.ts @@ -168,3 +168,21 @@ export const deleteStudyAPI = (studyId: number) => { return {} }) } + +interface ParamImportancesResponse { + target_name: string + param_importances: ParamImportance[] +} + +export const getParamImportances = ( + studyId: number +): Promise => { + return axiosInstance + .get( + `/api/studies/${studyId}/param_importances`, + {} + ) + .then((res) => { + return res.data + }) +} diff --git a/optuna_dashboard/static/components/HyperparameterImportances.tsx b/optuna_dashboard/static/components/HyperparameterImportances.tsx new file mode 100644 index 00000000..99280bf6 --- /dev/null +++ b/optuna_dashboard/static/components/HyperparameterImportances.tsx @@ -0,0 +1,88 @@ +import * as plotly from "plotly.js-dist" +import React, { FC, useEffect } from "react" +import { getParamImportances } from "../apiClient" +const plotDomId = "graph-hyperparameter-importances" + +// To match colors used by plot_param_importances in optuna. +const plotlyColorsSequentialBlues = [ + "rgb(247,251,255)", + "rgb(222,235,247)", + "rgb(198,219,239)", + "rgb(158,202,225)", + "rgb(107,174,214)", + "rgb(66,146,198)", + "rgb(33,113,181)", + "rgb(8,81,156)", + "rgb(8,48,107)", +] + +const distributionColors = { + UniformDistribution: plotlyColorsSequentialBlues.slice(-1)[0], + LogUniformDistribution: plotlyColorsSequentialBlues.slice(-1)[0], + DiscreteUniformDistribution: plotlyColorsSequentialBlues.slice(-1)[0], + IntUniformDistribution: plotlyColorsSequentialBlues.slice(-2)[0], + IntLogUniformDistribution: plotlyColorsSequentialBlues.slice(-2)[0], + CategoricalDistribution: plotlyColorsSequentialBlues.slice(-4)[0], +} + +export const HyperparameterImportances: FC<{ + studyId: number + numOfTrials: number +}> = ({ studyId, numOfTrials = 0 }) => { + useEffect(() => { + async function fetchAndPlotParamImportances(studyId: number) { + const paramsImportanceData = await getParamImportances(studyId) + plotParamImportances(paramsImportanceData) + } + fetchAndPlotParamImportances(studyId) + }, [numOfTrials]) + return
+} + +const plotParamImportances = (paramsImportanceData: ParamImportances) => { + if (document.getElementById(plotDomId) === null) { + return + } + const param_importances = paramsImportanceData.param_importances.reverse() + const importance_values = param_importances.map((p) => p.importance) + const param_names = param_importances.map((p) => p.name) + const param_colors = param_importances.map( + (p) => distributionColors[p.distribution] + ) + const param_hover_templates = param_importances.map( + (p) => `${p.name} (${p.distribution}): ${p.importance} ` + ) + + const layout: Partial = { + title: "Hyperparameter Importance", + xaxis: { + title: `Importance for ${paramsImportanceData.target_name}`, + }, + yaxis: { + title: "Hyperparameter", + }, + margin: { + l: 50, + r: 50, + b: 50, + }, + showlegend: false, + } + + const plotData: Partial[] = [ + { + type: "bar", + orientation: "h", + x: importance_values, + y: param_names, + text: importance_values.map((v) => String(v.toFixed(2))), + textposition: "outside", + hovertemplate: param_hover_templates, + marker: { + color: param_colors, + }, + }, + ] + + plotly.react(plotDomId, plotData, layout) +} diff --git a/optuna_dashboard/static/components/StudyDetail.tsx b/optuna_dashboard/static/components/StudyDetail.tsx index 4f1590d1..f1ba884f 100644 --- a/optuna_dashboard/static/components/StudyDetail.tsx +++ b/optuna_dashboard/static/components/StudyDetail.tsx @@ -20,6 +20,7 @@ import { Home, Cached } from "@material-ui/icons" import { DataGridColumn, DataGrid } from "./DataGrid" import { GraphParallelCoordinate } from "./GraphParallelCoordinate" +import { HyperparameterImportances } from "./HyperparameterImportances" import { GraphIntermediateValues } from "./GraphIntermediateValues" import { GraphSlice } from "./GraphSlice" import { GraphHistory } from "./GraphHistory" @@ -197,6 +198,20 @@ export const StudyDetail: FC = () => { ) : null} + {studyDetail !== null && isSingleObjectiveStudy(studyDetail) ? ( + + + + + + + + + + ) : null} {studyDetail !== null ? ( diff --git a/optuna_dashboard/static/types/index.d.ts b/optuna_dashboard/static/types/index.d.ts index cc320769..2ae3ee30 100644 --- a/optuna_dashboard/static/types/index.d.ts +++ b/optuna_dashboard/static/types/index.d.ts @@ -9,6 +9,13 @@ declare const URL_PREFIX: string type TrialState = "Running" | "Complete" | "Pruned" | "Fail" | "Waiting" type StudyDirection = "maximize" | "minimize" | "not_set" +type Distribution = + | "UniformDistribution" + | "LogUniformDistribution" + | "DiscreteUniformDistribution" + | "IntUniformDistribution" + | "IntLogUniformDistribution" + | "CategoricalDistribution" declare interface TrialIntermediateValue { step: number @@ -20,6 +27,12 @@ declare interface TrialParam { value: string } +declare interface ParamImportance { + name: string + importance: number + distribution: Distribution +} + declare interface Attribute { key: string value: string @@ -60,3 +73,8 @@ declare interface StudyDetail { declare interface StudyDetails { [study_id: string]: StudyDetail } + +declare interface ParamImportances { + target_name: string + param_importances: ParamImportance[] +} diff --git a/setup.cfg b/setup.cfg index cb21a27d..dc10a76f 100644 --- a/setup.cfg +++ b/setup.cfg @@ -30,6 +30,7 @@ install_requires = optuna>=2.4 bottle typing-extensions;python_version<'3.8' + scikit-learn [options.extras_require] lint =