From 981b30e9df20d1b92e222ef9d6d8b51ffb908bc0 Mon Sep 17 00:00:00 2001 From: Cheng Huzi Date: Wed, 10 Mar 2021 01:06:59 -0500 Subject: [PATCH 1/9] Add hyperparameter importances chart --- optuna_dashboard/app.py | 44 +++++++++- optuna_dashboard/static/apiClient.ts | 20 +++++ .../components/HyperparameterImportances.tsx | 88 +++++++++++++++++++ .../static/components/StudyDetail.tsx | 12 +++ optuna_dashboard/static/types/index.d.ts | 18 ++++ 5 files changed, 180 insertions(+), 2 deletions(-) create mode 100644 optuna_dashboard/static/components/HyperparameterImportances.tsx diff --git a/optuna_dashboard/app.py b/optuna_dashboard/app.py index e1af9ca9..d7891a42 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 + + 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: 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": i[0], + "importance": i[1], + "distribution": get_distribution_name(i[0], study), + } + for i 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..0e923dca 100644 --- a/optuna_dashboard/static/apiClient.ts +++ b/optuna_dashboard/static/apiClient.ts @@ -168,3 +168,23 @@ 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 + }) +} + +export default getParamImportances diff --git a/optuna_dashboard/static/components/HyperparameterImportances.tsx b/optuna_dashboard/static/components/HyperparameterImportances.tsx new file mode 100644 index 00000000..59904769 --- /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 +}> = ({ studyId }) => { + useEffect(() => { + async function fetchAndPlotParamImportances(studyId: number) { + const paramsImportanceData = await getParamImportances(studyId) + plotParamImportances(paramsImportanceData) + } + fetchAndPlotParamImportances(studyId) + }) + 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} ` + ) + console.log("param_colors", param_colors) + + 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 bd47d1ae..56e84507 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" @@ -196,6 +197,17 @@ 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[] +} From d228ba74df2744d22034c2bbf80c7ddea0a4f6cd Mon Sep 17 00:00:00 2001 From: Huzi Cheng Date: Thu, 11 Mar 2021 18:32:18 -0500 Subject: [PATCH 2/9] Update optuna_dashboard/app.py Co-authored-by: Masashi Shibata --- optuna_dashboard/app.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/optuna_dashboard/app.py b/optuna_dashboard/app.py index d7891a42..c6abf7a5 100644 --- a/optuna_dashboard/app.py +++ b/optuna_dashboard/app.py @@ -99,7 +99,7 @@ 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 + assert False, "Must not reach here." def create_app(storage: BaseStorage) -> Bottle: From 1de010590b5e564adc0365f7c6b1d9e07c3dd7ae Mon Sep 17 00:00:00 2001 From: Huzi Cheng Date: Thu, 11 Mar 2021 18:32:30 -0500 Subject: [PATCH 3/9] Update optuna_dashboard/app.py Co-authored-by: Masashi Shibata --- optuna_dashboard/app.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/optuna_dashboard/app.py b/optuna_dashboard/app.py index c6abf7a5..264db499 100644 --- a/optuna_dashboard/app.py +++ b/optuna_dashboard/app.py @@ -196,7 +196,7 @@ def create_app(storage: BaseStorage) -> Bottle: @app.get("/api/studies//param_importances") @handle_json_api_exception def get_param_importances(study_id: int) -> BottleViewReturn: - # TODO: add support for selecting params and targets via query parameters. + # 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) From b3fde3384401d945b7dc22cd2b6ba9ee2da25a70 Mon Sep 17 00:00:00 2001 From: Huzi Cheng Date: Thu, 11 Mar 2021 18:39:13 -0500 Subject: [PATCH 4/9] Update optuna_dashboard/static/apiClient.ts Co-authored-by: Masashi Shibata --- optuna_dashboard/static/apiClient.ts | 2 -- 1 file changed, 2 deletions(-) diff --git a/optuna_dashboard/static/apiClient.ts b/optuna_dashboard/static/apiClient.ts index 0e923dca..b3f888c4 100644 --- a/optuna_dashboard/static/apiClient.ts +++ b/optuna_dashboard/static/apiClient.ts @@ -186,5 +186,3 @@ export const getParamImportances = ( return res.data }) } - -export default getParamImportances From 64ddb452c86f2b43bc29da132507592bee6ca11c Mon Sep 17 00:00:00 2001 From: Cheng Huzi Date: Thu, 11 Mar 2021 18:49:24 -0500 Subject: [PATCH 5/9] Update install_requires --- setup.cfg | 1 + 1 file changed, 1 insertion(+) 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 = From 62dbcb0fb844ef5b015d96811d49cb192d76df5e Mon Sep 17 00:00:00 2001 From: Huzi Cheng Date: Sat, 13 Mar 2021 13:15:13 -0500 Subject: [PATCH 6/9] Update optuna_dashboard/static/components/HyperparameterImportances.tsx Co-authored-by: Masashi Shibata --- optuna_dashboard/static/components/HyperparameterImportances.tsx | 1 - 1 file changed, 1 deletion(-) diff --git a/optuna_dashboard/static/components/HyperparameterImportances.tsx b/optuna_dashboard/static/components/HyperparameterImportances.tsx index 59904769..f5969345 100644 --- a/optuna_dashboard/static/components/HyperparameterImportances.tsx +++ b/optuna_dashboard/static/components/HyperparameterImportances.tsx @@ -51,7 +51,6 @@ const plotParamImportances = (paramsImportanceData: ParamImportances) => { const param_hover_templates = param_importances.map( (p) => `${p.name} (${p.distribution}): ${p.importance} ` ) - console.log("param_colors", param_colors) const layout: Partial = { title: "Hyperparameter Importance", From 108c939d73c3f80769143827c97addff6dac3d56 Mon Sep 17 00:00:00 2001 From: Huzi Cheng Date: Sat, 13 Mar 2021 13:22:52 -0500 Subject: [PATCH 7/9] Update optuna_dashboard/app.py Co-authored-by: Masashi Shibata --- optuna_dashboard/app.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/optuna_dashboard/app.py b/optuna_dashboard/app.py index 264db499..542b7c0f 100644 --- a/optuna_dashboard/app.py +++ b/optuna_dashboard/app.py @@ -217,11 +217,11 @@ def create_app(storage: BaseStorage) -> Bottle: "target_name": target_name, "param_importances": [ { - "name": i[0], - "importance": i[1], - "distribution": get_distribution_name(i[0], study), + "name": name, + "importance": importance, + "distribution": get_distribution_name(name, study), } - for i in importances.items() + for name, importance in importances.items() ], } From 468bed28987cb4bff657afda810b3c71a9f6873a Mon Sep 17 00:00:00 2001 From: Cheng Huzi Date: Sat, 13 Mar 2021 14:18:09 -0500 Subject: [PATCH 8/9] Add dependency for updating parameter importance --- .../static/components/HyperparameterImportances.tsx | 6 +++--- optuna_dashboard/static/components/StudyDetail.tsx | 2 +- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/optuna_dashboard/static/components/HyperparameterImportances.tsx b/optuna_dashboard/static/components/HyperparameterImportances.tsx index f5969345..1af1c1ea 100644 --- a/optuna_dashboard/static/components/HyperparameterImportances.tsx +++ b/optuna_dashboard/static/components/HyperparameterImportances.tsx @@ -26,15 +26,15 @@ const distributionColors = { } export const HyperparameterImportances: FC<{ - studyId: number -}> = ({ studyId }) => { + studyId: number, numOfTrials: number +}> = ({ studyId, numOfTrials = 0 }) => { useEffect(() => { async function fetchAndPlotParamImportances(studyId: number) { const paramsImportanceData = await getParamImportances(studyId) plotParamImportances(paramsImportanceData) } fetchAndPlotParamImportances(studyId) - }) + }, [numOfTrials]) return
} diff --git a/optuna_dashboard/static/components/StudyDetail.tsx b/optuna_dashboard/static/components/StudyDetail.tsx index 56e84507..154a6bf0 100644 --- a/optuna_dashboard/static/components/StudyDetail.tsx +++ b/optuna_dashboard/static/components/StudyDetail.tsx @@ -202,7 +202,7 @@ export const StudyDetail: FC = () => { - + From be2af8ea9142c0a2763b20250a4427b85776d1a6 Mon Sep 17 00:00:00 2001 From: Cheng Huzi Date: Sat, 13 Mar 2021 14:56:12 -0500 Subject: [PATCH 9/9] Format code --- .../static/components/HyperparameterImportances.tsx | 3 ++- optuna_dashboard/static/components/StudyDetail.tsx | 5 ++++- 2 files changed, 6 insertions(+), 2 deletions(-) diff --git a/optuna_dashboard/static/components/HyperparameterImportances.tsx b/optuna_dashboard/static/components/HyperparameterImportances.tsx index 1af1c1ea..99280bf6 100644 --- a/optuna_dashboard/static/components/HyperparameterImportances.tsx +++ b/optuna_dashboard/static/components/HyperparameterImportances.tsx @@ -26,7 +26,8 @@ const distributionColors = { } export const HyperparameterImportances: FC<{ - studyId: number, numOfTrials: number + studyId: number + numOfTrials: number }> = ({ studyId, numOfTrials = 0 }) => { useEffect(() => { async function fetchAndPlotParamImportances(studyId: number) { diff --git a/optuna_dashboard/static/components/StudyDetail.tsx b/optuna_dashboard/static/components/StudyDetail.tsx index 154a6bf0..2b42afca 100644 --- a/optuna_dashboard/static/components/StudyDetail.tsx +++ b/optuna_dashboard/static/components/StudyDetail.tsx @@ -202,7 +202,10 @@ export const StudyDetail: FC = () => { - +