From 653b18d4c54bd8d7b2a8676fd6f09edd132a0667 Mon Sep 17 00:00:00 2001 From: keisuke-umezawa Date: Sun, 24 Sep 2023 14:29:25 +0900 Subject: [PATCH 01/24] Implement log scale of countor --- .../ts/components/GraphContour.tsx | 30 +++++++++++-------- 1 file changed, 18 insertions(+), 12 deletions(-) diff --git a/optuna_dashboard/ts/components/GraphContour.tsx b/optuna_dashboard/ts/components/GraphContour.tsx index c75e1500..b186fffe 100644 --- a/optuna_dashboard/ts/components/GraphContour.tsx +++ b/optuna_dashboard/ts/components/GraphContour.tsx @@ -162,19 +162,19 @@ const plotContour = ( return } - const xAxis = getAxisInfo(study, trials, xParam) - const yAxis = getAxisInfo(study, trials, yParam) + const xAxis = getAxisInfo(trials, xParam) + const yAxis = getAxisInfo(trials, yParam) const xIndices = xAxis.indices const yIndices = yAxis.indices const layout: Partial = { xaxis: { title: xParam.name, - type: xAxis.isCat ? "category" : undefined, + type: xAxis.isCat ? "category" : xAxis.isLog ? "log" : "linear", }, yaxis: { title: yParam.name, - type: yAxis.isCat ? "category" : undefined, + type: yAxis.isCat ? "category" : yAxis.isLog ? "log" : "linear", }, margin: { l: 50, @@ -278,9 +278,19 @@ const getAxisInfoForNumericalParams = ( paramName: string, distribution: FloatDistribution | IntDistribution ): AxisInfo => { - const padding = (distribution.high - distribution.low) * PADDING_RATIO - const min = distribution.low - padding - const max = distribution.high + padding + let min = 0 + let max = 0 + if (distribution.log) { + const padding = + (Math.log10(distribution.high) - Math.log10(distribution.low)) * + PADDING_RATIO + min = Math.pow(10, Math.log10(distribution.low) - padding) + max = Math.pow(10, Math.log10(distribution.high) + padding) + } else { + const padding = (distribution.high - distribution.low) * PADDING_RATIO + min = distribution.low - padding + max = distribution.high + padding + } const values = trials.map( (trial) => @@ -341,11 +351,7 @@ const getAxisInfoForCategoricalParams = ( } } -const getAxisInfo = ( - study: StudyDetail, - trials: Trial[], - param: SearchSpaceItem -): AxisInfo => { +const getAxisInfo = (trials: Trial[], param: SearchSpaceItem): AxisInfo => { if (param.distribution.type === "CategoricalDistribution") { return getAxisInfoForCategoricalParams( trials, From 5c987c3dec6ac78d400fd8ca716cac22ce03c898 Mon Sep 17 00:00:00 2001 From: keisuke-umezawa Date: Sun, 24 Sep 2023 14:48:01 +0900 Subject: [PATCH 02/24] Implement log scale of parcoodes --- .../ts/components/GraphParallelCoordinate.tsx | 41 +++++++++++++++---- 1 file changed, 34 insertions(+), 7 deletions(-) diff --git a/optuna_dashboard/ts/components/GraphParallelCoordinate.tsx b/optuna_dashboard/ts/components/GraphParallelCoordinate.tsx index 6726ba2f..02b3b1f6 100644 --- a/optuna_dashboard/ts/components/GraphParallelCoordinate.tsx +++ b/optuna_dashboard/ts/components/GraphParallelCoordinate.tsx @@ -169,6 +169,20 @@ const plotCoordinate = ( .join("") } + const calculateLogScale = (values: number[]) => { + const logValues = values.map((v) => { + return Math.log10(v) + }) + const minValue = Math.min(...logValues) + const maxValue = Math.max(...logValues) + const tickvals = Array.from( + { length: Math.ceil(maxValue) - Math.floor(minValue) + 1 }, + (_, i) => i + Math.floor(minValue) + ) + const ticktext = tickvals.map((x) => `${Math.pow(10, x).toPrecision(3)}`) + return { logValues, tickvals, ticktext } + } + const dimensions = targets.map((target) => { if (target.kind === "objective" || target.kind === "user_attr") { const values: number[] = trials.map( @@ -187,13 +201,7 @@ const plotCoordinate = ( const values: number[] = trials.map( (t) => target.getTargetValue(t) as number ) - if (s.distribution.type !== "CategoricalDistribution") { - return { - label: breakLabelIfTooLong(s.name), - values: values, - range: [s.distribution.low, s.distribution.high], - } - } else { + if (s.distribution.type === "CategoricalDistribution") { // categorical const vocabArr: string[] = s.distribution.choices.map((c) => c.value) const tickvals: number[] = vocabArr.map((v, i) => i) @@ -205,6 +213,25 @@ const plotCoordinate = ( tickvals: tickvals, ticktext: vocabArr, } + } else if (s.distribution.log) { + // numerical and log + const values = trials.map((t) => { + return target.getTargetValue(t) as number + }) + const { logValues, tickvals, ticktext } = calculateLogScale(values) + return { + label: breakLabelIfTooLong(s.name), + values: logValues, + tickvals: tickvals, + ticktext: ticktext, + } + } else { + // numerical and non-log + return { + label: breakLabelIfTooLong(s.name), + values: values, + range: [s.distribution.low, s.distribution.high], + } } } }) From da7aa6acf2f99164bcdaf2f5c27221ad45d447df Mon Sep 17 00:00:00 2001 From: keisuke-umezawa Date: Sun, 24 Sep 2023 14:56:39 +0900 Subject: [PATCH 03/24] Add range of plot --- .../ts/components/GraphParallelCoordinate.tsx | 14 +++++++------- 1 file changed, 7 insertions(+), 7 deletions(-) diff --git a/optuna_dashboard/ts/components/GraphParallelCoordinate.tsx b/optuna_dashboard/ts/components/GraphParallelCoordinate.tsx index 02b3b1f6..5d5b1222 100644 --- a/optuna_dashboard/ts/components/GraphParallelCoordinate.tsx +++ b/optuna_dashboard/ts/components/GraphParallelCoordinate.tsx @@ -175,12 +175,13 @@ const plotCoordinate = ( }) const minValue = Math.min(...logValues) const maxValue = Math.max(...logValues) + const range = [minValue, maxValue] const tickvals = Array.from( { length: Math.ceil(maxValue) - Math.floor(minValue) + 1 }, (_, i) => i + Math.floor(minValue) ) const ticktext = tickvals.map((x) => `${Math.pow(10, x).toPrecision(3)}`) - return { logValues, tickvals, ticktext } + return { logValues, range, tickvals, ticktext } } const dimensions = targets.map((target) => { @@ -215,15 +216,14 @@ const plotCoordinate = ( } } else if (s.distribution.log) { // numerical and log - const values = trials.map((t) => { - return target.getTargetValue(t) as number - }) - const { logValues, tickvals, ticktext } = calculateLogScale(values) + const { logValues, range, tickvals, ticktext } = + calculateLogScale(values) return { label: breakLabelIfTooLong(s.name), values: logValues, - tickvals: tickvals, - ticktext: ticktext, + range, + tickvals, + ticktext, } } else { // numerical and non-log From a90b89834f935b96d904c0d5c1c1aae8e3b7e7b8 Mon Sep 17 00:00:00 2001 From: Contramundum Date: Fri, 29 Sep 2023 11:16:59 +0900 Subject: [PATCH 04/24] Simplify EP implementation --- optuna_dashboard/preferential/samplers/gp.py | 34 +++++++++++--------- 1 file changed, 19 insertions(+), 15 deletions(-) diff --git a/optuna_dashboard/preferential/samplers/gp.py b/optuna_dashboard/preferential/samplers/gp.py index d023e741..33a7512e 100644 --- a/optuna_dashboard/preferential/samplers/gp.py +++ b/optuna_dashboard/preferential/samplers/gp.py @@ -156,6 +156,18 @@ def _truncnorm_mean_var_logz(alpha: Tensor) -> tuple[Tensor, Tensor, Tensor]: return mean, var, logz +def _observation(var0: Tensor, mean0: Tensor, noise_var: Tensor) -> tuple[Tensor, Tensor, Tensor]: + obs_var = var0 + noise_var + obs_sigma = torch.sqrt(obs_var) + alpha = -mean0 / torch.clamp_min(obs_sigma, min=1e-20) + mean_norm, var_norm, logz = _truncnorm_mean_var_logz(alpha) + + denom_factor = 1 / (noise_var + var_norm * var0) + da = (1 - var_norm) * denom_factor + db = (mean0 * (1 - var_norm) + obs_sigma * mean_norm) * denom_factor + return (da, db, logz) + + def _orthants_MVN_EP( cov0: Tensor, preferences: Tensor, noise_var: Tensor, cycles: int ) -> tuple[Tensor, Tensor, Tensor]: @@ -176,25 +188,17 @@ def _orthants_MVN_EP( r0 = (1 - var1 * virtual_obs_a[i]).reciprocal() var0 = var1 * r0 - mean0 = (mean1 + var1 * virtual_obs_b[i]) * r0 + mean0 = (mean1 - var1 * virtual_obs_b[i]) * r0 - obs_var = var0 + noise_var - obs_sigma = torch.sqrt(obs_var) - alpha = -mean0 / torch.clamp_min(obs_sigma, min=1e-20) - mean_norm, var_norm, logz = _truncnorm_mean_var_logz(alpha) + virtual_obs_a2, virtual_obs_b2, logz = _observation(var0, mean0, noise_var) - kalman_factor = var0 / torch.clamp_min(obs_var, min=1e-20) - mean2 = mean0 + obs_sigma * mean_norm * kalman_factor - var2 = kalman_factor * (noise_var + var_norm * var0) - - var1_var2_inv = torch.clamp_min(var1 * var2, min=1e-20).reciprocal() - db = (mean1 * var2 - mean2 * var1) * var1_var2_inv - da = (var1 - var2) * var1_var2_inv - virtual_obs_b[i] = virtual_obs_b[i] + db - virtual_obs_a[i] = virtual_obs_a[i] + da + da = virtual_obs_a2 - virtual_obs_a[i] + db = virtual_obs_b2 - virtual_obs_b[i] + virtual_obs_a[i] = virtual_obs_a2 + virtual_obs_b[i] = virtual_obs_b2 dr = (1 + var1 * da).reciprocal() - mu = mu - Sxy * ((db + mean1 * da) * dr) + mu = mu + Sxy * ((db - mean1 * da) * dr) cov = cov - (Sxy[:, None] * (da * dr)) @ Sxy[None, :] log_zs[i] = logz return mu, cov, torch.sum(log_zs) From f2f233c75356e531e1aaf0a9a6e25bd023503066 Mon Sep 17 00:00:00 2001 From: Contramundum Date: Fri, 29 Sep 2023 11:42:59 +0900 Subject: [PATCH 05/24] Add some measures to prevent nan when noise_var==0 --- optuna_dashboard/preferential/samplers/gp.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/optuna_dashboard/preferential/samplers/gp.py b/optuna_dashboard/preferential/samplers/gp.py index 33a7512e..5c3d849b 100644 --- a/optuna_dashboard/preferential/samplers/gp.py +++ b/optuna_dashboard/preferential/samplers/gp.py @@ -162,7 +162,7 @@ def _observation(var0: Tensor, mean0: Tensor, noise_var: Tensor) -> tuple[Tensor alpha = -mean0 / torch.clamp_min(obs_sigma, min=1e-20) mean_norm, var_norm, logz = _truncnorm_mean_var_logz(alpha) - denom_factor = 1 / (noise_var + var_norm * var0) + denom_factor = 1 / torch.clamp_min(noise_var + var_norm * var0, min=1e-20) da = (1 - var_norm) * denom_factor db = (mean0 * (1 - var_norm) + obs_sigma * mean_norm) * denom_factor return (da, db, logz) From 8a8cd29f21410a65f3c577c43fc27e77ed1b12af Mon Sep 17 00:00:00 2001 From: Contramundum Date: Fri, 29 Sep 2023 12:34:35 +0900 Subject: [PATCH 06/24] Do exhaustive evaluation of acquisition function when the search space is small enough --- optuna_dashboard/preferential/samplers/gp.py | 47 ++++++++++++++++++-- 1 file changed, 44 insertions(+), 3 deletions(-) diff --git a/optuna_dashboard/preferential/samplers/gp.py b/optuna_dashboard/preferential/samplers/gp.py index d023e741..f4415a38 100644 --- a/optuna_dashboard/preferential/samplers/gp.py +++ b/optuna_dashboard/preferential/samplers/gp.py @@ -5,6 +5,7 @@ import math from typing import Any from typing import Callable from typing import cast +import warnings import botorch.acquisition.analytic import botorch.models.model @@ -17,6 +18,8 @@ import numpy as np import optuna import optuna._transform from optuna.distributions import CategoricalDistribution +from optuna.distributions import FloatDistribution +from optuna.distributions import IntDistribution import torch from torch import Tensor @@ -359,11 +362,45 @@ class PreferentialGPSampler(optuna.samplers.BaseSampler): ) # TODO: Make it possible to apply it on mixed search space - if all(isinstance(dist, CategoricalDistribution) for dist in search_space.values()): + def get_all_possible_params(dist: optuna.distributions.BaseDistribution) -> list[Any]: + if isinstance(dist, CategoricalDistribution): + return dist.choices + elif isinstance(dist, (IntDistribution, FloatDistribution)): + return list(np.arange(dist.low, dist.high, dist.step)) + else: + return [] + + all_possible_params = { + name: get_all_possible_params(dist) for name, dist in search_space.items() + } + + has_categorical = any( + isinstance(dist, CategoricalDistribution) for dist in search_space.values() + ) + is_all_discrete = all( + len(possible_params) > 0 for possible_params in all_possible_params.values() + ) + can_evaluate_all = ( + is_all_discrete + and np.prod( + [len(possible_params) for possible_params in all_possible_params.values()] + ) + <= 1e6 + ) + print(has_categorical, is_all_discrete, can_evaluate_all) + + if has_categorical and not can_evaluate_all: + warnings.warn( + "The objective function has categorical parameters, " + "but the total search space is too large to be enumerated. " + "This may result in significantly bad performance." + ) + + if is_all_discrete and can_evaluate_all: all_param_combinations = itertools.product( *[ - [(name, choice) for choice in cast(CategoricalDistribution, dist).choices] - for name, dist in search_space.items() + [(name, choice) for choice in possible_params] + for name, possible_params in all_possible_params.items() ] ) choices = torch.tensor( @@ -395,6 +432,10 @@ class PreferentialGPSampler(optuna.samplers.BaseSampler): param_name: str, param_distribution: optuna.distributions.BaseDistribution, ) -> Any: + warnings.warn( + f"Dynamic search space detected. Falling back to {self.independent_sampler}." + ) + return self.independent_sampler.sample_independent( study, trial, param_name, param_distribution ) From 5a4f1ed195f3c9d4f4a724defaa4b23bce951e5a Mon Sep 17 00:00:00 2001 From: Contramundum Date: Fri, 29 Sep 2023 12:40:30 +0900 Subject: [PATCH 07/24] Fix linter --- optuna_dashboard/preferential/samplers/gp.py | 1 - 1 file changed, 1 deletion(-) diff --git a/optuna_dashboard/preferential/samplers/gp.py b/optuna_dashboard/preferential/samplers/gp.py index f4415a38..7e5e2c35 100644 --- a/optuna_dashboard/preferential/samplers/gp.py +++ b/optuna_dashboard/preferential/samplers/gp.py @@ -4,7 +4,6 @@ import itertools import math from typing import Any from typing import Callable -from typing import cast import warnings import botorch.acquisition.analytic From 483b49bb9755c51c14230220f2948a0359b4b212 Mon Sep 17 00:00:00 2001 From: Contramundum Date: Fri, 29 Sep 2023 12:42:43 +0900 Subject: [PATCH 08/24] Fix mypy --- optuna_dashboard/preferential/samplers/gp.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/optuna_dashboard/preferential/samplers/gp.py b/optuna_dashboard/preferential/samplers/gp.py index 7e5e2c35..44427d9b 100644 --- a/optuna_dashboard/preferential/samplers/gp.py +++ b/optuna_dashboard/preferential/samplers/gp.py @@ -363,7 +363,7 @@ class PreferentialGPSampler(optuna.samplers.BaseSampler): # TODO: Make it possible to apply it on mixed search space def get_all_possible_params(dist: optuna.distributions.BaseDistribution) -> list[Any]: if isinstance(dist, CategoricalDistribution): - return dist.choices + return list(dist.choices) elif isinstance(dist, (IntDistribution, FloatDistribution)): return list(np.arange(dist.low, dist.high, dist.step)) else: From 87c41cbf0d1f40405a80a5ced84142165cf502fa Mon Sep 17 00:00:00 2001 From: Contramundum Date: Fri, 29 Sep 2023 17:37:49 +0900 Subject: [PATCH 09/24] Fix warning message --- optuna_dashboard/preferential/samplers/gp.py | 44 ++++++++++++-------- 1 file changed, 27 insertions(+), 17 deletions(-) diff --git a/optuna_dashboard/preferential/samplers/gp.py b/optuna_dashboard/preferential/samplers/gp.py index 44427d9b..cb03fb8f 100644 --- a/optuna_dashboard/preferential/samplers/gp.py +++ b/optuna_dashboard/preferential/samplers/gp.py @@ -373,27 +373,36 @@ class PreferentialGPSampler(optuna.samplers.BaseSampler): name: get_all_possible_params(dist) for name, dist in search_space.items() } - has_categorical = any( - isinstance(dist, CategoricalDistribution) for dist in search_space.values() - ) is_all_discrete = all( len(possible_params) > 0 for possible_params in all_possible_params.values() ) - can_evaluate_all = ( - is_all_discrete - and np.prod( - [len(possible_params) for possible_params in all_possible_params.values()] - ) - <= 1e6 + search_space_size = np.prod( + [len(possible_params) for possible_params in all_possible_params.values()] ) - print(has_categorical, is_all_discrete, can_evaluate_all) + # TODO(contramundum53): Fix this arbitrarily chosen limit. + size_limit = 1e6 + can_evaluate_all = is_all_discrete and search_space_size <= size_limit - if has_categorical and not can_evaluate_all: - warnings.warn( - "The objective function has categorical parameters, " - "but the total search space is too large to be enumerated. " - "This may result in significantly bad performance." - ) + if ( + any(isinstance(dist, CategoricalDistribution) for dist in search_space.values()) + and not can_evaluate_all + ): + if is_all_discrete: + warnings.warn( + "The objective function has categorical parameters, " + "but the total search space is too large to be enumerated. " + f"(Search space size: {search_space_size} > limit: {size_limit})" + "This may result in significantly bad performance." + ) + else: + warnings.warn( + "The objective function has categorical parameters, " + "but the search space cannot be enumerated because " + "it also contains continuous parameters. " + "This may result in significantly bad performance. " + "You can work around this problem by specifying 'step' " + "in each continuous parameter." + ) if is_all_discrete and can_evaluate_all: all_param_combinations = itertools.product( @@ -432,7 +441,8 @@ class PreferentialGPSampler(optuna.samplers.BaseSampler): param_distribution: optuna.distributions.BaseDistribution, ) -> Any: warnings.warn( - f"Dynamic search space detected. Falling back to {self.independent_sampler}." + "Dynamic search space detected. " + f"Falling back to {self.independent_sampler.__class__.__name__}." ) return self.independent_sampler.sample_independent( From 9c29c24bd4713eff0224eaf35a31a57b16e5f59c Mon Sep 17 00:00:00 2001 From: Contramundum Date: Fri, 29 Sep 2023 18:24:19 +0900 Subject: [PATCH 10/24] Fix typo --- optuna_dashboard/ts/components/AppDrawer.tsx | 4 ++-- optuna_dashboard/ts/components/StudyDetail.tsx | 4 ++-- optuna_dashboard/ts/state.ts | 2 +- 3 files changed, 5 insertions(+), 5 deletions(-) diff --git a/optuna_dashboard/ts/components/AppDrawer.tsx b/optuna_dashboard/ts/components/AppDrawer.tsx index 8740c68a..31eeec69 100644 --- a/optuna_dashboard/ts/components/AppDrawer.tsx +++ b/optuna_dashboard/ts/components/AppDrawer.tsx @@ -17,7 +17,7 @@ import ListItemText from "@mui/material/ListItemText" import { drawerOpenState, reloadIntervalState, - useStudyIsPreferencial, + useStudyIsPreferential, } from "../state" import { Link } from "react-router-dom" import AutoGraphIcon from "@mui/icons-material/AutoGraph" @@ -130,7 +130,7 @@ export const AppDrawer: FC<{ const [open, setOpen] = useRecoilState(drawerOpenState) const reloadInterval = useRecoilValue(reloadIntervalState) const isPreferential = - studyId !== undefined ? useStudyIsPreferencial(studyId) : null + studyId !== undefined ? useStudyIsPreferential(studyId) : null const styleListItem = { display: "block", diff --git a/optuna_dashboard/ts/components/StudyDetail.tsx b/optuna_dashboard/ts/components/StudyDetail.tsx index ab37d02b..1a8ced33 100644 --- a/optuna_dashboard/ts/components/StudyDetail.tsx +++ b/optuna_dashboard/ts/components/StudyDetail.tsx @@ -18,7 +18,7 @@ import { actionCreator } from "../action" import { reloadIntervalState, useStudyDetailValue, - useStudyIsPreferencial, + useStudyIsPreferential, useStudyName, } from "../state" import { TrialTable } from "./TrialTable" @@ -54,7 +54,7 @@ export const StudyDetail: FC<{ const studyDetail = useStudyDetailValue(studyId) const reloadInterval = useRecoilValue(reloadIntervalState) const studyName = useStudyName(studyId) - const isPreferential = useStudyIsPreferencial(studyId) + const isPreferential = useStudyIsPreferential(studyId) const title = studyName !== null ? `${studyName} (id=${studyId})` : `Study #${studyId}` diff --git a/optuna_dashboard/ts/state.ts b/optuna_dashboard/ts/state.ts index 98100d49..3c6f0fd7 100644 --- a/optuna_dashboard/ts/state.ts +++ b/optuna_dashboard/ts/state.ts @@ -87,7 +87,7 @@ export const useStudyDirections = ( return studyDetail?.directions || studySummary?.directions || null } -export const useStudyIsPreferencial = (studyId: number): boolean | null => { +export const useStudyIsPreferential = (studyId: number): boolean | null => { const studyDetail = useStudyDetailValue(studyId) const studySummary = useStudySummaryValue(studyId) return studyDetail?.is_preferential || studySummary?.is_preferential || null From cc1c31c9a929df2754c09ef985fa05fbbb90329c Mon Sep 17 00:00:00 2001 From: keisuke umezawa Date: Sat, 30 Sep 2023 13:07:06 +0900 Subject: [PATCH 11/24] Add contribution-welcome and good-first-issue for exempt-issue-labels in stale.yml --- .github/workflows/stale.yml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.github/workflows/stale.yml b/.github/workflows/stale.yml index 821fe15f..bb820455 100644 --- a/.github/workflows/stale.yml +++ b/.github/workflows/stale.yml @@ -24,6 +24,6 @@ jobs: days-before-pr-close: 7 # default number stale-issue-label: 'stale' stale-pr-label: 'stale' - exempt-issue-labels: 'no-stale' + exempt-issue-labels: 'no-stale,good-first-issue,contribution-welcome' exempt-pr-labels: 'no-stale' operations-per-run: 1000 From e128d0fd8d98e885f807db0d3a9dfcdc969cc6e2 Mon Sep 17 00:00:00 2001 From: YuigaWada Date: Sat, 30 Sep 2023 13:29:02 +0900 Subject: [PATCH 12/24] Fix `isSupportedSchema` --- standalone_app/src/sqlite3.ts | 7 +++---- 1 file changed, 3 insertions(+), 4 deletions(-) diff --git a/standalone_app/src/sqlite3.ts b/standalone_app/src/sqlite3.ts index 00570f9a..e9b5c963 100644 --- a/standalone_app/src/sqlite3.ts +++ b/standalone_app/src/sqlite3.ts @@ -51,13 +51,12 @@ export const loadStorage = ( const isSupportedSchema = (db: SQLite3DB): boolean => { let supported = true + let supportedVersions = ["v3.2.0.a"] db.exec({ - sql: "SELECT schema_version FROM version_info LIMIT 1", + sql: "SELECT version_num FROM alembic_version LIMIT 1", // eslint-disable-next-line @typescript-eslint/no-explicit-any callback: (vals: any[]) => { - if (vals[0] != 12) { - supported = false - } + supported = supportedVersions.includes(vals[0]) }, }) return supported From 4250c59203359d31c8d3691b8677cf6494a5f109 Mon Sep 17 00:00:00 2001 From: keisuke-umezawa Date: Sat, 30 Sep 2023 14:00:45 +0900 Subject: [PATCH 13/24] Follow review comments --- optuna_dashboard/ts/components/GraphParallelCoordinate.tsx | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/optuna_dashboard/ts/components/GraphParallelCoordinate.tsx b/optuna_dashboard/ts/components/GraphParallelCoordinate.tsx index 5d5b1222..71f2b32d 100644 --- a/optuna_dashboard/ts/components/GraphParallelCoordinate.tsx +++ b/optuna_dashboard/ts/components/GraphParallelCoordinate.tsx @@ -175,7 +175,7 @@ const plotCoordinate = ( }) const minValue = Math.min(...logValues) const maxValue = Math.max(...logValues) - const range = [minValue, maxValue] + const range = [Math.floor(minValue), Math.ceil(maxValue)] const tickvals = Array.from( { length: Math.ceil(maxValue) - Math.floor(minValue) + 1 }, (_, i) => i + Math.floor(minValue) @@ -226,7 +226,7 @@ const plotCoordinate = ( ticktext, } } else { - // numerical and non-log + // numerical and linear return { label: breakLabelIfTooLong(s.name), values: values, From 9fc8932f9ab232f371fdc1dd55d383a219801940 Mon Sep 17 00:00:00 2001 From: YuigaWada Date: Sat, 30 Sep 2023 14:19:39 +0900 Subject: [PATCH 14/24] Modified the logic for checking supported schema --- standalone_app/src/sqlite3.ts | 26 +++++++++++++++++++------- 1 file changed, 19 insertions(+), 7 deletions(-) diff --git a/standalone_app/src/sqlite3.ts b/standalone_app/src/sqlite3.ts index e9b5c963..2372f162 100644 --- a/standalone_app/src/sqlite3.ts +++ b/standalone_app/src/sqlite3.ts @@ -38,7 +38,8 @@ export const loadStorage = ( ) db.checkRc(rc) try { - if (!isSupportedSchema(db)) { + const schemaVersion = getSchemaVersion(db) + if (!isSupportedSchema(schemaVersion)) { return } const studies = getStudies(db) @@ -49,20 +50,31 @@ export const loadStorage = ( }) } -const isSupportedSchema = (db: SQLite3DB): boolean => { - let supported = true - let supportedVersions = ["v3.2.0.a"] +const getSchemaVersion = (db: SQLite3DB): string => { + let schemaVersion = "" db.exec({ sql: "SELECT version_num FROM alembic_version LIMIT 1", // eslint-disable-next-line @typescript-eslint/no-explicit-any callback: (vals: any[]) => { - supported = supportedVersions.includes(vals[0]) + schemaVersion = vals[0] }, }) - return supported + return schemaVersion } -const getStudies = (db: SQLite3DB): Study[] => { +const isSupportedSchema = (schemaVersion: string): boolean => { + let lowestVersion = "v3.0.0.d" // OK: "v3.2.0.a", "v3.0.0.d" + if (schemaVersion == lowestVersion) return true + return isGreaterSchemaVersion(schemaVersion,lowestVersion) +} + +const isGreaterSchemaVersion = (leftVersion: string, rightVersion: string): boolean => { // return leftVersion > rightVersion + leftVersion = leftVersion.replace(/\D/g,'') + rightVersion = rightVersion.replace(/\D/g,'') + return Number(leftVersion) > Number(rightVersion) +} + +const getStudies = (db: SQLite3DB, schemaVersion: string): Study[] => { const studies: Study[] = [] db.exec({ sql: From 6ca09a2ff489c536a1848898a0e48763da885030 Mon Sep 17 00:00:00 2001 From: YuigaWada Date: Sat, 30 Sep 2023 14:35:04 +0900 Subject: [PATCH 15/24] add support for loading v2.6.0.a --- standalone_app/src/sqlite3.ts | 120 +++++++++++++++++++++------------- 1 file changed, 76 insertions(+), 44 deletions(-) diff --git a/standalone_app/src/sqlite3.ts b/standalone_app/src/sqlite3.ts index 2372f162..c27a2117 100644 --- a/standalone_app/src/sqlite3.ts +++ b/standalone_app/src/sqlite3.ts @@ -42,7 +42,7 @@ export const loadStorage = ( if (!isSupportedSchema(schemaVersion)) { return } - const studies = getStudies(db) + const studies = getStudies(db, schemaVersion) setter((prev) => [...prev, ...studies]) } finally { db.close() @@ -63,7 +63,7 @@ const getSchemaVersion = (db: SQLite3DB): string => { } const isSupportedSchema = (schemaVersion: string): boolean => { - let lowestVersion = "v3.0.0.d" // OK: "v3.2.0.a", "v3.0.0.d" + let lowestVersion = "v2.6.0.a" // OK: "v3.2.0.a", "v3.0.0.d", "v2.6.0.a" if (schemaVersion == lowestVersion) return true return isGreaterSchemaVersion(schemaVersion,lowestVersion) } @@ -89,7 +89,7 @@ const getStudies = (db: SQLite3DB, schemaVersion: string): Study[] => { vals[2] === "MINIMIZE" ? "minimize" : "maximize" const objective = vals[3] - const trials = getTrials(db, studyId) + const trials = getTrials(db, studyId, schemaVersion) const union_search_space: SearchSpaceItem[] = [] const union_user_attrs: AttributeSpec[] = [] let intersection_search_space: Set = new Set() @@ -147,7 +147,7 @@ const getStudies = (db: SQLite3DB, schemaVersion: string): Study[] => { return studies } -const getTrials = (db: SQLite3DB, studyId: number): Trial[] => { +const getTrials = (db: SQLite3DB, studyId: number, schemaVersion: string): Trial[] => { const trials: Trial[] = [] db.exec({ sql: @@ -171,8 +171,8 @@ const getTrials = (db: SQLite3DB, studyId: number): Trial[] => { number: vals[1], study_id: studyId, state: state, - values: getTrialValues(db, trialId), - intermediate_values: getTrialIntermediateValues(db, trialId), + values: getTrialValues(db, trialId, schemaVersion), + intermediate_values: getTrialIntermediateValues(db, trialId, schemaVersion), params: [], // Set this column later user_attrs: [], // Set this column later datetime_start: vals[3], @@ -184,24 +184,37 @@ const getTrials = (db: SQLite3DB, studyId: number): Trial[] => { return trials } -const getTrialValues = (db: SQLite3DB, trialId: number): TrialValueNumber[] => { +const getTrialValues = (db: SQLite3DB, trialId: number, schemaVersion: string): TrialValueNumber[] => { const values: TrialValueNumber[] = [] - db.exec({ - sql: - "SELECT value, value_type" + - ` FROM trial_values WHERE trial_id = ${trialId}` + - " ORDER BY objective", - // eslint-disable-next-line @typescript-eslint/no-explicit-any - callback: (vals: any[]) => { - values.push( - vals[1] === "INF_NEG" - ? "-inf" - : vals[1] === "INF_POS" - ? "+inf" - : vals[0] - ) - }, - }) + if (isGreaterSchemaVersion(schemaVersion,"v2.6.0.a")) { + db.exec({ + sql: + "SELECT value, value_type" + + ` FROM trial_values WHERE trial_id = ${trialId}` + + " ORDER BY objective", + // eslint-disable-next-line @typescript-eslint/no-explicit-any + callback: (vals: any[]) => { + values.push( + vals[1] === "INF_NEG" + ? "-inf" + : vals[1] === "INF_POS" + ? "+inf" + : vals[0] + ) + }, + }) + } else { + db.exec({ + sql: + "SELECT value" + + ` FROM trial_values WHERE trial_id = ${trialId}` + + " ORDER BY objective", + // eslint-disable-next-line @typescript-eslint/no-explicit-any + callback: (vals: any[]) => { + values.push(vals[0]) + }, + }) + } return values } @@ -244,7 +257,9 @@ const paramInternalValueToExternalValue = ( const parseDistributionJSON = (t: string): Distribution => { const parsed = JSON.parse(t) - if (parsed.name === "FloatDistribution") { + const floatDistributionList = ["FloatDistribution", "UniformDistribution", "LogUniformDistribution", "DiscreteUniformDistribution"] + const intDistributionList = ["IntDistribution", "IntUniformDistribution", "IntLogUniformDistribution"] + if (floatDistributionList.includes(parsed.name)) { return { type: "FloatDistribution", low: parsed.attributes.low as number, @@ -252,7 +267,7 @@ const parseDistributionJSON = (t: string): Distribution => { step: parsed.attributes.step as number, log: parsed.attributes.log as boolean, } - } else if (parsed.name === "IntDistribution") { + } else if (intDistributionList.includes(parsed.name)) { return { type: "IntDistribution", low: parsed.attributes.low as number, @@ -298,26 +313,43 @@ const getTrialUserAttributes = ( const getTrialIntermediateValues = ( db: SQLite3DB, - trialId: number + trialId: number, + schemaVersion: string ): TrialIntermediateValue[] => { const values: TrialIntermediateValue[] = [] - db.exec({ - sql: - "SELECT step, intermediate_value, intermediate_value_type" + - ` FROM trial_intermediate_values WHERE trial_id = ${trialId}` + - " ORDER BY step", - // eslint-disable-next-line @typescript-eslint/no-explicit-any - callback: (vals: any[]) => { - values.push({ - step: vals[0], - value: - vals[2] === "INF_NEG" - ? "-inf" - : vals[2] === "INF_POS" - ? "+inf" - : vals[1], - }) - }, - }) + if (isGreaterSchemaVersion(schemaVersion,"v2.6.0.a")) { + db.exec({ + sql: + "SELECT step, intermediate_value, intermediate_value_type" + + ` FROM trial_intermediate_values WHERE trial_id = ${trialId}` + + " ORDER BY step", + // eslint-disable-next-line @typescript-eslint/no-explicit-any + callback: (vals: any[]) => { + values.push({ + step: vals[0], + value: + vals[2] === "INF_NEG" + ? "-inf" + : vals[2] === "INF_POS" + ? "+inf" + : vals[1], // TODO: NANの対応 + }) + }, + }) + } else { + db.exec({ + sql: + "SELECT step, intermediate_value" + + ` FROM trial_intermediate_values WHERE trial_id = ${trialId}` + + " ORDER BY step", + // eslint-disable-next-line @typescript-eslint/no-explicit-any + callback: (vals: any[]) => { + values.push({ + step: vals[0], + value: vals[1], + }) + }, + }) + } return values } From 61f9da432ec305775f9b2ea759bf91fb54f06e46 Mon Sep 17 00:00:00 2001 From: YuigaWada Date: Sat, 30 Sep 2023 15:05:41 +0900 Subject: [PATCH 16/24] Add a missing option for `getTrialIntermediateValues` --- standalone_app/src/sqlite3.ts | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/standalone_app/src/sqlite3.ts b/standalone_app/src/sqlite3.ts index c27a2117..2eb92a33 100644 --- a/standalone_app/src/sqlite3.ts +++ b/standalone_app/src/sqlite3.ts @@ -332,7 +332,9 @@ const getTrialIntermediateValues = ( ? "-inf" : vals[2] === "INF_POS" ? "+inf" - : vals[1], // TODO: NANの対応 + : vals[2] === "NAN" + ? "nan" + : vals[1] }) }, }) From 4fa25dfd54cb0f3c14c7549f42b78d6f94e89e7b Mon Sep 17 00:00:00 2001 From: YuigaWada Date: Sat, 30 Sep 2023 15:22:16 +0900 Subject: [PATCH 17/24] Format `standalone_app/src/sqlite3.ts` --- standalone_app/src/sqlite3.ts | 51 ++++++++++++++++++++++++++--------- 1 file changed, 38 insertions(+), 13 deletions(-) diff --git a/standalone_app/src/sqlite3.ts b/standalone_app/src/sqlite3.ts index 2eb92a33..f6516347 100644 --- a/standalone_app/src/sqlite3.ts +++ b/standalone_app/src/sqlite3.ts @@ -63,14 +63,18 @@ const getSchemaVersion = (db: SQLite3DB): string => { } const isSupportedSchema = (schemaVersion: string): boolean => { - let lowestVersion = "v2.6.0.a" // OK: "v3.2.0.a", "v3.0.0.d", "v2.6.0.a" + let lowestVersion = "v2.6.0.a" // supported: "v3.2.0.a", "v3.0.0.{a,b,c,d}", "v2.6.0.a" if (schemaVersion == lowestVersion) return true - return isGreaterSchemaVersion(schemaVersion,lowestVersion) + return isGreaterSchemaVersion(schemaVersion, lowestVersion) } -const isGreaterSchemaVersion = (leftVersion: string, rightVersion: string): boolean => { // return leftVersion > rightVersion - leftVersion = leftVersion.replace(/\D/g,'') - rightVersion = rightVersion.replace(/\D/g,'') +const isGreaterSchemaVersion = ( + leftVersion: string, + rightVersion: string +): boolean => { + // return leftVersion > rightVersion + leftVersion = leftVersion.replace(/\D/g, "") + rightVersion = rightVersion.replace(/\D/g, "") return Number(leftVersion) > Number(rightVersion) } @@ -147,7 +151,11 @@ const getStudies = (db: SQLite3DB, schemaVersion: string): Study[] => { return studies } -const getTrials = (db: SQLite3DB, studyId: number, schemaVersion: string): Trial[] => { +const getTrials = ( + db: SQLite3DB, + studyId: number, + schemaVersion: string +): Trial[] => { const trials: Trial[] = [] db.exec({ sql: @@ -172,7 +180,11 @@ const getTrials = (db: SQLite3DB, studyId: number, schemaVersion: string): Trial study_id: studyId, state: state, values: getTrialValues(db, trialId, schemaVersion), - intermediate_values: getTrialIntermediateValues(db, trialId, schemaVersion), + intermediate_values: getTrialIntermediateValues( + db, + trialId, + schemaVersion + ), params: [], // Set this column later user_attrs: [], // Set this column later datetime_start: vals[3], @@ -184,9 +196,13 @@ const getTrials = (db: SQLite3DB, studyId: number, schemaVersion: string): Trial return trials } -const getTrialValues = (db: SQLite3DB, trialId: number, schemaVersion: string): TrialValueNumber[] => { +const getTrialValues = ( + db: SQLite3DB, + trialId: number, + schemaVersion: string +): TrialValueNumber[] => { const values: TrialValueNumber[] = [] - if (isGreaterSchemaVersion(schemaVersion,"v2.6.0.a")) { + if (isGreaterSchemaVersion(schemaVersion, "v2.6.0.a")) { db.exec({ sql: "SELECT value, value_type" + @@ -257,8 +273,17 @@ const paramInternalValueToExternalValue = ( const parseDistributionJSON = (t: string): Distribution => { const parsed = JSON.parse(t) - const floatDistributionList = ["FloatDistribution", "UniformDistribution", "LogUniformDistribution", "DiscreteUniformDistribution"] - const intDistributionList = ["IntDistribution", "IntUniformDistribution", "IntLogUniformDistribution"] + const floatDistributionList = [ + "FloatDistribution", + "UniformDistribution", + "LogUniformDistribution", + "DiscreteUniformDistribution", + ] + const intDistributionList = [ + "IntDistribution", + "IntUniformDistribution", + "IntLogUniformDistribution", + ] if (floatDistributionList.includes(parsed.name)) { return { type: "FloatDistribution", @@ -317,7 +342,7 @@ const getTrialIntermediateValues = ( schemaVersion: string ): TrialIntermediateValue[] => { const values: TrialIntermediateValue[] = [] - if (isGreaterSchemaVersion(schemaVersion,"v2.6.0.a")) { + if (isGreaterSchemaVersion(schemaVersion, "v2.6.0.a")) { db.exec({ sql: "SELECT step, intermediate_value, intermediate_value_type" + @@ -334,7 +359,7 @@ const getTrialIntermediateValues = ( ? "+inf" : vals[2] === "NAN" ? "nan" - : vals[1] + : vals[1], }) }, }) From 87f2c1bbb5acd07e0a914d88e7a82cdb2ef17672 Mon Sep 17 00:00:00 2001 From: YuigaWada Date: Sat, 30 Sep 2023 15:44:32 +0900 Subject: [PATCH 18/24] Change `lowestVersions` from let to const --- standalone_app/src/sqlite3.ts | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/standalone_app/src/sqlite3.ts b/standalone_app/src/sqlite3.ts index f6516347..30bdc620 100644 --- a/standalone_app/src/sqlite3.ts +++ b/standalone_app/src/sqlite3.ts @@ -63,7 +63,7 @@ const getSchemaVersion = (db: SQLite3DB): string => { } const isSupportedSchema = (schemaVersion: string): boolean => { - let lowestVersion = "v2.6.0.a" // supported: "v3.2.0.a", "v3.0.0.{a,b,c,d}", "v2.6.0.a" + const lowestVersion = "v2.6.0.a" // supported: "v3.2.0.a", "v3.0.0.{a,b,c,d}", "v2.6.0.a" if (schemaVersion == lowestVersion) return true return isGreaterSchemaVersion(schemaVersion, lowestVersion) } From f533761b446c74ad8cce467037a6ebf9d8728898 Mon Sep 17 00:00:00 2001 From: YuigaWada Date: Sat, 30 Sep 2023 16:14:32 +0900 Subject: [PATCH 19/24] Fix: add support for v3.0.0.{a,b,c} --- standalone_app/src/sqlite3.ts | 12 +++++++++--- 1 file changed, 9 insertions(+), 3 deletions(-) diff --git a/standalone_app/src/sqlite3.ts b/standalone_app/src/sqlite3.ts index 30bdc620..c9f2c13d 100644 --- a/standalone_app/src/sqlite3.ts +++ b/standalone_app/src/sqlite3.ts @@ -73,9 +73,15 @@ const isGreaterSchemaVersion = ( rightVersion: string ): boolean => { // return leftVersion > rightVersion + const leftSuffix = leftVersion.split(".").reverse()[0] + const rightSuffix = rightVersion.split(".").reverse()[0] leftVersion = leftVersion.replace(/\D/g, "") rightVersion = rightVersion.replace(/\D/g, "") - return Number(leftVersion) > Number(rightVersion) + + const left = Number(leftVersion) + const right = Number(rightVersion) + if (left == right) return leftSuffix > rightSuffix + return left > right } const getStudies = (db: SQLite3DB, schemaVersion: string): Study[] => { @@ -202,7 +208,7 @@ const getTrialValues = ( schemaVersion: string ): TrialValueNumber[] => { const values: TrialValueNumber[] = [] - if (isGreaterSchemaVersion(schemaVersion, "v2.6.0.a")) { + if (isGreaterSchemaVersion(schemaVersion, "v3.0.0.c")) { db.exec({ sql: "SELECT value, value_type" + @@ -342,7 +348,7 @@ const getTrialIntermediateValues = ( schemaVersion: string ): TrialIntermediateValue[] => { const values: TrialIntermediateValue[] = [] - if (isGreaterSchemaVersion(schemaVersion, "v2.6.0.a")) { + if (isGreaterSchemaVersion(schemaVersion, "v3.0.0.c")) { db.exec({ sql: "SELECT step, intermediate_value, intermediate_value_type" + From 2e44f9cdb3d21d8f1cdef9854b4088c460fe4802 Mon Sep 17 00:00:00 2001 From: YuigaWada Date: Sat, 30 Sep 2023 16:45:00 +0900 Subject: [PATCH 20/24] Patch distribution management bug --- standalone_app/src/sqlite3.ts | 55 ++++++++++++++++++++++------- standalone_app/src/types/index.d.ts | 4 +-- 2 files changed, 44 insertions(+), 15 deletions(-) diff --git a/standalone_app/src/sqlite3.ts b/standalone_app/src/sqlite3.ts index c9f2c13d..3a04c163 100644 --- a/standalone_app/src/sqlite3.ts +++ b/standalone_app/src/sqlite3.ts @@ -279,18 +279,7 @@ const paramInternalValueToExternalValue = ( const parseDistributionJSON = (t: string): Distribution => { const parsed = JSON.parse(t) - const floatDistributionList = [ - "FloatDistribution", - "UniformDistribution", - "LogUniformDistribution", - "DiscreteUniformDistribution", - ] - const intDistributionList = [ - "IntDistribution", - "IntUniformDistribution", - "IntLogUniformDistribution", - ] - if (floatDistributionList.includes(parsed.name)) { + if (parsed.name === "FloatDistribution") { return { type: "FloatDistribution", low: parsed.attributes.low as number, @@ -298,7 +287,31 @@ const parseDistributionJSON = (t: string): Distribution => { step: parsed.attributes.step as number, log: parsed.attributes.log as boolean, } - } else if (intDistributionList.includes(parsed.name)) { + } else if (parsed.name === "UniformDistribution") { + return { + type: "FloatDistribution", + low: parsed.attributes.low as number, + high: parsed.attributes.high as number, + step: null, + log: false, + } + } else if (parsed.name === "LogUniformDistribution") { + return { + type: "FloatDistribution", + low: parsed.attributes.low as number, + high: parsed.attributes.high as number, + step: null, + log: true, + } + } else if (parsed.name === "DiscreteUniformDistribution") { + return { + type: "FloatDistribution", + low: parsed.attributes.low as number, + high: parsed.attributes.high as number, + step: parsed.attributes.q, + log: false, + } + } else if (parsed.name === "IntDistribution") { return { type: "IntDistribution", low: parsed.attributes.low as number, @@ -306,6 +319,22 @@ const parseDistributionJSON = (t: string): Distribution => { step: parsed.attributes.step as number, log: parsed.attributes.log as boolean, } + } else if (parsed.name === "IntUniformDistribution") { + return { + type: "IntDistribution", + low: parsed.attributes.low as number, + high: parsed.attributes.high as number, + step: parsed.attributes.step as number, + log: false, + } + } else if (parsed.name === "IntLogUniformDistribution") { + return { + type: "IntDistribution", + low: parsed.attributes.low as number, + high: parsed.attributes.high as number, + step: parsed.attributes.step as number, + log: true, + } } else { // eslint-disable-next-line @typescript-eslint/no-explicit-any const choices = parsed.attributes.choices.map((value: any) => { diff --git a/standalone_app/src/types/index.d.ts b/standalone_app/src/types/index.d.ts index 14daaaca..da00a7aa 100644 --- a/standalone_app/src/types/index.d.ts +++ b/standalone_app/src/types/index.d.ts @@ -10,7 +10,7 @@ type FloatDistribution = { type: "FloatDistribution" low: number high: number - step: number + step: number | null log: boolean } @@ -18,7 +18,7 @@ type IntDistribution = { type: "IntDistribution" low: number high: number - step: number + step: number | null log: boolean } From a9f080da1a42d29fe0d2f800f3f242f2c69ab50a Mon Sep 17 00:00:00 2001 From: YuigaWada Date: Sat, 30 Sep 2023 17:23:59 +0900 Subject: [PATCH 21/24] Add support for videos on ArtifactCardMedia --- optuna_dashboard/ts/components/ArtifactCardMedia.tsx | 12 ++++++++++++ 1 file changed, 12 insertions(+) diff --git a/optuna_dashboard/ts/components/ArtifactCardMedia.tsx b/optuna_dashboard/ts/components/ArtifactCardMedia.tsx index ebbbd751..6326cb16 100644 --- a/optuna_dashboard/ts/components/ArtifactCardMedia.tsx +++ b/optuna_dashboard/ts/components/ArtifactCardMedia.tsx @@ -21,6 +21,18 @@ export const ArtifactCardMedia: FC<{ filetype={artifact.filename.split(".").pop()} /> ) + } else if (artifact.mimetype.startsWith("video")) { + return ( + + ) } else if (artifact.mimetype.startsWith("audio")) { return (